| |
| import os |
| import cv2 |
| import numpy as np |
| import axengine as axe |
|
|
| |
| |
| |
| AXMODEL_PATH = "dncnn_color_blind_416x416_sim.axmodel" |
| OUTPUT_DIR = "./" |
| |
| INPUT_DTYPE = "uint8" |
| NOISE_SIGMA = 25 |
|
|
| |
|
|
|
|
| def load_image(img_path): |
| """读取图像,返回 BGR uint8 原始图""" |
| img_bgr = cv2.imread(img_path, cv2.IMREAD_UNCHANGED) |
| if img_bgr is None: |
| raise FileNotFoundError(f"无法读取图像: {img_path}") |
| if img_bgr.ndim == 2: |
| img_bgr = cv2.cvtColor(img_bgr, cv2.COLOR_GRAY2BGR) |
| return img_bgr |
|
|
|
|
| def add_awgn(img_bgr, sigma): |
| """对 BGR uint8 图像施加 AWGN,返回加噪后的 BGR uint8 图像""" |
| noise = np.random.randn(*img_bgr.shape).astype(np.float32) * sigma |
| noisy = img_bgr.astype(np.float32) + noise |
| noisy = np.clip(noisy, 0, 255).astype(np.uint8) |
| return noisy |
|
|
|
|
| def preprocess(img_bgr, model_h, model_w): |
| """ |
| 预处理: BGR->RGB, resize to model size, 按 INPUT_DTYPE 做归一化或保持 uint8 |
| 返回: 模型输入张量, 原始尺寸 (h, w) |
| """ |
| orig_h, orig_w = img_bgr.shape[:2] |
| img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB) |
| img_resized = cv2.resize(img_rgb, (model_w, model_h), interpolation=cv2.INTER_LINEAR) |
|
|
| if INPUT_DTYPE == "float32": |
| img_float = img_resized.astype(np.float32) / 255.0 |
| tensor = np.transpose(img_float, (2, 0, 1))[np.newaxis, ...].astype(np.float32) |
| else: |
| tensor = np.transpose(img_resized, (2, 0, 1))[np.newaxis, ...].astype(np.uint8) |
| return tensor, (orig_h, orig_w) |
|
|
|
|
| def postprocess(tensor, orig_h, orig_w): |
| """ |
| 后处理: squeeze batch, CHW->HWC, clip, uint8, resize 回原始尺寸 |
| 返回: BGR uint8 图像 |
| """ |
| if tensor.dtype == np.uint8: |
| arr = np.squeeze(tensor, axis=0) |
| arr = np.transpose(arr, (1, 2, 0)) |
| arr_uint = arr |
| else: |
| arr = np.clip(np.squeeze(tensor, axis=0), 0.0, 1.0) |
| arr = np.transpose(arr, (1, 2, 0)) |
| arr_uint = (arr * 255.0).round().astype(np.uint8) |
| arr_uint = cv2.resize(arr_uint, (orig_w, orig_h), interpolation=cv2.INTER_LINEAR) |
| result_bgr = cv2.cvtColor(arr_uint, cv2.COLOR_RGB2BGR) |
| return result_bgr |
|
|
|
|
| def make_concat(original_bgr, noisy_bgr, result_bgr, |
| label_left="Original", label_mid="Noisy", label_right="DnCNN AXModel"): |
| """左中右拼接原图、加噪图与结果图,并添加顶部标注""" |
| imgs = [original_bgr, noisy_bgr, result_bgr] |
| labels = [label_left, label_mid, label_right] |
| hs = [im.shape[0] for im in imgs] |
| h = max(hs) |
|
|
| resized = [] |
| for im in imgs: |
| hh, ww = im.shape[:2] |
| if hh != h: |
| im = cv2.resize(im, (int(ww * h / hh), h), interpolation=cv2.INTER_LINEAR) |
| resized.append(im) |
|
|
| ws = [im.shape[1] for im in resized] |
| total_w = sum(ws) |
| label_h = max(30, h // 25) |
| canvas = np.full((h + label_h, total_w, 3), 255, dtype=np.uint8) |
|
|
| x = 0 |
| for idx, im in enumerate(resized): |
| canvas[label_h:, x:x + ws[idx]] = im |
| x += ws[idx] |
|
|
| font = cv2.FONT_HERSHEY_SIMPLEX |
| font_scale = label_h / 30.0 |
| thickness = max(1, int(font_scale)) |
| color = (0, 0, 0) |
|
|
| x = 0 |
| for idx, (label, w) in enumerate(zip(labels, ws)): |
| (tw, th), _ = cv2.getTextSize(label, font, font_scale, thickness) |
| cv2.putText(canvas, label, (x + w // 2 - tw // 2, label_h - (label_h - th) // 2), |
| font, font_scale, color, thickness, cv2.LINE_AA) |
| x += w |
|
|
| return canvas |
|
|
|
|
| def main(): |
| img_path = './3096.png' |
| os.makedirs(OUTPUT_DIR, exist_ok=True) |
|
|
| |
| session = axe.InferenceSession(AXMODEL_PATH, providers=["AxEngineExecutionProvider"]) |
| input_name = session.get_inputs()[0].name |
| _, _, model_h, model_w = session.get_inputs()[0].shape |
|
|
| |
| img_bgr = load_image(img_path) |
|
|
| |
| if NOISE_SIGMA > 0: |
| noisy_bgr = add_awgn(img_bgr, NOISE_SIGMA) |
| else: |
| noisy_bgr = img_bgr |
|
|
| |
| tensor, (orig_h, orig_w) = preprocess(noisy_bgr, model_h, model_w) |
|
|
| |
| result = session.run(None, {input_name: tensor})[0] |
|
|
| |
| result_bgr = postprocess(result, orig_h, orig_w) |
|
|
| |
| label_mid = f"Noisy (sigma={NOISE_SIGMA})" if NOISE_SIGMA > 0 else "Input" |
| concat = make_concat(img_bgr, noisy_bgr, result_bgr, label_mid=label_mid) |
|
|
| name = os.path.splitext(os.path.basename(img_path))[0] |
| out_path = os.path.join(OUTPUT_DIR, "axmodel_res.png") |
| cv2.imwrite(out_path, concat) |
| print(f"输出已保存: {out_path}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|