File size: 11,095 Bytes
73900de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31ffec6
73900de
31ffec6
 
24b13fe
 
 
 
73900de
 
 
 
24b13fe
73900de
 
 
 
 
 
 
24b13fe
 
 
 
 
 
 
 
73900de
 
 
 
24b13fe
73900de
24b13fe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
73900de
 
24b13fe
31ffec6
 
 
 
24b13fe
 
 
 
 
73900de
31ffec6
 
73900de
 
31ffec6
 
 
 
24b13fe
 
 
73900de
31ffec6
73900de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24b13fe
 
 
73900de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24b13fe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
73900de
 
 
24b13fe
31ffec6
73900de
 
 
 
24b13fe
 
 
 
 
 
 
 
73900de
 
24b13fe
 
 
 
73900de
 
24b13fe
73900de
24b13fe
73900de
 
 
 
 
 
 
24b13fe
 
73900de
 
 
 
24b13fe
 
73900de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24b13fe
73900de
 
 
 
 
 
24b13fe
73900de
 
 
24b13fe
 
73900de
 
 
 
24b13fe
73900de
 
 
24b13fe
73900de
 
 
 
 
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
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
"""Minimal inference example for MobileSAM INT8 using ExecuTorch.

Loads a quantized .pte model and runs box-prompted instance segmentation
on a single image. The model takes an image and a box prompt (two corner
points in 1024-pixel space) and returns a binary segmentation mask.

Usage:
    python example.py

The script saves sample_output.png (mask overlay) and segmentation.json
to the same directory as this file.
"""

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, ImageDraw

# ── Configuration ──────────────────────────────────────────────────────────────
MODEL_PATH = "mobile_sam_raspberry_executorch_optimized.pte"
IMAGE_PATH = "sample_input.jpg"
INPUT_SIZE = 1024  # Fixed input resolution the model was exported at

# Box prompt in the native pixel space of sample_input.jpg (960x1280), COCO
# (x, y, w, h) convention: the dark SUV in the centre of the road.
PROMPT_BOX = (512.0, 895.0, 188.0, 180.0)


# ── Model Loading ──────────────────────────────────────────────────────────────
def load_model(pte_path: str):
    """Load ExecuTorch .pte model and return the forward method."""
    script_dir = Path(__file__).resolve().parent
    runtime = Runtime.get()
    program = runtime.load_program(str(script_dir / pte_path))
    return program.load_method("forward")


# ── Preprocessing ──────────────────────────────────────────────────────────────
def preprocess(image_path: str) -> tuple[torch.Tensor, Image.Image]:
    """Load an image and resize it to INPUT_SIZE x INPUT_SIZE in [0, 1] float32.

    Scaling and resizing run exactly as they do in the evaluation pipeline that
    produced the reported metrics: the raw uint8 pixels are divided by 255 and
    then resized with ``F.interpolate(mode="bilinear", align_corners=False)``.
    Torchvision's ``Resize`` is deliberately not used β€” it resamples the PIL
    image with antialiasing on downscale, which yields different pixels and
    would shift the mask away from the numbers in the model card.

    ImageNet normalization (mean/std) is applied inside the exported graph via
    registered buffers, so callers must NOT normalize the input externally.
    """
    script_dir = Path(__file__).resolve().parent
    image = Image.open(script_dir / image_path).convert("RGB")
    # np.array (not np.asarray) returns a writable copy; torch.from_numpy warns
    # on a read-only array. The pixel values are identical either way.
    tensor = (
        torch.from_numpy(np.array(image))
        .permute(2, 0, 1)  # HWC β†’ CHW
        .float()
        .div(255.0)
        .unsqueeze(0)  # [1, 3, H, W]
    )
    resized = F.interpolate(
        tensor,
        size=(INPUT_SIZE, INPUT_SIZE),
        mode="bilinear",
        align_corners=False,
    )
    return resized.contiguous(), image


def box_prompt_from_native(
    bbox_xywh: tuple[float, float, float, float],
    original_size: tuple[int, int],
    input_size: int = INPUT_SIZE,
) -> torch.Tensor:
    """Convert a COCO (x, y, w, h) box in native-image pixel space to a model
    box prompt in input_size-pixel space.

    Scales each axis independently and clamps to the input square, matching the
    prompt scaling the evaluation pipeline applies to ground-truth boxes.

    Returns a [1, 1, 2, 2] tensor with two corner points (top-left, bottom-right).
    This matches the MobileSAM prompt encoder convention:
    point_coords.view(1, 4) β†’ [x_tl, y_tl, x_br, y_br] as a box.
    """
    orig_w, orig_h = original_size
    x, y, w, h = bbox_xywh
    scale_x = input_size / orig_w
    scale_y = input_size / orig_h
    x1, y1 = max(0.0, x * scale_x), max(0.0, y * scale_y)
    x2 = min(float(input_size), (x + w) * scale_x)
    y2 = min(float(input_size), (y + h) * scale_y)
    return torch.tensor(
        [[[[x1, y1], [x2, y2]]]],
        dtype=torch.float32,
    ).contiguous()  # [1, 1, 2, 2]


# ── Inference ──────────────────────────────────────────────────────────────────
def run_inference(
    method,
    image_tensor: torch.Tensor,
    point_coords: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Run forward pass. Returns (low_res_masks, iou_predictions)."""
    outputs = method.execute([image_tensor, point_coords])
    low_res_masks = outputs[0]    # [1, 3, 256, 256]
    iou_predictions = outputs[1]  # [1, 3]
    return low_res_masks, iou_predictions


# ── Postprocessing ─────────────────────────────────────────────────────────────
def postprocess(
    low_res_masks: torch.Tensor,
    iou_predictions: torch.Tensor,
    target_size: int = INPUT_SIZE,
) -> torch.Tensor:
    """Select best mask proposal, upsample to target_size, and threshold.

    Runs outside the exported graph, as it does in the pipeline, so that PT2E
    canonicalization of argmax / gather / bilinear cannot affect the result.

    Steps:
    1. Pick proposal index with highest IoU score.
    2. Bilinear-upsample the selected logit map from 256x256 to target_size.
    3. Threshold at 0 to produce a binary mask ([1, H, W], bool).
    """
    best_idx = int(iou_predictions[0].argmax().item())
    best_logit = low_res_masks[0, best_idx].unsqueeze(0).unsqueeze(0)  # [1, 1, 256, 256]
    if best_logit.shape[-1] != target_size:
        best_logit = F.interpolate(
            best_logit,
            size=(target_size, target_size),
            mode="bilinear",
            align_corners=False,
        )
    return (best_logit.squeeze(0) > 0).bool()  # [1, H, W]


def mask_to_native(binary_mask: torch.Tensor, native_size: tuple[int, int]) -> np.ndarray:
    """Resize the binary mask back to the image's native resolution.

    Nearest-neighbour, as in the pipeline: a mask is a label map, and bilinear
    resampling would invent intermediate values that are not foreground or
    background.
    """
    native_w, native_h = native_size
    resized = F.interpolate(
        binary_mask.float().unsqueeze(0),  # [1, 1, H, W]
        size=(native_h, native_w),
        mode="nearest",
    )
    return resized.squeeze(0).squeeze(0).numpy().astype(bool)


# ── Save Results ───────────────────────────────────────────────────────────────
def save_results(
    original_image: Image.Image,
    mask_native: np.ndarray,
    box_prompt: torch.Tensor,
    iou_predictions: torch.Tensor,
    best_idx: int,
) -> None:
    """Save a color mask overlay (sample_output.png) and segmentation.json."""
    script_dir = Path(__file__).resolve().parent

    # Overlay at native resolution, so the saved image is the input image with
    # the mask drawn on it rather than a stretched 1024x1024 copy.
    overlay = np.array(original_image).astype(np.float32)
    overlay[mask_native, 0] = overlay[mask_native, 0] * 0.5          # reduce red
    overlay[mask_native, 1] = overlay[mask_native, 1] * 0.5 + 127.5  # boost green
    overlay[mask_native, 2] = overlay[mask_native, 2] * 0.5          # reduce blue
    output_image = Image.fromarray(overlay.astype(np.uint8))

    # Draw the box prompt, mapped back from 1024-pixel space to native pixels.
    native_w, native_h = original_image.size
    (px1, py1), (px2, py2) = box_prompt[0, 0].tolist()
    sx, sy = native_w / INPUT_SIZE, native_h / INPUT_SIZE
    draw = ImageDraw.Draw(output_image)
    draw.rectangle(
        [px1 * sx, py1 * sy, px2 * sx, py2 * sy],
        outline=(255, 165, 0),  # orange box prompt indicator
        width=4,
    )

    output_image_path = script_dir / "sample_output.png"
    output_image.save(output_image_path)
    print(f"Saved mask overlay: {output_image_path}")

    # Save structured segmentation result
    mask_pixels = int(mask_native.sum())
    total_pixels = int(mask_native.size)
    iou_scores = iou_predictions[0].tolist()
    segmentation_result = {
        "model": "mobile_sam_int8_executorch",
        "input_size": INPUT_SIZE,
        "native_size": [native_w, native_h],
        "box_prompt": [px1, py1, px2, py2],
        "best_proposal_index": best_idx,
        "iou_scores": [round(s, 4) for s in iou_scores],
        "best_iou_score": round(iou_scores[best_idx], 4),
        "mask_pixels": mask_pixels,
        "total_pixels": total_pixels,
        "mask_coverage_pct": round(100.0 * mask_pixels / total_pixels, 2),
    }

    json_path = script_dir / "segmentation.json"
    with open(json_path, "w") as f:
        json.dump(segmentation_result, f, indent=2)
    print(f"Saved segmentation results: {json_path}")


# ── Main ───────────────────────────────────────────────────────────────────────
def main() -> None:
    print("Loading MobileSAM INT8 ExecuTorch model...")
    method = load_model(MODEL_PATH)

    print(f"Preprocessing image: {IMAGE_PATH}")
    image_tensor, original_image = preprocess(IMAGE_PATH)
    point_coords = box_prompt_from_native(PROMPT_BOX, original_image.size)

    print("Running inference...")
    low_res_masks, iou_predictions = run_inference(method, image_tensor, point_coords)

    print("Postprocessing outputs...")
    binary_mask = postprocess(low_res_masks, iou_predictions, target_size=INPUT_SIZE)
    mask_native = mask_to_native(binary_mask, original_image.size)

    best_idx = int(iou_predictions[0].argmax().item())
    iou_scores = iou_predictions[0].tolist()
    mask_pixels = int(mask_native.sum())
    coverage_pct = 100.0 * mask_pixels / mask_native.size

    print("\n── Segmentation Results ───────────────────────────────────────")
    print(f"  Best proposal:   index {best_idx}  (IoU score: {iou_scores[best_idx]:.4f})")
    print(f"  All IoU scores:  {[round(s, 4) for s in iou_scores]}")
    print(f"  Mask pixels:     {mask_pixels:,} / {mask_native.size:,}")
    print(f"  Mask coverage:   {coverage_pct:.2f}%")
    print("───────────────────────────────────────────────────────────────\n")

    save_results(original_image, mask_native, point_coords, iou_predictions, best_idx)
    print("Done.")


if __name__ == "__main__":
    main()