| """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() |
|
|