aorabdel's picture
Sync model repo (text/metadata)
67d0363 verified
Raw
History Blame Contribute Delete
4.96 kB
"""Minimal inference example for DeiT-Tiny INT8 using ExecuTorch.
Loads a quantized .pte model and runs inference on a single image,
printing the top-5 ImageNet class predictions with probabilities.
"""
import json
from pathlib import Path
import torch
from executorch.runtime import Runtime
from PIL import Image
from torchvision import transforms
from torchvision.models import GoogLeNet_Weights
# ── Configuration ──────────────────────────────────────────────────────────────
MODEL_PATH = "deit_raspberry_executorch_optimized.pte"
IMAGE_PATH = "sample_input.jpg"
INPUT_SIZE = (224, 224)
RESIZE_SIZE = 256 # resize shorter edge before center crop
MEAN = [0.485, 0.456, 0.406]
STD = [0.229, 0.224, 0.225]
TOP_K = 5
# ── Model Loading ──────────────────────────────────────────────────────────────
def load_model(pte_path: str):
"""Load ExecuTorch .pte model and return the forward method."""
runtime = Runtime.get()
program = runtime.load_program(pte_path)
return program.load_method("forward")
# ── Preprocessing ──────────────────────────────────────────────────────────────
def preprocess(image_path: str) -> torch.Tensor:
"""Load and preprocess image for model input.
Pipeline:
1. Resize shortest edge to 256 px (bicubic)
2. Center-crop to 224x224
3. Convert to float32 tensor in [0, 1]
4. Normalize with ImageNet mean/std
"""
transform = transforms.Compose(
[
transforms.Resize(
RESIZE_SIZE, interpolation=transforms.InterpolationMode.BICUBIC
),
transforms.CenterCrop(INPUT_SIZE),
transforms.ToTensor(),
transforms.Normalize(mean=MEAN, std=STD),
]
)
image = Image.open(image_path).convert("RGB")
tensor = transform(image)
return tensor.unsqueeze(0) # [1, 3, 224, 224]
# ── Inference ──────────────────────────────────────────────────────────────────
def run_inference(method, input_tensor: torch.Tensor) -> torch.Tensor:
"""Run forward pass and return the raw logits tensor."""
outputs = method.execute([input_tensor])
return outputs[0]
# ── Postprocessing ─────────────────────────────────────────────────────────────
def postprocess(raw_output: torch.Tensor, labels: list) -> list[dict]:
"""Convert raw logits to top-k class predictions.
Applies softmax over the 1000-class logit vector, then returns the
top-k (class name, probability) pairs sorted by descending probability.
"""
logits = raw_output.squeeze(0) # [1000]
probabilities = torch.softmax(logits, dim=-1)
top_probs, top_indices = torch.topk(probabilities, TOP_K)
return [
{"class": labels[idx.item()], "probability": round(prob.item(), 6)}
for prob, idx in zip(top_probs, top_indices, strict=False)
]
# ── Save Results ───────────────────────────────────────────────────────────────
def save_results(results: list[dict], script_dir: Path) -> None:
"""Save top-k predictions to predictions.json in the script's directory."""
output_path = script_dir / "predictions.json"
with open(output_path, "w") as f:
json.dump(results, f, indent=2)
print(f"Saved predictions to {output_path}")
# ── Main ───────────────────────────────────────────────────────────────────────
def main() -> None:
script_dir = Path(__file__).parent
# ImageNet class labels (1000 classes)
WEIGHTS = GoogLeNet_Weights.IMAGENET1K_V1
IMAGENET_CLASSES = WEIGHTS.meta["categories"]
# Load model, run inference, decode results
method = load_model(str(script_dir / MODEL_PATH))
input_tensor = preprocess(str(script_dir / IMAGE_PATH))
raw_output = run_inference(method, input_tensor)
results = postprocess(raw_output, IMAGENET_CLASSES)
# Print top-k predictions
print(f"Top-{TOP_K} predictions:")
for i, r in enumerate(results, 1):
print(f" {i}. {r['class']} ({r['probability'] * 100:.2f}%)")
save_results(results, script_dir)
if __name__ == "__main__":
main()