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