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
File size: 6,361 Bytes
814a093 bd5b2c4 814a093 bd5b2c4 814a093 bd5b2c4 814a093 bd5b2c4 814a093 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 | """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)
|