Download example.py from Arm/efficient-sam-s-int8-xnnpack-executorch: direct link, hf CLI and curl.
- Browser
- Download file 5.91 kB
-
https://huggingface.co/Arm/efficient-sam-s-int8-xnnpack-executorch/resolve/main/example.py
- Command line
-
hf download hf://Arm/efficient-sam-s-int8-xnnpack-executorch/example.py
-
curl -L -o example.py https://huggingface.co/Arm/efficient-sam-s-int8-xnnpack-executorch/resolve/main/example.py
5.91 kB
| """Run box-prompted segmentation with EfficientSAM-S and ExecuTorch.""" | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| from executorch.runtime import Runtime | |
| from PIL import Image, ImageDraw | |
| from torchvision import transforms | |
| from torchvision.transforms import InterpolationMode | |
| MODEL_DIR = Path(__file__).resolve().parent | |
| MODEL_PATH = MODEL_DIR / "efficient_sam_s_graviton_executorch_optimized.pte" | |
| SAMPLE_IMAGE_PATH = MODEL_DIR / "sample_input.jpg" | |
| INPUT_SIZE = 1024 | |
| SAMPLE_BOX_PROMPT = [410, 510, 600, 830] | |
| def required_file(path: Path) -> Path: | |
| if not path.is_file(): | |
| raise FileNotFoundError(f"required model file not found: {path}") | |
| return path | |
| def load_model() -> object: | |
| runtime = Runtime.get() | |
| if not runtime.backend_registry.is_available("XnnpackBackend"): | |
| available = ", ".join(runtime.backend_registry.registered_backend_names) or "none" | |
| raise RuntimeError(f"XnnpackBackend is unavailable; registered backends: {available}") | |
| program = runtime.load_program(required_file(MODEL_PATH)) | |
| if program.method_names != {"forward"}: | |
| raise RuntimeError(f"expected PTE method 'forward', found: {sorted(program.method_names)}") | |
| return program.load_method("forward") | |
| def preprocess(image_path: Path) -> torch.Tensor: | |
| transform = transforms.Compose( | |
| [ | |
| transforms.Resize( | |
| (INPUT_SIZE, INPUT_SIZE), | |
| interpolation=InterpolationMode.BILINEAR, | |
| antialias=True, | |
| ), | |
| transforms.ToTensor(), | |
| ] | |
| ) | |
| with Image.open(required_file(image_path)) as image: | |
| return transform(image.convert("RGB")).unsqueeze(0).contiguous() | |
| def make_box_prompt(box: list[int]) -> tuple[torch.Tensor, torch.Tensor]: | |
| x1, y1, x2, y2 = box | |
| point_coords = torch.tensor([[[[x1, y1], [x2, y2]]]], dtype=torch.float32) | |
| point_labels = torch.tensor([[[2, 3]]], dtype=torch.float32) | |
| return point_coords, point_labels | |
| def validate_box(box: list[int]) -> list[int]: | |
| x1, y1, x2, y2 = box | |
| if not (0 <= x1 < x2 <= INPUT_SIZE and 0 <= y1 < y2 <= INPUT_SIZE): | |
| raise ValueError( | |
| f"box must satisfy 0 <= x1 < x2 <= {INPUT_SIZE} and " | |
| f"0 <= y1 < y2 <= {INPUT_SIZE}; received {box}" | |
| ) | |
| return box | |
| def postprocess(raw_output: torch.Tensor, box: list[int]) -> tuple[np.ndarray, dict[str, object]]: | |
| if tuple(raw_output.shape) != (1, INPUT_SIZE, INPUT_SIZE): | |
| raise ValueError( | |
| f"expected output shape (1, {INPUT_SIZE}, {INPUT_SIZE}), " | |
| f"found {tuple(raw_output.shape)}" | |
| ) | |
| if not torch.isfinite(raw_output).all(): | |
| raise ValueError("model output contains non-finite mask values") | |
| mask = (raw_output.squeeze(0) > 0.5).cpu().numpy().astype(np.uint8) | |
| pixel_count = int(mask.sum()) | |
| total_pixels = int(mask.size) | |
| result = { | |
| "pixel_count": pixel_count, | |
| "total_pixels": total_pixels, | |
| "coverage_pct": round(100.0 * pixel_count / total_pixels, 2), | |
| "box_prompt": box, | |
| } | |
| return mask, result | |
| def create_overlay(mask: np.ndarray, image_path: Path, box: list[int]) -> Image.Image: | |
| with Image.open(required_file(image_path)) as image: | |
| source = image.convert("RGB") | |
| source_width, source_height = source.size | |
| source_mask = Image.fromarray(mask * 255).resize( | |
| source.size, | |
| resample=Image.Resampling.NEAREST, | |
| ) | |
| source_mask_array = np.asarray(source_mask, dtype=np.uint8) > 0 | |
| mask_rgba = np.zeros((source_height, source_width, 4), dtype=np.uint8) | |
| mask_rgba[source_mask_array] = [0, 200, 0, 200] | |
| overlay = Image.fromarray(mask_rgba) | |
| composited = Image.alpha_composite(source.convert("RGBA"), overlay).convert("RGB") | |
| x1, y1, x2, y2 = box | |
| scaled_box = [ | |
| round(x1 * source_width / INPUT_SIZE), | |
| round(y1 * source_height / INPUT_SIZE), | |
| round(x2 * source_width / INPUT_SIZE), | |
| round(y2 * source_height / INPUT_SIZE), | |
| ] | |
| draw = ImageDraw.Draw(composited) | |
| draw.rectangle(scaled_box, outline=(255, 0, 0), width=3) | |
| return composited | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument( | |
| "--image", | |
| type=Path, | |
| default=SAMPLE_IMAGE_PATH, | |
| help="input image (default: sample_input.jpg)", | |
| ) | |
| parser.add_argument( | |
| "--box", | |
| nargs=4, | |
| type=int, | |
| metavar=("X1", "Y1", "X2", "Y2"), | |
| default=SAMPLE_BOX_PROMPT, | |
| help="box prompt in resized 1024x1024 coordinates (default: sample box)", | |
| ) | |
| parser.add_argument("--output-dir", type=Path, help="optional output directory") | |
| return parser.parse_args() | |
| def main() -> None: | |
| args = parse_args() | |
| image_path = required_file(args.image) | |
| box = validate_box(args.box) | |
| image_tensor = preprocess(image_path) | |
| point_coords, point_labels = make_box_prompt(box) | |
| method = load_model() | |
| raw_output = method.execute([image_tensor, point_coords, point_labels])[0] | |
| mask, result = postprocess(raw_output, box) | |
| print(f"Input: {image_path.name}; tensor shape: {tuple(image_tensor.shape)}") | |
| print(f"Box prompt: {box}") | |
| print( | |
| f"Segmented pixels: {result['pixel_count']:,} / {result['total_pixels']:,} " | |
| f"({result['coverage_pct']:.2f}%)" | |
| ) | |
| if args.output_dir is not None: | |
| args.output_dir.mkdir(parents=True, exist_ok=True) | |
| overlay_path = args.output_dir / "sample_output.png" | |
| result_path = args.output_dir / "segmentation.json" | |
| create_overlay(mask, image_path, box).save(overlay_path) | |
| result_path.write_text(json.dumps(result, indent=2) + "\n", encoding="utf-8") | |
| print(f"Saved: {overlay_path}") | |
| print(f"Saved: {result_path}") | |
| if __name__ == "__main__": | |
| main() | |