File size: 7,853 Bytes
e657676
3aea121
e657676
 
 
4376c3e
 
 
 
 
 
 
 
3aea121
 
60d2da1
3aea121
 
 
 
 
 
4376c3e
3aea121
e657676
 
3aea121
e657676
4376c3e
 
e657676
 
 
3aea121
4376c3e
 
 
 
 
 
 
 
3aea121
e657676
 
3aea121
e657676
3aea121
 
e657676
3aea121
 
 
e657676
 
 
 
3aea121
e657676
 
 
3aea121
4376c3e
e657676
4376c3e
 
 
3aea121
4376c3e
3aea121
 
 
e657676
3aea121
 
 
 
e657676
 
3aea121
 
 
 
 
 
 
 
bbf3023
 
 
3aea121
e657676
 
 
 
 
3aea121
 
 
 
e657676
 
3aea121
 
 
 
 
e657676
 
3aea121
 
 
 
 
 
 
 
 
 
e657676
3aea121
 
e657676
 
 
3aea121
 
 
 
bbf3023
 
3aea121
e657676
3aea121
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e657676
3aea121
 
 
 
 
 
 
 
e657676
3aea121
 
bbf3023
 
 
3aea121
 
bbf3023
 
 
3aea121
 
e657676
3aea121
 
 
 
 
 
 
 
bbf3023
 
 
 
 
 
3aea121
e657676
3aea121
 
bbf3023
3aea121
 
 
 
 
 
 
 
 
 
 
 
 
 
 
60d2da1
 
 
 
 
 
3aea121
 
e657676
3aea121
 
e657676
3aea121
 
e657676
 
3aea121
e657676
 
 
3aea121
60d2da1
 
 
 
3aea121
 
 
 
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
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
"""Minimal inference example for google/vit-base-patch16-224 (ViT-Base) using ExecuTorch.

Loads a quantized INT8 .pte model and runs image classification on a single
image, printing the top-5 ImageNet-1k predictions.

Preprocessing: resize the shorter edge to 232px (bilinear), center-crop to
224x224, convert to RGB, scale pixel values to [0, 1], then normalize with
mean=[0.5, 0.5, 0.5] and std=[0.5, 0.5, 0.5] (the stats reported by the
checkpoint's image processor). This matches the resize-232/center-crop-224
recipe used to calibrate and evaluate this model in the Arm optimization
pipeline.

Requires: executorch, torch, torchvision, pillow, json
"""

import argparse
import json
from pathlib import Path

import torch
from executorch.runtime import Runtime
from PIL import Image, ImageDraw, ImageFont
from torchvision import transforms

# ── Configuration ────────────────────────────────────────────────────────────
MODEL_PATH = "google__vit-base-patch16-224_raspberry_executorch_optimized.pte"
IMAGE_PATH = "sample_input.jpg"
LABELS_PATH = "imagenet_classes.json"
RESIZE_SIZE = 232  # shorter edge, before center crop
CROP_SIZE = 224
MEAN = (0.5, 0.5, 0.5)
STD = (0.5, 0.5, 0.5)
TOP_K = 5

_TRANSFORM = transforms.Compose(
    [
        transforms.Resize(RESIZE_SIZE),
        transforms.CenterCrop(CROP_SIZE),
        transforms.ToTensor(),
    ]
)


def load_model(model_path: str):
    """Load an ExecuTorch .pte model and return its forward method.

    The method is loaded once and reused across calls.
    """
    runtime = Runtime.get()
    program = runtime.load_program(model_path)
    return program.load_method("forward")


def load_labels(labels_path: str) -> dict[str, str]:
    """Load the ImageNet-1k index-to-label mapping."""
    with open(labels_path) as f:
        return json.load(f)


def preprocess(image_path: str) -> torch.Tensor:
    """Load an image and prepare it as a normalized [1, 3, 224, 224] tensor."""
    image = Image.open(image_path).convert("RGB")
    tensor = _TRANSFORM(image)  # [3, 224, 224] in [0, 1]

    mean = torch.tensor(MEAN, dtype=torch.float32).view(3, 1, 1)
    std = torch.tensor(STD, dtype=torch.float32).view(3, 1, 1)
    tensor = (tensor - mean) / std

    return tensor.unsqueeze(0).contiguous()  # [1, 3, H, W]


def run_inference(method, input_tensor: torch.Tensor) -> torch.Tensor:
    """Run a forward pass and return the raw logits tensor [1, num_classes]."""
    outputs = method.execute([input_tensor])
    return outputs[0]


def postprocess(logits: torch.Tensor, labels: dict[str, str]) -> list[dict]:
    """Apply softmax and return the top-k class predictions."""
    if logits.dim() == 1:
        logits = logits.unsqueeze(0)

    probabilities = torch.softmax(logits, dim=-1)
    k = min(TOP_K, probabilities.shape[-1])
    top_scores, top_indices = probabilities.topk(k, dim=-1)

    results = []
    for score, idx in zip(
        top_scores[0].tolist(), top_indices[0].tolist(), strict=False
    ):
        results.append(
            {
                "class_index": idx,
                "label": labels.get(str(idx), f"class_{idx}"),
                "probability": round(score, 6),
            }
        )
    return results


def save_results(results: list[dict], output_path: Path) -> None:
    """Save the top-k predictions as JSON next to the script."""
    with open(output_path, "w") as f:
        json.dump(results, f, indent=2)
    print(f"Predictions saved to: {output_path}")


def _load_fonts() -> tuple:
    """Load DejaVuSans fonts, falling back to the built-in bitmap font."""
    font_dir = Path("/usr/share/fonts/truetype/dejavu")
    try:
        regular = ImageFont.truetype(str(font_dir / "DejaVuSans.ttf"), 14)
        bold = ImageFont.truetype(str(font_dir / "DejaVuSans-Bold.ttf"), 15)
    except OSError:
        regular = ImageFont.load_default()
        bold = regular
    return regular, bold


def save_output_image(image_path: str, results: list[dict], output_path: Path) -> None:
    """Render the input image alongside its top-k predictions and save as sample_output.jpg.

    Layout (700 x 320 px):
      left panel: input image (280x280)
      right panel: top-5 prediction bars with class names and probabilities
    """
    MARGIN = 20
    IMG_SZ = 280
    PANEL_W = 360
    CANVAS_W = MARGIN + IMG_SZ + MARGIN + PANEL_W + MARGIN  # 700
    CANVAS_H = MARGIN + IMG_SZ + MARGIN  # 320
    TITLE_H = 28
    ROW_H = (IMG_SZ - TITLE_H - 8) // TOP_K

    BG_COLOR = "#f5f5f5"
    PANEL_BG = "#ffffff"
    TOP1_BG = "#e3f2fd"
    BAR_TOP1 = "#1565c0"
    BAR_REST = "#90caf9"
    TEXT_DARK = "#212121"
    TEXT_GRAY = "#616161"
    BORDER = "#bdbdbd"

    font, font_bold = _load_fonts()

    canvas = Image.new("RGB", (CANVAS_W, CANVAS_H), BG_COLOR)
    draw = ImageDraw.Draw(canvas)

    # Left panel: input image
    orig = Image.open(image_path).convert("RGB").resize((IMG_SZ, IMG_SZ), Image.LANCZOS)
    canvas.paste(orig, (MARGIN, MARGIN))
    draw.rectangle(
        [MARGIN, MARGIN, MARGIN + IMG_SZ - 1, MARGIN + IMG_SZ - 1],
        outline=BORDER,
        width=1,
    )

    # Right panel: predictions
    px = MARGIN + IMG_SZ + MARGIN
    py = MARGIN
    draw.rectangle(
        [px, py, px + PANEL_W, py + IMG_SZ], fill=PANEL_BG, outline=BORDER, width=1
    )

    draw.text((px + 10, py + 6), "Top-5 Predictions", font=font_bold, fill=TEXT_DARK)
    draw.line(
        [px + 1, py + TITLE_H, px + PANEL_W - 1, py + TITLE_H], fill=BORDER, width=1
    )

    max_prob = results[0]["probability"] if results else 1.0
    BAR_MAX_W = PANEL_W - 24

    for rank, pred in enumerate(results):
        ry = py + TITLE_H + 8 + rank * ROW_H

        if rank == 0:
            draw.rectangle([px + 2, ry, px + PANEL_W - 2, ry + ROW_H - 2], fill=TOP1_BG)

        badge_text = f"#{rank + 1}"
        draw.text(
            (px + 8, ry + 4),
            badge_text,
            font=font_bold,
            fill=BAR_TOP1 if rank == 0 else TEXT_GRAY,
        )

        label = pred["label"]
        max_chars = 32
        if len(label) > max_chars:
            label = label[: max_chars - 1] + "…"
        draw.text((px + 36, ry + 4), label, font=font, fill=TEXT_DARK)

        bar_w = max(4, int((pred["probability"] / max_prob) * BAR_MAX_W))
        bar_y = ry + ROW_H - 14
        bar_color = BAR_TOP1 if rank == 0 else BAR_REST
        draw.rectangle([px + 8, bar_y, px + 8 + bar_w, bar_y + 8], fill=bar_color)

        pct_text = f"{pred['probability'] * 100:.2f}%"
        draw.text((px + PANEL_W - 56, ry + 4), pct_text, font=font, fill=TEXT_GRAY)

    canvas.save(str(output_path), quality=95)
    print(f"Output image saved to: {output_path}")


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output", type=Path, help="Write predictions as JSON")
    parser.add_argument("--output-image", type=Path, help="Write the rendered result image")
    args = parser.parse_args()

    script_dir = Path(__file__).resolve().parent
    model_path = str(script_dir / MODEL_PATH)
    image_path = str(script_dir / IMAGE_PATH)
    labels_path = str(script_dir / LABELS_PATH)

    method = load_model(model_path)
    labels = load_labels(labels_path)

    input_tensor = preprocess(image_path)
    logits = run_inference(method, input_tensor)
    results = postprocess(logits, labels)

    print(f"Top-{TOP_K} predictions for {IMAGE_PATH}:")
    for rank, pred in enumerate(results, start=1):
        print(f"  {rank}. {pred['label']} ({pred['probability']:.4f})")

    if args.output:
        save_results(results, args.output)
    if args.output_image:
        save_output_image(image_path, results, args.output_image)


if __name__ == "__main__":
    main()