"""Minimal inference example for Swin Tiny INT8 using ExecuTorch.""" 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 = "swin_tiny_dynamic_raspberry_executorch_optimized.pte" IMAGE_PATH = "sample_input.jpg" INPUT_SIZE = (224, 224) # Height, Width expected by the model TOP_K = 5 # Number of top predictions to return # Normalization constants from the Swin image processor (ImageNet stats) IMAGENET_MEAN = [0.485, 0.456, 0.406] IMAGENET_STD = [0.229, 0.224, 0.225] SCRIPT_DIR = Path(__file__).resolve().parent # ImageNet class labels (1000 classes), in model output order with (SCRIPT_DIR / "imagenet_classes.json").open(encoding="utf-8") as file: IMAGENET_CLASSES = json.load(file) # Bar colours per rank (blue → green → yellow → orange → red) BAR_COLORS = [ (52, 152, 219), (46, 204, 113), (241, 196, 15), (230, 126, 34), (231, 76, 60), ] def load_model(pte_path: str): """Load ExecuTorch .pte model and return the forward method.""" runtime = Runtime.get() program = runtime.load_program(str(SCRIPT_DIR / pte_path)) return program.load_method("forward") def preprocess(image_path: str) -> torch.Tensor: """Load and preprocess image for Swin model input. Pipeline: Resize(232) -> CenterCrop(224, 224) -> ToTensor -> Normalize Input values are in [0, 1] after ToTensor, then normalized with ImageNet stats. """ image = Image.open(str(SCRIPT_DIR / image_path)).convert("RGB") transform = transforms.Compose( [ transforms.Resize(232), transforms.CenterCrop(INPUT_SIZE), transforms.ToTensor(), transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD), ] ) tensor = transform(image) # Add batch dimension: [C, H, W] -> [1, C, H, W] return tensor.unsqueeze(0) def get_display_image(image_path: str) -> Image.Image: """Return the same 224×224 center-crop used during preprocessing.""" image = Image.open(str(SCRIPT_DIR / image_path)).convert("RGB") return transforms.Compose( [transforms.Resize(232), transforms.CenterCrop(INPUT_SIZE)] )(image) def run_inference(method, input_tensor: torch.Tensor) -> torch.Tensor: """Run forward pass and return raw logits tensor [1, 1000].""" outputs = method.execute([input_tensor]) return outputs[0] def postprocess(raw_output: torch.Tensor, labels: dict[str, str]) -> list[dict]: """Decode raw logits into top-k class predictions. Applies softmax to convert logits to probabilities, then returns the top-k predictions with class names and scores. """ # Ensure 2D: [1, num_classes] logits = raw_output if logits.dim() == 1: logits = logits.unsqueeze(0) # Softmax to get probabilities probabilities = torch.softmax(logits, dim=-1) # Top-k predictions 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 ): label = labels[str(idx)] results.append( {"class_index": idx, "class": label, "probability": round(score, 6)} ) return results def _load_fonts(sizes: tuple[int, int]) -> tuple: """Load DejaVu fonts, falling back to PIL default.""" bold = "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf" regular = "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf" try: return ( ImageFont.truetype(bold, sizes[0]), ImageFont.truetype(regular, sizes[1]), ImageFont.truetype(bold, sizes[1]), ) except OSError: default = ImageFont.load_default() return default, default, default def save_output_image(image_path: str, results: list[dict]) -> None: """Render input image + top-5 prediction bars and save as sample_output.jpg.""" img = get_display_image(image_path) # Layout: 224px image | 12px gap | 270px predictions panel img_w, img_h = 224, 224 gap = 12 panel_w = 270 canvas_w = img_w + gap + panel_w canvas_h = img_h + 12 # small top/bottom margin canvas = Image.new("RGB", (canvas_w, canvas_h), (245, 245, 245)) canvas.paste(img, (0, (canvas_h - img_h) // 2)) draw = ImageDraw.Draw(canvas) font_title, font_label, font_pct = _load_fonts((13, 11)) x0 = img_w + gap bar_w = 190 # width of the probability bar pct_x = x0 + bar_w + 6 y = 10 draw.text((x0, y), "Top-5 Predictions", fill=(30, 30, 30), font=font_title) y += 22 for i, pred in enumerate(results): label = pred["class"] prob = pred["probability"] # Truncate long class names display = label if len(label) <= 24 else label[:23] + "…" draw.text((x0, y), f"{i + 1}. {display}", fill=(50, 50, 50), font=font_label) y += 15 # Background bar draw.rectangle([x0, y, x0 + bar_w, y + 13], fill=(210, 210, 210)) # Filled bar proportional to probability fill_w = max(1, int(bar_w * prob)) draw.rectangle([x0, y, x0 + fill_w, y + 13], fill=BAR_COLORS[i]) # Percentage label draw.text( (pct_x, y + 1), f"{prob * 100:.1f}%", fill=(60, 60, 60), font=font_pct ) y += 20 if i < len(results) - 1 else 0 out_path = SCRIPT_DIR / "sample_output.jpg" canvas.save(out_path, quality=95) print(f"Saved output image to {out_path}") def save_predictions_json(results: list[dict]) -> None: """Persist top-k predictions as JSON next to this script.""" out_path = SCRIPT_DIR / "predictions.json" with open(out_path, "w") as f: json.dump(results, f, indent=2) print(f"Saved predictions to {out_path}") def main() -> None: labels = IMAGENET_CLASSES # Load model print(f"Loading model from: {SCRIPT_DIR / MODEL_PATH}") method = load_model(MODEL_PATH) # Preprocess input image print(f"Preprocessing image: {SCRIPT_DIR / IMAGE_PATH}") input_tensor = preprocess(IMAGE_PATH) # Run inference print("Running inference...") raw_output = run_inference(method, input_tensor) # Postprocess outputs results = postprocess(raw_output, labels) # Print top-k predictions print(f"\nTop-{TOP_K} predictions:") for i, pred in enumerate(results, 1): print( f" {i}. {pred['class']} — {pred['probability']:.4f} ({pred['probability']*100:.2f}%)" ) # Save predictions JSON and annotated output image save_predictions_json(results) save_output_image(IMAGE_PATH, results) if __name__ == "__main__": main()