File size: 3,060 Bytes
525e655 | 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 | import os
import numpy as np
import axengine as axe
from PIL import Image, ImageDraw, ImageFont
# ============================================================
# Default parameters (fixed values)
# ============================================================
MODEL_PATH = "./SCI_TPAMI_600_400.axmodel"
IMAGE_PATH = "./pic/00001.png"
INPUT_WIDTH = 600
INPUT_HEIGHT = 400
OUTPUT_PATH = "./axmodel_res.png"
# ============================================================
# Label text
# ============================================================
LEFT_LABEL = "Original"
RIGHT_LABEL = "SCI++ Enhanced (AXEngine)"
def load_and_preprocess(image_path, width, height):
"""Load original image (full size) and a resized copy for model input (uint8)."""
original = Image.open(image_path).convert("RGB")
resized = original.resize((width, height), Image.BICUBIC)
img_np = np.asarray(resized, dtype=np.uint8)
img_nchw = np.transpose(img_np, (2, 0, 1))[None, ...] # NCHW, uint8
return img_nchw, original
def run_axmodel(session, input_np):
input_name = session.get_inputs()[0].name
return session.run(None, {input_name: input_np})[0]
def tensor_to_image(tensor):
"""NCHW (0–1 range float) → PIL Image."""
img_np = np.clip(tensor[0], 0.0, 1.0)
img_np = np.transpose(img_np, (1, 2, 0))
return Image.fromarray((img_np * 255.0).round().astype(np.uint8))
def stitch_comparison(original, enhanced, output_path):
"""Horizontally stitch original + enhanced, add labels above."""
ow, oh = original.size
ew, eh = enhanced.size
assert (ow, oh) == (ew, eh), f"Size mismatch: original={original.size} enhanced={enhanced.size}"
canvas_w = ow * 2
canvas_h = oh + 40
canvas = Image.new("RGB", (canvas_w, canvas_h), (255, 255, 255))
canvas.paste(original, (0, 40))
canvas.paste(enhanced, (ow, 40))
draw = ImageDraw.Draw(canvas)
try:
font = ImageFont.truetype("/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", 18)
except Exception:
font = ImageFont.load_default()
draw.text((ow // 2, 8), LEFT_LABEL, fill=(0, 0, 0), font=font, anchor="mt")
draw.text((ow + ow // 2, 8), RIGHT_LABEL, fill=(0, 0, 0), font=font, anchor="mt")
os.makedirs(os.path.dirname(output_path) or ".", exist_ok=True)
canvas.save(output_path)
print(f"Saved comparison → {output_path}")
def main():
# load
input_np, original_img = load_and_preprocess(IMAGE_PATH, INPUT_WIDTH, INPUT_HEIGHT)
original_size = original_img.size # (W, H)
# infer
session = axe.InferenceSession(MODEL_PATH, providers=["AxEngineExecutionProvider"])
output = run_axmodel(session, input_np)
# restore to original size then stitch
enhanced_img = tensor_to_image(output)
enhanced_img = enhanced_img.resize(original_size, Image.BICUBIC)
stitch_comparison(original_img, enhanced_img, OUTPUT_PATH)
print(f"Input: {original_size} → model input: ({INPUT_WIDTH},{INPUT_HEIGHT}) → restored: {original_size}")
if __name__ == "__main__":
main()
|