swin-tiny-int8-litert / example.py
aorabdel's picture
Sync model repo (text/metadata)
ca740ff verified
Raw History Blame Contribute Delete
4.67 kB
"""
Inference example for microsoft/swin-tiny-patch4-window7-224 — INT8 LiteRT (.tflite)
Requirements:
Declared in pyproject.toml and pinned in uv.lock. Install them with:
uv python install && uv sync --frozen
Usage:
python example.py --model <path/to/microsoft__swin-tiny-patch4-window7-224_android_litert_optimized.tflite> --image <path/to/image.jpg>
Outputs (saved to the same directory as this script):
predictions.json — top-5 class predictions with scores
sample_output.jpg — input image annotated with the top prediction
"""
import argparse
import json
import os
import numpy as np
from PIL import Image, ImageDraw, ImageFont
from torchvision.models import Swin_T_Weights
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
# ImageNet preprocessing constants
MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
RESIZE_EDGE = 232
CROP_SIZE = (224, 224)
TOP_K = 5
def preprocess(image_path: str) -> np.ndarray:
img = Image.open(image_path).convert("RGB")
# Resize shortest edge to 232, matching the evaluation pipeline
w, h = img.size
scale = RESIZE_EDGE / min(w, h)
img = img.resize((int(w * scale), int(h * scale)), Image.BILINEAR)
# Center crop to 224×224
cw, ch = img.size
left = (cw - CROP_SIZE[1]) // 2
top = (ch - CROP_SIZE[0]) // 2
img = img.crop((left, top, left + CROP_SIZE[1], top + CROP_SIZE[0]))
arr = np.array(img, dtype=np.float32) / 255.0
arr = (arr - MEAN) / STD
arr = arr.transpose(2, 0, 1)
return arr[np.newaxis, :, :, :].astype(np.float32)
def softmax(x: np.ndarray) -> np.ndarray:
e = np.exp(x - x.max())
return e / e.sum()
def run_inference(model_path: str, input_tensor: np.ndarray) -> np.ndarray:
from ai_edge_litert.interpreter import Interpreter
interpreter = Interpreter(model_path=model_path)
interpreter.allocate_tensors()
input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()
interpreter.set_tensor(input_details[0]["index"], input_tensor)
interpreter.invoke()
return interpreter.get_tensor(output_details[0]["index"])[0]
def annotate_image(image_path: str, label: str, score: float, output_path: str) -> None:
img = Image.open(image_path).convert("RGB")
draw = ImageDraw.Draw(img)
text = f"{label}: {score:.2%}"
try:
font = ImageFont.truetype("/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf", 18)
except OSError:
font = ImageFont.load_default()
bbox = draw.textbbox((0, 0), text, font=font)
text_w = bbox[2] - bbox[0]
text_h = bbox[3] - bbox[1]
margin = 6
draw.rectangle([0, 0, text_w + 2 * margin, text_h + 2 * margin], fill=(0, 0, 0, 180))
draw.text((margin, margin), text, fill=(255, 255, 255), font=font)
img.save(output_path)
def main() -> None:
parser = argparse.ArgumentParser(description="Swin-Tiny INT8 LiteRT inference")
parser.add_argument(
"--model",
default=os.path.join(
SCRIPT_DIR,
"microsoft__swin-tiny-patch4-window7-224_android_litert_optimized.tflite",
),
help="Path to the optimized .tflite model",
)
parser.add_argument(
"--image",
default=os.path.join(SCRIPT_DIR, "sample_input.jpg"),
help="Path to the input image",
)
args = parser.parse_args()
# ImageNet class labels (1000 classes)
WEIGHTS = Swin_T_Weights.IMAGENET1K_V1
IMAGENET_CLASSES = WEIGHTS.meta["categories"]
print(f"Running inference on: {args.image}")
input_tensor = preprocess(args.image)
logits = run_inference(args.model, input_tensor)
probs = softmax(logits)
top_indices = np.argsort(probs)[::-1][:TOP_K]
predictions = [
{
"rank": int(i + 1),
"class_index": int(idx),
"label": IMAGENET_CLASSES[int(idx)],
"score": float(probs[idx]),
}
for i, idx in enumerate(top_indices)
]
print("\nTop-5 predictions:")
for pred in predictions:
print(f" {pred['rank']}. {pred['label']:<40s} {pred['score']:.4f}")
predictions_path = os.path.join(SCRIPT_DIR, "predictions.json")
with open(predictions_path, "w") as f:
json.dump(predictions, f, indent=2)
print(f"\nPredictions saved to: {predictions_path}")
output_image_path = os.path.join(SCRIPT_DIR, "sample_output.jpg")
annotate_image(args.image, predictions[0]["label"], predictions[0]["score"], output_image_path)
print(f"Annotated image saved to: {output_image_path}")
if __name__ == "__main__":
main()