| import os |
|
|
| import numpy as np |
| import axengine as axe |
| from PIL import Image, ImageDraw, ImageFont |
|
|
| |
| |
| |
| MODEL_PATH = "./SCI_TPAMI_600_400.axmodel" |
| IMAGE_PATH = "./pic/00001.png" |
| INPUT_WIDTH = 600 |
| INPUT_HEIGHT = 400 |
| OUTPUT_PATH = "./axmodel_res.png" |
|
|
| |
| |
| |
| 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, ...] |
| 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(): |
| |
| input_np, original_img = load_and_preprocess(IMAGE_PATH, INPUT_WIDTH, INPUT_HEIGHT) |
| original_size = original_img.size |
|
|
| |
| session = axe.InferenceSession(MODEL_PATH, providers=["AxEngineExecutionProvider"]) |
| output = run_axmodel(session, input_np) |
|
|
| |
| 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() |
|
|