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()