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