aorabdel's picture
Sync model repo (text/metadata)
60d2da1 verified
Raw History Blame Contribute Delete
7.85 kB
"""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()