Instructions to use Arm/deit-tiny-int8-litert with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- LiteRT
How to use Arm/deit-tiny-int8-litert with LiteRT:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
| """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) | |