aorabdel's picture
Sync model repo (text/metadata)
1c7a596 verified
Raw History Blame Contribute Delete
6.71 kB
"""Minimal inference example for INT8 PIDNet-S semantic segmentation with ExecuTorch.
Loads the quantized ``.pte`` model, segments a single street image into the 19
Cityscapes trainId classes, and writes a colour-coded overlay plus a per-class
pixel summary next to this script.
Input pipeline: RGB resized to 2048x1024, scaled to [0, 1], ImageNet mean/std
normalized. The network emits logits at 1/8 resolution, which are bilinearly
upsampled with ``align_corners=True`` before softmax and argmax.
"""
import json
from pathlib import Path
import numpy as np
import torch
import torch.nn.functional as F
from executorch.runtime import Runtime
from PIL import Image
# ── Configuration ──────────────────────────────────────────
MODEL_PATH = "pidnet_s_raspberry_executorch_optimized.pte"
IMAGE_PATH = "sample_input.jpg"
INPUT_SIZE = (1024, 2048) # (H, W)
OUTPUT_STRIDE = 8 # logits are emitted at 1/8 input resolution
IMAGENET_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
IMAGENET_STD = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)
OVERLAY_ALPHA = 0.5
# ── Class labels and colours (19 Cityscapes trainId classes) ───────────────────
CLASS_NAMES = (
"road",
"sidewalk",
"building",
"wall",
"fence",
"pole",
"traffic light",
"traffic sign",
"vegetation",
"terrain",
"sky",
"person",
"rider",
"car",
"truck",
"bus",
"train",
"motorcycle",
"bicycle",
)
CLASS_COLORS = (
(128, 64, 128),
(244, 35, 232),
(70, 70, 70),
(102, 102, 156),
(190, 153, 153),
(153, 153, 153),
(250, 170, 30),
(220, 220, 0),
(107, 142, 35),
(152, 251, 152),
(70, 130, 180),
(220, 20, 60),
(255, 0, 0),
(0, 0, 142),
(0, 0, 70),
(0, 60, 100),
(0, 80, 100),
(0, 0, 230),
(119, 11, 32),
)
# ── Model Loading ──────────────────────────────────────────
def load_model(model_path: str):
"""Load the .pte model and return its forward method.
The method is loaded once here and reused for every inference call.
"""
runtime = Runtime.get()
program = runtime.load_program(model_path)
return program.load_method("forward")
# ── Preprocessing ──────────────────────────────────────────
def preprocess(image_path: str) -> torch.Tensor:
"""Load an image and turn it into the normalized [1, 3, 1024, 2048] input."""
image = Image.open(image_path).convert("RGB")
# PIL takes (width, height).
image = image.resize((INPUT_SIZE[1], INPUT_SIZE[0]), Image.BILINEAR)
array = np.asarray(image, dtype=np.uint8).transpose(2, 0, 1).copy()
tensor = torch.from_numpy(array).float().div_(255.0).unsqueeze(0)
return (tensor - IMAGENET_MEAN) / IMAGENET_STD
# ── Inference ──────────────────────────────────────────────
def run_inference(method, input_tensor: torch.Tensor):
"""Run one forward pass and return the raw logits."""
outputs = method.execute([input_tensor])
return outputs[0]
# ── Postprocessing ─────────────────────────────────────────
def postprocess(raw_output) -> torch.Tensor:
"""Upsample the 1/8-resolution logits and argmax into a class map.
Bilinear upsampling of the logits before the argmax (rather than resizing
the class map afterwards) is what the published accuracy protocol does, and
``align_corners=True`` is part of that protocol β€” changing it shifts the
sampling grid and costs accuracy.
"""
if isinstance(raw_output, (list, tuple)):
raw_output = raw_output[-1]
logits = F.interpolate(
raw_output,
scale_factor=OUTPUT_STRIDE,
mode="bilinear",
align_corners=True,
)
probabilities = torch.softmax(logits, dim=1)
return probabilities.argmax(dim=1)
def summarize(class_map: torch.Tensor) -> list[dict]:
"""Count pixels per present class, sorted by descending share."""
ids, counts = torch.unique(class_map, return_counts=True)
total = int(class_map.numel())
summary = [
{
"class_id": int(class_id),
"class_name": CLASS_NAMES[int(class_id)],
"pixel_count": int(count),
"percentage": round(100.0 * int(count) / total, 4),
}
for class_id, count in zip(ids.tolist(), counts.tolist(), strict=True)
]
summary.sort(key=lambda entry: entry["pixel_count"], reverse=True)
return summary
# ── Save Results ──────────────────────────────────────────
def save_results(image_path: str, class_map: torch.Tensor, summary: list[dict]) -> None:
"""Write the blended mask overlay and the per-class pixel JSON."""
script_dir = Path(__file__).parent
palette = np.asarray(CLASS_COLORS, dtype=np.uint8)
color_mask = palette[class_map[0].numpy()]
original = Image.open(image_path).convert("RGB")
original = original.resize((INPUT_SIZE[1], INPUT_SIZE[0]), Image.BILINEAR)
blended = Image.blend(original, Image.fromarray(color_mask), OVERLAY_ALPHA)
blended.save(script_dir / "sample_output.png")
(script_dir / "segmentation.json").write_text(json.dumps(summary, indent=2))
# ── Main ───────────────────────────────────────────────────
def main() -> None:
script_dir = Path(__file__).parent
model_path = str(script_dir / MODEL_PATH)
image_path = str(script_dir / IMAGE_PATH)
method = load_model(model_path)
input_tensor = preprocess(image_path)
raw_output = run_inference(method, input_tensor)
class_map = postprocess(raw_output)
summary = summarize(class_map)
print(
f"Segmented {INPUT_SIZE[1]}x{INPUT_SIZE[0]} image into "
f"{len(summary)} of {len(CLASS_NAMES)} classes"
)
print("Top classes by pixel share:")
for entry in summary[:5]:
print(
f" {entry['class_name']:<14} {entry['percentage']:6.2f}% "
f"({entry['pixel_count']} px)"
)
save_results(image_path, class_map, summary)
print("Saved sample_output.png and segmentation.json")
if __name__ == "__main__":
main()