"""Refine coarse target/obstacle masks with Segment Anything box prompts.""" from __future__ import annotations import argparse import json import sys from pathlib import Path import cv2 import numpy as np from PIL import Image, ImageOps def read_rgb(path: str | Path) -> np.ndarray: return np.array(ImageOps.exif_transpose(Image.open(path)).convert('RGB')) def read_mask(path: str | Path, shape: tuple[int, int]) -> np.ndarray: mask = np.array(ImageOps.exif_transpose(Image.open(path)).convert('L')) > 127 h, w = shape if mask.shape != (h, w): raise ValueError( f'Mask/RGB raster mismatch for {path}: mask={mask.shape}, rgb={(h, w)}. ' 'Refusing to resize because this can hide EXIF-orientation misalignment.' ) return mask def save_mask(path: str | Path, mask: np.ndarray) -> None: Image.fromarray((mask.astype(np.uint8) * 255)).save(path) def component_boxes(mask: np.ndarray, keep: int, min_area: int, pad: int) -> np.ndarray: num, labels, stats, _ = cv2.connectedComponentsWithStats(mask.astype(np.uint8), connectivity=8) boxes = [] areas = [] h, w = mask.shape for idx in range(1, num): area = int(stats[idx, cv2.CC_STAT_AREA]) if area < min_area: continue x = int(stats[idx, cv2.CC_STAT_LEFT]) y = int(stats[idx, cv2.CC_STAT_TOP]) bw = int(stats[idx, cv2.CC_STAT_WIDTH]) bh = int(stats[idx, cv2.CC_STAT_HEIGHT]) boxes.append([max(0, x - pad), max(0, y - pad), min(w - 1, x + bw + pad), min(h - 1, y + bh + pad)]) areas.append(area) if not boxes: return np.empty((0, 4), dtype=np.float32) order = np.argsort(np.array(areas))[::-1][:keep] return np.array([boxes[i] for i in order], dtype=np.float32) def refine_one_mask(predictor, mask: np.ndarray, keep: int, min_area: int, pad: int) -> np.ndarray: boxes = component_boxes(mask, keep=keep, min_area=min_area, pad=pad) if boxes.size == 0: return mask import torch transformed = predictor.transform.apply_boxes_torch( torch.as_tensor(boxes, dtype=torch.float32, device=predictor.device), mask.shape, ) masks, scores, _ = predictor.predict_torch( point_coords=None, point_labels=None, boxes=transformed, multimask_output=True, ) refined = np.zeros_like(mask, dtype=bool) masks_np = masks.detach().cpu().numpy() scores_np = scores.detach().cpu().numpy() for i in range(masks_np.shape[0]): best = int(np.argmax(scores_np[i])) refined |= masks_np[i, best].astype(bool) return refined def overlay(rgb: np.ndarray, masks: list[tuple[np.ndarray, tuple[int, int, int], float]]) -> np.ndarray: out = rgb.astype(np.float32).copy() for mask, color, alpha in masks: if mask.any(): out[mask] = out[mask] * (1.0 - alpha) + np.array(color, dtype=np.float32) * alpha return np.clip(out, 0, 255).astype(np.uint8) def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description='Refine binary masks with SAM using mask-derived box prompts.') parser.add_argument('--image', required=True) parser.add_argument('--target-mask', required=True) parser.add_argument('--obstacle-mask', required=True) parser.add_argument('--output-dir', required=True) parser.add_argument('--sam-repo', default='../amodal/segment-anything', help='Path containing the segment_anything package.') parser.add_argument('--sam-checkpoint', required=True) parser.add_argument('--sam-model-type', choices=['vit_h', 'vit_l', 'vit_b', 'default'], default='vit_h') parser.add_argument('--device', default='auto') parser.add_argument('--keep-target-components', type=int, default=8) parser.add_argument('--keep-obstacle-components', type=int, default=4) parser.add_argument('--min-area', type=int, default=64) return parser def main() -> None: args = build_parser().parse_args() rgb = read_rgb(args.image) shape = rgb.shape[:2] target = read_mask(args.target_mask, shape) obstacle = read_mask(args.obstacle_mask, shape) sam_repo = Path(args.sam_repo).resolve() checkpoint = Path(args.sam_checkpoint).resolve() if not checkpoint.exists(): raise FileNotFoundError(f'SAM checkpoint not found: {checkpoint}') if not sam_repo.exists(): raise FileNotFoundError(f'SAM repo not found: {sam_repo}') sys.path.insert(0, str(sam_repo)) import torch from segment_anything import SamPredictor, sam_model_registry if args.device == 'auto': device = 'cuda' if torch.cuda.is_available() else 'cpu' else: device = args.device sam = sam_model_registry[args.sam_model_type](checkpoint=str(checkpoint)).to(device=device) predictor = SamPredictor(sam) predictor.set_image(rgb) refined_target = refine_one_mask( predictor, target, keep=args.keep_target_components, min_area=args.min_area, pad=args.box_pad, ) refined_obstacle = refine_one_mask( predictor, obstacle, keep=args.keep_obstacle_components, min_area=args.min_area, pad=args.box_pad, ) output_dir = Path(args.output_dir) output_dir.mkdir(parents=True, exist_ok=True) target_path = output_dir / 'target_visible_mask_sam.png' obstacle_path = output_dir / 'obstacle_mask_sam.png' overlay_path = output_dir / 'sam_refine_overlay.png' manifest_path = output_dir / 'sam_refine_manifest.json' save_mask(target_path, refined_target) save_mask(obstacle_path, refined_obstacle) Image.fromarray(overlay(rgb, [ (refined_target, (0, 220, 80), 0.45), (refined_obstacle, (255, 60, 20), 0.55), ])).save(overlay_path) manifest_path.write_text(json.dumps({ 'image': args.image, 'sam_repo': str(sam_repo), 'sam_checkpoint': str(checkpoint), 'sam_model_type': args.sam_model_type, 'device': device, 'target_output': str(target_path), 'obstacle_output': str(obstacle_path), 'overlay': str(overlay_path), }, indent=2), encoding='utf-8') print(f'Wrote SAM-refined masks to {output_dir}') if __name__ == '__main__': main()