aorabdel's picture
Sync model repo (text/metadata)
0e9a1b5 verified
Raw History Blame Contribute Delete
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()