#!/usr/bin/env python """ AXModel 单图推理脚本 (不依赖 PyTorch / ONNX Runtime)。 功能:读取单张图片 → 预处理到模型尺寸 → axengine 推理 → 还原到原图尺寸 → 保存「原图 + 去噪结果」横向拼接图。 - 原图 ≤ 模型尺寸:reflect pad 到模型尺寸,推理后裁回原图 - 原图 > 模型尺寸:resize 到模型尺寸,推理后 resize 回原图 所有参数均以固定默认值写在脚本顶部,直接运行即可。 """ import os import sys import cv2 import numpy as np # ═══════════════════════════════════════════════════════════════ # 固定默认参数(按需修改) # ═══════════════════════════════════════════════════════════════ DEFAULT_AXMODEL = "./Restormer_real_denoising_224x224_sim.axmodel" DEFAULT_INPUT_IMAGE = "noisy.png" DEFAULT_OUTPUT_IMAGE = "axmodel_result.png" # 模型导出时的固定输入尺寸 (H, W),必须能被 8 整除 MODEL_HEIGHT = 224 MODEL_WIDTH = 224 # 是否保存纯去噪结果(不拼接) SAVE_RESTORED_ONLY = False RESTORED_ONLY_PATH = "axmodel_result_restored.png" # ═══════════════════════════════════════════════════════════════ def _load_image(path): """用 OpenCV 读取图片,返回 RGB uint8 numpy (H,W,3) 或 (H,W,1)。""" if not os.path.isfile(path): raise FileNotFoundError(f"输入图片不存在: {path}") img = cv2.imread(path, cv2.IMREAD_COLOR) if img is None: img = cv2.imread(path, cv2.IMREAD_GRAYSCALE) if img is None: raise RuntimeError(f"无法读取图片: {path}") return img[:, :, None] # (H,W,1) return cv2.cvtColor(img, cv2.COLOR_BGR2RGB) def _save_image(path, img): """保存 uint8 numpy 图片,自动处理灰度/彩色。""" os.makedirs(os.path.dirname(path) or ".", exist_ok=True) if img.ndim == 2: cv2.imwrite(path, img) elif img.shape[2] == 1: cv2.imwrite(path, img[:, :, 0]) else: cv2.imwrite(path, cv2.cvtColor(img, cv2.COLOR_RGB2BGR)) def _image_to_nchw(img): """uint8 (H,W,C) → float32 NCHW [0,1]。""" # return np.transpose(img.astype(np.float32) / 255.0, (2, 0, 1))[np.newaxis, ...] return np.transpose(img.astype(np.float32), (2, 0, 1))[np.newaxis, ...] def _nchw_to_image(arr): """float32 NCHW [0,1] → uint8 (H,W,C)。""" arr = arr.squeeze(0) arr = np.transpose(arr, (1, 2, 0)) arr = np.clip(arr, 0.0, 1.0) return (arr * 255.0).round().clip(0, 255).astype(np.uint8) def _prepare_for_model(tensor, model_h, model_w): """ 将 NCHW tensor 处理到模型固定尺寸: - 原图 ≤ 模型尺寸:reflect pad → 返回 (padded, (orig_h, orig_w), "pad") - 原图 > 模型尺寸:resize → 返回 (resized, (orig_h, orig_w), "resize") """ _, c, orig_h, orig_w = tensor.shape if orig_h <= model_h and orig_w <= model_w: pad_h = model_h - orig_h pad_w = model_w - orig_w if pad_h or pad_w: tensor = np.pad(tensor, ((0, 0), (0, 0), (0, pad_h), (0, pad_w)), mode="reflect") return tensor.astype(np.uint8), (orig_h, orig_w), "pad" else: resized = np.zeros((1, c, model_h, model_w), dtype=np.float32) for b in range(tensor.shape[0]): for ch in range(c): resized[b, ch] = cv2.resize(tensor[b, ch], (model_w, model_h), interpolation=cv2.INTER_LINEAR) return resized.astype(np.uint8), (orig_h, orig_w), "resize" def _restore_from_model(output_arr, orig_h, orig_w, mode): """ 将模型输出还原到原图尺寸: - mode="pad":裁掉 padding 区域 - mode="resize":resize 回原图尺寸 """ if mode == "pad": return output_arr[:, :, :orig_h, :orig_w] else: _, c, _, _ = output_arr.shape restored = np.zeros((1, c, orig_h, orig_w), dtype=np.float32) for b in range(output_arr.shape[0]): for ch in range(c): restored[b, ch] = cv2.resize(output_arr[b, ch], (orig_w, orig_h), interpolation=cv2.INTER_LINEAR) return restored def main(): # ----- 命令行可覆盖前三个参数(可选) ----- axmodel_path = sys.argv[1] if len(sys.argv) > 1 else DEFAULT_AXMODEL input_path = sys.argv[2] if len(sys.argv) > 2 else DEFAULT_INPUT_IMAGE output_path = sys.argv[3] if len(sys.argv) > 3 else DEFAULT_OUTPUT_IMAGE if MODEL_HEIGHT % 8 != 0 or MODEL_WIDTH % 8 != 0: raise ValueError(f"MODEL_HEIGHT({MODEL_HEIGHT}) 和 MODEL_WIDTH({MODEL_WIDTH}) 必须能被 8 整除") # ----- 加载 axengine ----- try: import axengine as axe except ImportError: raise ImportError("请先安装 axengine") session = axe.InferenceSession(axmodel_path, providers=["AxEngineExecutionProvider"]) input_name = session.get_inputs()[0].name output_name = session.get_outputs()[0].name # ----- 读取原图 ----- original_img = _load_image(input_path) # uint8 HWC RGB orig_h, orig_w = original_img.shape[:2] is_gray = original_img.ndim == 2 or original_img.shape[2] == 1 # ----- 预处理到模型尺寸 ----- tensor = _image_to_nchw(original_img) tensor, (crop_h, crop_w), prep_mode = _prepare_for_model(tensor, MODEL_HEIGHT, MODEL_WIDTH) print(f"[AXModel 推理] 模型: {axmodel_path}") print(f"[AXModel 推理] 输入图片: {input_path} (原图尺寸: {orig_h}x{orig_w}, 模型输入尺寸: {MODEL_HEIGHT}x{MODEL_WIDTH}, 预处理模式: {prep_mode})") # ----- axengine 推理 ----- ax_output = session.run([output_name], {input_name: tensor})[0] # ----- 后处理:还原到原图尺寸 ----- restored_tensor = _restore_from_model(ax_output, crop_h, crop_w, prep_mode) restored_img = _nchw_to_image(restored_tensor) # uint8 HWC RGB # ----- 拼接:原图 || 去噪结果 ----- if is_gray: orig_display = original_img[:, :, 0] if original_img.ndim == 3 else original_img rest_display = restored_img[:, :, 0] if restored_img.ndim == 3 else restored_img comparison = np.concatenate([orig_display, rest_display], axis=1) else: comparison = np.concatenate([original_img, restored_img], axis=1) # ----- 保存结果 ----- _save_image(output_path, comparison) print(f"[AXModel 推理] 拼接结果已保存到: {output_path} (左=原图, 右=去噪)") if SAVE_RESTORED_ONLY: _save_image(RESTORED_ONLY_PATH, restored_img) print(f"[AXModel 推理] 纯去噪结果已保存到: {RESTORED_ONLY_PATH}") if __name__ == "__main__": main()