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