""" Inference example for microsoft/swin-tiny-patch4-window7-224 — INT8 LiteRT (.tflite) Requirements: Declared in pyproject.toml and pinned in uv.lock. Install them with: uv python install && uv sync --frozen Usage: python example.py --model --image Outputs (saved to the same directory as this script): predictions.json — top-5 class predictions with scores sample_output.jpg — input image annotated with the top prediction """ import argparse import json import os import numpy as np from PIL import Image, ImageDraw, ImageFont from torchvision.models import Swin_T_Weights SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__)) # ImageNet preprocessing constants MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32) STD = np.array([0.229, 0.224, 0.225], dtype=np.float32) RESIZE_EDGE = 232 CROP_SIZE = (224, 224) TOP_K = 5 def preprocess(image_path: str) -> np.ndarray: img = Image.open(image_path).convert("RGB") # Resize shortest edge to 232, matching the evaluation pipeline w, h = img.size scale = RESIZE_EDGE / min(w, h) img = img.resize((int(w * scale), int(h * scale)), Image.BILINEAR) # Center crop to 224×224 cw, ch = img.size left = (cw - CROP_SIZE[1]) // 2 top = (ch - CROP_SIZE[0]) // 2 img = img.crop((left, top, left + CROP_SIZE[1], top + CROP_SIZE[0])) arr = np.array(img, dtype=np.float32) / 255.0 arr = (arr - MEAN) / STD arr = arr.transpose(2, 0, 1) return arr[np.newaxis, :, :, :].astype(np.float32) def softmax(x: np.ndarray) -> np.ndarray: e = np.exp(x - x.max()) return e / e.sum() def run_inference(model_path: str, input_tensor: np.ndarray) -> np.ndarray: from ai_edge_litert.interpreter import Interpreter interpreter = Interpreter(model_path=model_path) interpreter.allocate_tensors() input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() interpreter.set_tensor(input_details[0]["index"], input_tensor) interpreter.invoke() return interpreter.get_tensor(output_details[0]["index"])[0] def annotate_image(image_path: str, label: str, score: float, output_path: str) -> None: img = Image.open(image_path).convert("RGB") draw = ImageDraw.Draw(img) text = f"{label}: {score:.2%}" try: font = ImageFont.truetype("/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf", 18) except OSError: font = ImageFont.load_default() bbox = draw.textbbox((0, 0), text, font=font) text_w = bbox[2] - bbox[0] text_h = bbox[3] - bbox[1] margin = 6 draw.rectangle([0, 0, text_w + 2 * margin, text_h + 2 * margin], fill=(0, 0, 0, 180)) draw.text((margin, margin), text, fill=(255, 255, 255), font=font) img.save(output_path) def main() -> None: parser = argparse.ArgumentParser(description="Swin-Tiny INT8 LiteRT inference") parser.add_argument( "--model", default=os.path.join( SCRIPT_DIR, "microsoft__swin-tiny-patch4-window7-224_android_litert_optimized.tflite", ), help="Path to the optimized .tflite model", ) parser.add_argument( "--image", default=os.path.join(SCRIPT_DIR, "sample_input.jpg"), help="Path to the input image", ) args = parser.parse_args() # ImageNet class labels (1000 classes) WEIGHTS = Swin_T_Weights.IMAGENET1K_V1 IMAGENET_CLASSES = WEIGHTS.meta["categories"] print(f"Running inference on: {args.image}") input_tensor = preprocess(args.image) logits = run_inference(args.model, input_tensor) probs = softmax(logits) top_indices = np.argsort(probs)[::-1][:TOP_K] predictions = [ { "rank": int(i + 1), "class_index": int(idx), "label": IMAGENET_CLASSES[int(idx)], "score": float(probs[idx]), } for i, idx in enumerate(top_indices) ] print("\nTop-5 predictions:") for pred in predictions: print(f" {pred['rank']}. {pred['label']:<40s} {pred['score']:.4f}") predictions_path = os.path.join(SCRIPT_DIR, "predictions.json") with open(predictions_path, "w") as f: json.dump(predictions, f, indent=2) print(f"\nPredictions saved to: {predictions_path}") output_image_path = os.path.join(SCRIPT_DIR, "sample_output.jpg") annotate_image(args.image, predictions[0]["label"], predictions[0]["score"], output_image_path) print(f"Annotated image saved to: {output_image_path}") if __name__ == "__main__": main()