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