Download example.py from Arm/vit-base-int8-xnnpack-executorch: direct link, hf CLI and curl.
- Browser
- Download file 7.85 kB
-
https://huggingface.co/Arm/vit-base-int8-xnnpack-executorch/resolve/main/example.py
- Command line
-
hf download hf://Arm/vit-base-int8-xnnpack-executorch/example.py
-
curl -L -o example.py https://huggingface.co/Arm/vit-base-int8-xnnpack-executorch/resolve/main/example.py
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() | |