"""Inference example for DeiT-Tiny INT8 (LiteRT .tflite). Requirements: pip install ai-edge-litert pillow torch numpy Usage: python example.py # uses sample_input.jpg python example.py --image path/to/img.jpg """ from __future__ import annotations import argparse import json import math from pathlib import Path from torchvision.models import GoogLeNet_Weights import numpy as np from PIL import Image # --------------------------------------------------------------------------- # Paths # --------------------------------------------------------------------------- SCRIPT_DIR = Path(__file__).parent.resolve() _model_in_parent = SCRIPT_DIR.parent / "facebook__deit-tiny-patch16-224_litert_optimized.tflite" MODEL_PATH = _model_in_parent if _model_in_parent.exists() else SCRIPT_DIR / "facebook__deit-tiny-patch16-224_litert_optimized.tflite" IMAGE_PATH = SCRIPT_DIR / "sample_input.jpg" # --------------------------------------------------------------------------- # Constants # --------------------------------------------------------------------------- INPUT_SIZE = 224 RESIZE_SIZE = 256 IMAGENET_MEAN = [0.485, 0.456, 0.406] IMAGENET_STD = [0.229, 0.224, 0.225] TOP_K = 5 # --------------------------------------------------------------------------- # Model loading # --------------------------------------------------------------------------- def load_model(model_path: Path): """Load LiteRT .tflite model and return allocated interpreter.""" from ai_edge_litert.interpreter import Interpreter # noqa: PLC0415 interpreter = Interpreter(model_path=str(model_path)) interpreter.allocate_tensors() return interpreter # --------------------------------------------------------------------------- # Preprocessing # --------------------------------------------------------------------------- def preprocess(image_path: Path) -> np.ndarray: """Load and preprocess an image for DeiT-Tiny input. Steps: 1. Resize the shorter edge to 256 using bicubic interpolation. 2. Center crop to 224x224. 3. Normalize to [0, 1] and apply ImageNet mean/std. Returns: Array of shape (1, 3, 224, 224), dtype float32. """ img = Image.open(image_path).convert("RGB") w, h = img.size scale = RESIZE_SIZE / min(w, h) new_w, new_h = math.ceil(w * scale), math.ceil(h * scale) img = img.resize((new_w, new_h), Image.BICUBIC) left = (new_w - INPUT_SIZE) // 2 top = (new_h - INPUT_SIZE) // 2 img = img.crop((left, top, left + INPUT_SIZE, top + INPUT_SIZE)) arr = np.array(img, dtype=np.float32) / 255.0 mean = np.array(IMAGENET_MEAN, dtype=np.float32).reshape(1, 1, 3) std = np.array(IMAGENET_STD, dtype=np.float32).reshape(1, 1, 3) arr = (arr - mean) / std return arr.transpose(2, 0, 1)[np.newaxis, ...] # (1, 3, H, W) # --------------------------------------------------------------------------- # Postprocessing # --------------------------------------------------------------------------- def postprocess(logits: np.ndarray, top_k: int = TOP_K) -> list[dict]: """Apply softmax and return top-k predictions. Args: logits: Array of shape (1, 1000) with raw class scores. top_k: Number of top predictions to return. Returns: List of dicts: {"rank": int, "class": str, "probability": float}. """ logits = logits[0] exp = np.exp(logits - logits.max()) probs = exp / exp.sum() top_indices = np.argsort(probs)[::-1][:top_k] # ImageNet class labels (1000 classes) WEIGHTS = GoogLeNet_Weights.IMAGENET1K_V1 classes = WEIGHTS.meta["categories"] return [ {"rank": i + 1, "class": classes[idx], "probability": float(probs[idx])} for i, idx in enumerate(top_indices) ] # --------------------------------------------------------------------------- # Inference # --------------------------------------------------------------------------- def run_inference(interpreter, input_array: np.ndarray) -> np.ndarray: """Run forward pass and return dequantized logits. If the output tensor is INT8 (as produced by PT2E static quantization), it is dequantized to float32 using the tensor's scale and zero_point before being returned. Float32 outputs are returned as-is. """ input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() interpreter.set_tensor(input_details[0]["index"], input_array) interpreter.invoke() raw = interpreter.get_tensor(output_details[0]["index"]) if raw.dtype == np.int8: quant_params = output_details[0].get("quantization_parameters", {}) scales = quant_params.get("scales", None) zero_points = quant_params.get("zero_points", None) if scales is not None and len(scales) > 0: scale = float(scales[0]) zero_point = int(zero_points[0]) else: # Fall back to the legacy (scale, zero_point) tuple scale, zero_point = output_details[0]["quantization"] return (raw.astype(np.float32) - zero_point) * scale return raw.astype(np.float32) # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- def main(image_path: Path) -> None: print(f"Loading model from {MODEL_PATH} ...") interpreter = load_model(MODEL_PATH) print(f"Preprocessing image: {image_path}") input_array = preprocess(image_path) print("Running inference ...") logits = run_inference(interpreter, input_array) predictions = postprocess(logits) print("\nTop-5 predictions:") for pred in predictions: print(f" Top-{pred['rank']}: {pred['class']} ({pred['probability']:.4f})") out_path = SCRIPT_DIR / "predictions.json" with open(out_path, "w") as f: json.dump({"predictions": predictions}, f, indent=2) print(f"\nPredictions saved to {out_path}") if __name__ == "__main__": parser = argparse.ArgumentParser(description="DeiT-Tiny INT8 LiteRT inference") parser.add_argument( "--image", type=Path, default=IMAGE_PATH, help="Path to input image (default: sample_input.jpg next to this script)", ) args = parser.parse_args() main(args.image)