""" CLIP ViT-B/32 INT8 -- Zero-Shot ImageNet Classification Inference Example """ import json import sys from pathlib import Path import torch import torchvision.transforms as T from torchvision.models import GoogLeNet_Weights from PIL import Image, ImageDraw, ImageFont # --------------------------------------------------------------------------- # Constants # --------------------------------------------------------------------------- # ImageNet class labels (1000 classes) WEIGHTS = GoogLeNet_Weights.IMAGENET1K_V1 IMAGENET_CLASSES = WEIGHTS.meta["categories"] PROMPT_TEMPLATE = "a photo of a {}" CLIP_MEAN = [0.48145466, 0.4578275, 0.40821073] CLIP_STD = [0.26862954, 0.26130258, 0.27577711] IMAGE_PATH = Path(__file__).parent / "sample_input.jpg" MODEL_PATH = Path(__file__).parent / "clip_raspberry_executorch_optimized.pte" OUTPUT_IMAGE_PATH = Path(__file__).parent / "sample_output.jpg" OUTPUT_JSON_PATH = Path(__file__).parent / "predictions.json" # Overlay appearance PANEL_WIDTH = 340 BAR_BG = (28, 40, 51) TEXT_COLOR = (255, 255, 255) HEADER_COLOR = (240, 240, 240) RANK_COLORS = [ (255, 215, 0), # gold (192, 192, 192), # silver (205, 127, 50), # bronze (160, 160, 160), (130, 130, 130), ] # --------------------------------------------------------------------------- # Preprocessing # --------------------------------------------------------------------------- _transform = T.Compose([ T.Resize(224, interpolation=T.InterpolationMode.BICUBIC), T.CenterCrop(224), T.ToTensor(), T.Normalize(mean=CLIP_MEAN, std=CLIP_STD), ]) def preprocess(image_path: Path) -> torch.Tensor: """Load and preprocess an image for the CLIP vision encoder. Returns float32 tensor [1, 3, 224, 224]. """ img = Image.open(image_path).convert("RGB") return _transform(img).unsqueeze(0) # --------------------------------------------------------------------------- # Tokenisation # --------------------------------------------------------------------------- def tokenize(prompt: str) -> dict[str, torch.Tensor]: """Tokenize a single prompt. Returns input_ids and attention_mask as Long [1, 77].""" try: from transformers import CLIPTokenizer except ImportError as exc: raise ImportError( "transformers is required. Install with: pip install transformers" ) from exc tok = tokenize._tokenizer enc = tok(prompt, return_tensors="pt", padding="max_length", max_length=77, truncation=True) return { "input_ids": enc["input_ids"].to(torch.long), "attention_mask": enc["attention_mask"].to(torch.long), } def _init_tokenizer() -> None: try: from transformers import CLIPTokenizer tokenize._tokenizer = CLIPTokenizer.from_pretrained("openai/clip-vit-base-patch32") except ImportError as exc: raise ImportError( "transformers is required. Install with: pip install transformers" ) from exc # --------------------------------------------------------------------------- # Output image rendering # --------------------------------------------------------------------------- def _load_font(size: int): for name in ("DejaVuSans.ttf", "Arial.ttf", "FreeSans.ttf"): try: return ImageFont.truetype(name, size) except (OSError, IOError): pass return ImageFont.load_default() def render_output_image(source_path: Path, predictions: list[dict], output_path: Path) -> None: """Render the input image with a top-5 predictions panel and save it.""" img = Image.open(source_path).convert("RGB") target_h = 420 scale = target_h / img.height img = img.resize((int(img.width * scale), target_h), Image.LANCZOS) total_w = img.width + PANEL_WIDTH canvas = Image.new("RGB", (total_w, target_h), BAR_BG) canvas.paste(img, (0, 0)) draw = ImageDraw.Draw(canvas) font_title = _load_font(15) font_label = _load_font(13) font_score = _load_font(12) draw.text((img.width + 12, 10), "Top-5 Predictions (ImageNet)", font=font_title, fill=HEADER_COLOR) draw.text((img.width + 12, 27), f"({len(IMAGENET_CLASSES)} classes · single prompt)", font=font_score, fill=(140, 140, 140)) draw.line([(img.width + 8, 44), (total_w - 8, 44)], fill=(60, 60, 60), width=1) row_h = (target_h - 52) // 5 bar_max_w = PANEL_WIDTH - 72 for i, pred in enumerate(predictions[:5]): label = pred["class"].replace("_", " ").title() score = pred["score"] pct = score * 100 color = RANK_COLORS[i] y_top = 50 + i * row_h # Rank badge draw.ellipse([img.width + 8, y_top + 4, img.width + 24, y_top + 20], fill=color) draw.text((img.width + 12, y_top + 4), str(i + 1), font=font_score, fill=(20, 20, 20)) # Label draw.text((img.width + 30, y_top + 3), label, font=font_label, fill=TEXT_COLOR) # Bar bar_y = y_top + 22 draw.rectangle([img.width + 8, bar_y, img.width + 8 + bar_max_w, bar_y + 10], fill=(55, 55, 55)) filled = int(bar_max_w * min(score * 5, 1.0)) # scale up low probs for visibility if filled > 0: draw.rectangle([img.width + 8, bar_y, img.width + 8 + filled, bar_y + 10], fill=color) # Score draw.text((img.width + 8 + bar_max_w + 6, bar_y - 1), f"{pct:.2f}%", font=font_score, fill=color) canvas.save(output_path, quality=92) print(f"Output image saved to {output_path}") # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- def main() -> None: image_path = IMAGE_PATH if "--image" in sys.argv: image_path = Path(sys.argv[sys.argv.index("--image") + 1]) if not image_path.exists(): print(f"Image not found: {image_path}") print("Place a JPEG/PNG at sample_input.jpg or pass --image ") sys.exit(1) if not MODEL_PATH.exists(): print(f"Model not found: {MODEL_PATH}") sys.exit(1) print(f"ImageNet classes loaded: {len(IMAGENET_CLASSES)}") # Load ExecuTorch model try: from executorch.runtime import Runtime from executorch.extension.pybindings._portable_lib import Verification except ImportError as exc: raise ImportError( "executorch is required. " "See https://pytorch.org/executorch/stable/getting-started-setup.html" ) from exc print(f"Loading model from {MODEL_PATH} ...") runtime = Runtime.get() program = runtime.load_program(str(MODEL_PATH), verification=Verification.Minimal) encode_image = program.load_method("encode_image") encode_text = program.load_method("encode_text") print("Model loaded.") _init_tokenizer() # Encode image print(f"Preprocessing image: {image_path}") pixel_values = preprocess(image_path) image_embeds = encode_image.execute([pixel_values])[0] # [1, 512] n_classes = len(IMAGENET_CLASSES) print(f"Encoding {n_classes} class prompts ...") text_embeds_list: list[torch.Tensor] = [] for c_idx, cls in enumerate(IMAGENET_CLASSES): prompt = PROMPT_TEMPLATE.format(cls) toks = tokenize(prompt) embed = encode_text.execute([toks["input_ids"], toks["attention_mask"]])[0] # [1, 512] text_embeds_list.append(embed) if (c_idx + 1) % 100 == 0: print(f" {c_idx + 1}/{n_classes}") class_embeds = torch.cat(text_embeds_list, dim=0) # [1000, 512] # Cosine similarity → softmax → top-5 logits = (image_embeds @ class_embeds.T).squeeze(0) # [1000] probs = torch.softmax(logits * 100.0, dim=-1) # temperature from CLIP paper top5_values, top5_indices = torch.topk(probs, k=5) print("\nTop-5 Zero-Shot Predictions (ImageNet):") print("-" * 45) predictions = [] for rank, (idx, prob) in enumerate(zip(top5_indices.tolist(), top5_values.tolist()), 1): label = IMAGENET_CLASSES[idx] print(f" {rank}. {label:<30} {prob * 100:.2f}%") predictions.append({"rank": rank, "class": label, "score": round(float(prob), 6)}) # Save predictions.json with open(OUTPUT_JSON_PATH, "w") as f: json.dump({"image": str(image_path), "model": str(MODEL_PATH), "dataset": "imagenet-1k", "top5": predictions}, f, indent=2) print(f"\nPredictions saved to {OUTPUT_JSON_PATH}") # Render output image render_output_image(image_path, predictions, OUTPUT_IMAGE_PATH) if __name__ == "__main__": main()