File size: 5,223 Bytes
1266aec | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 | #!/usr/bin/env python3
import os
import cv2
import numpy as np
import axengine as axe
# ============================================================
# 默认参数(可按需修改)
# ============================================================
AXMODEL_PATH = "dncnn_color_blind_416x416_sim.axmodel"
OUTPUT_DIR = "./"
# AX 量化模型输入格式: "float32" 或 "uint8",取决于编译时的量化配置
INPUT_DTYPE = "uint8" # 量化模型通常为 uint8;float 模型设为 "float32"
NOISE_SIGMA = 25 # 对原图施加的 AWGN 噪声强度(uint8 尺度,设为 0 则不加噪)
# ============================================================
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)
# ---- 加载 AX 模型 ----
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)
# ---- AX 推理 ----
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()
|