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)