deit-tiny-int8-litert / example.py
aorabdel's picture
Full catalogue sync: example.py
bd5b2c4 verified
Raw
History Blame Contribute Delete
6.36 kB
"""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)