"""Train the SAMPolyBuild-style polygon head for marine ecological features.""" from __future__ import annotations import argparse import csv import json import math import random import sys from dataclasses import asdict, dataclass from pathlib import Path import torch import numpy as np from PIL import Image, ImageDraw from torch import Tensor, nn import torch.nn.functional as F from torch.utils.data import DataLoader, Dataset from torchvision.transforms import functional as TF ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.append(str(ROOT)) from marine_sampoly_polygon_model import MarineSAMPolyModel, PolygonModelConfig, cyclic_l1_distance # noqa: E402 IMAGE_SUFFIXES = {".jpg", ".jpeg", ".png", ".tif", ".tiff"} @dataclass class PolygonMetrics: images: int mask_iou: float vertex_iou: float boundary_iou: float polygon_positive_queries: int def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--data-root", required=True) parser.add_argument("--vit-weights", required=True) parser.add_argument("--convnext-weights", required=True) parser.add_argument("--output-dir", required=True) parser.add_argument("--epochs", type=int, default=20) parser.add_argument("--imgsz", type=int, default=512) parser.add_argument("--batch", type=int, default=1) parser.add_argument("--workers", type=int, default=0) parser.add_argument("--device", default="cuda") parser.add_argument("--lr", type=float, default=1e-4) parser.add_argument("--backbone-lr", type=float, default=1e-5) parser.add_argument("--weight-decay", type=float, default=1e-4) parser.add_argument("--num-queries", type=int, default=100) parser.add_argument("--vertices-per-polygon", type=int, default=32) parser.add_argument("--decoder-layers", type=int, default=4) parser.add_argument("--decoder-heads", type=int, default=8) parser.add_argument("--mask-weight", type=float, default=2.0) parser.add_argument("--boundary-weight", type=float, default=1.0) parser.add_argument("--vertex-weight", type=float, default=1.0) parser.add_argument("--polygon-weight", type=float, default=2.0) parser.add_argument("--no-object-weight", type=float, default=0.1) parser.add_argument("--threshold", type=float, default=0.5) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--no-pretrained", action="store_true") parser.add_argument("--data-parallel", action="store_true") return parser.parse_args() def image_paths_for_split(root: Path, split: str) -> list[Path]: image_dir = root / "images" / split return sorted(path for path in image_dir.iterdir() if path.suffix.lower() in IMAGE_SUFFIXES) def mask_path_for_image(root: Path, image_path: Path, split: str) -> Path: mask_dir = root / "masks" / split for suffix in (".png", ".tif", ".tiff", ".jpg", ".jpeg"): path = mask_dir / f"{image_path.stem}{suffix}" if path.exists(): return path return mask_dir / f"{image_path.stem}.png" def coco_ann_path(root: Path, split: str) -> Path: for name in (f"{split}.json", "ann.json", "annotations.json"): path = root / "annotations" / name if path.exists(): return path return root / "annotations" / f"{split}.json" def resample_polygon(points: list[tuple[float, float]], n: int) -> list[tuple[float, float]]: if len(points) < 3: return [(0.0, 0.0)] * n closed = points + [points[0]] lengths = [] total = 0.0 for a, b in zip(closed[:-1], closed[1:]): seg = math.hypot(b[0] - a[0], b[1] - a[1]) lengths.append(seg) total += seg if total <= 0: return [points[0]] * n samples = [] cursor = 0.0 seg_idx = 0 seg_start = 0.0 for k in range(n): target = total * k / n while seg_idx < len(lengths) - 1 and seg_start + lengths[seg_idx] < target: seg_start += lengths[seg_idx] seg_idx += 1 a = closed[seg_idx] b = closed[seg_idx + 1] t = (target - seg_start) / max(lengths[seg_idx], 1e-8) samples.append((a[0] + (b[0] - a[0]) * t, a[1] + (b[1] - a[1]) * t)) cursor = target return samples def draw_targets(polygons: list[Tensor], size: int) -> tuple[Tensor, Tensor, Tensor]: mask_img = Image.new("L", (size, size), 0) boundary_img = Image.new("L", (size, size), 0) vertex_img = Image.new("L", (size, size), 0) mask_draw = ImageDraw.Draw(mask_img) boundary_draw = ImageDraw.Draw(boundary_img) vertex_draw = ImageDraw.Draw(vertex_img) for poly in polygons: pts = [(float(x * size), float(y * size)) for x, y in poly.tolist()] if len(pts) < 3: continue mask_draw.polygon(pts, fill=255) boundary_draw.line(pts + [pts[0]], fill=255, width=max(2, size // 128)) radius = max(1, size // 192) for x, y in pts: vertex_draw.ellipse((x - radius, y - radius, x + radius, y + radius), fill=255) mask = TF.to_tensor(mask_img) boundary = TF.to_tensor(boundary_img) vertex = TF.to_tensor(vertex_img) return mask, boundary, vertex def polygons_from_binary_mask(mask: Image.Image, vertices_per_polygon: int) -> list[Tensor]: mask_np = np.asarray(mask.convert("L")) binary = (mask_np > 0).astype(np.uint8) if binary.max() == 0: return [] try: import cv2 contours, _ = cv2.findContours(binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) polygons = [] h, w = binary.shape min_area = max(4.0, 0.0005 * h * w) for contour in contours: if cv2.contourArea(contour) < min_area: continue pts = [(float(p[0][0]) / max(w - 1, 1), float(p[0][1]) / max(h - 1, 1)) for p in contour] sampled = resample_polygon(pts, vertices_per_polygon) polygons.append(torch.tensor(sampled, dtype=torch.float32).clamp(0, 1)) return polygons except Exception: ys, xs = np.where(binary > 0) if len(xs) == 0: return [] h, w = binary.shape x1, x2 = xs.min() / max(w - 1, 1), xs.max() / max(w - 1, 1) y1, y2 = ys.min() / max(h - 1, 1), ys.max() / max(h - 1, 1) sampled = resample_polygon([(x1, y1), (x2, y1), (x2, y2), (x1, y2)], vertices_per_polygon) return [torch.tensor(sampled, dtype=torch.float32).clamp(0, 1)] def boundary_from_mask(mask_tensor: Tensor) -> Tensor: pooled_max = F.max_pool2d(mask_tensor.unsqueeze(0), kernel_size=3, stride=1, padding=1) pooled_min = -F.max_pool2d(-mask_tensor.unsqueeze(0), kernel_size=3, stride=1, padding=1) return (pooled_max - pooled_min).squeeze(0).clamp(0, 1) class PolygonDataset(Dataset): def __init__(self, root: str | Path, split: str, image_size: int, vertices_per_polygon: int) -> None: self.root = Path(root) self.split = split self.image_size = image_size self.vertices_per_polygon = vertices_per_polygon self.images = image_paths_for_split(self.root, split) self.coco_by_file = self._load_coco_polygons() def _load_coco_polygons(self) -> dict[str, list[list[tuple[float, float]]]]: path = coco_ann_path(self.root, self.split) if not path.exists(): return {} data = json.loads(path.read_text(encoding="utf-8")) image_by_id = {item["id"]: item for item in data.get("images", [])} grouped: dict[str, list[list[tuple[float, float]]]] = {} for ann in data.get("annotations", []): image = image_by_id.get(ann.get("image_id")) if not image: continue width = float(image.get("width", 1)) height = float(image.get("height", 1)) for seg in ann.get("segmentation", []): if not isinstance(seg, list) or len(seg) < 6: continue pts = [(seg[i] / width, seg[i + 1] / height) for i in range(0, len(seg), 2)] grouped.setdefault(Path(image["file_name"]).name, []).append(pts) return grouped def __len__(self) -> int: return len(self.images) def __getitem__(self, idx: int) -> dict[str, object]: image_path = self.images[idx] image = Image.open(image_path).convert("RGB") image = image.resize((self.image_size, self.image_size), Image.BILINEAR) tensor = TF.to_tensor(image) raw_polygons = self.coco_by_file.get(image_path.name, []) polygons = [ torch.tensor(resample_polygon(poly, self.vertices_per_polygon), dtype=torch.float32).clamp(0, 1) for poly in raw_polygons ] mask_path = mask_path_for_image(self.root, image_path, self.split) if not polygons and mask_path.exists(): mask = Image.open(mask_path).convert("L").resize((self.image_size, self.image_size), Image.NEAREST) mask_tensor = (TF.to_tensor(mask) > 0.5).float() polygons = polygons_from_binary_mask(mask, self.vertices_per_polygon) if polygons: _, boundary, vertex = draw_targets(polygons, self.image_size) else: boundary = boundary_from_mask(mask_tensor) vertex = torch.zeros_like(mask_tensor) else: mask_tensor, boundary, vertex = draw_targets(polygons, self.image_size) return { "image": tensor, "mask": mask_tensor, "boundary": boundary, "vertex": vertex, "polygons": polygons, "path": str(image_path), } def collate(batch: list[dict[str, object]]) -> dict[str, object]: return { "image": torch.stack([item["image"] for item in batch]), # type: ignore[index] "mask": torch.stack([item["mask"] for item in batch]), # type: ignore[index] "boundary": torch.stack([item["boundary"] for item in batch]), # type: ignore[index] "vertex": torch.stack([item["vertex"] for item in batch]), # type: ignore[index] "polygons": [item["polygons"] for item in batch], "path": [item["path"] for item in batch], } def dice_loss(logits: Tensor, target: Tensor) -> Tensor: prob = logits.sigmoid() inter = (prob * target).sum(dim=(1, 2, 3)) denom = prob.sum(dim=(1, 2, 3)) + target.sum(dim=(1, 2, 3)) return (1 - (2 * inter + 1) / (denom + 1)).mean() def polygon_loss(poly_logits: Tensor, polygons: Tensor, targets: list[list[Tensor]], no_object_weight: float) -> Tensor: object_target = torch.zeros_like(poly_logits) losses = [] for b, target_list in enumerate(targets): n = min(len(target_list), polygons.shape[1]) if n == 0: continue target = torch.stack(target_list[:n]).to(polygons.device) object_target[b, :n] = 1.0 losses.append(cyclic_l1_distance(polygons[b, :n], target).mean()) weight = torch.where(object_target > 0, torch.ones_like(object_target), torch.full_like(object_target, no_object_weight)) objectness = F.binary_cross_entropy_with_logits(poly_logits, object_target, weight=weight) if losses: return objectness + torch.stack(losses).mean() return objectness def total_loss(outputs: dict[str, Tensor], batch: dict[str, object], args: argparse.Namespace) -> tuple[Tensor, dict[str, float]]: mask = batch["mask"].to(outputs["mask_logits"].device) # type: ignore[union-attr] boundary = batch["boundary"].to(outputs["mask_logits"].device) # type: ignore[union-attr] vertex = batch["vertex"].to(outputs["mask_logits"].device) # type: ignore[union-attr] mask_loss = F.binary_cross_entropy_with_logits(outputs["mask_logits"], mask) + dice_loss(outputs["mask_logits"], mask) boundary_loss = F.binary_cross_entropy_with_logits(outputs["boundary_logits"], boundary) + dice_loss( outputs["boundary_logits"], boundary ) vertex_loss = F.binary_cross_entropy_with_logits(outputs["vertex_logits"], vertex) + dice_loss( outputs["vertex_logits"], vertex ) poly_loss = polygon_loss(outputs["poly_logits"], outputs["polygons"], batch["polygons"], args.no_object_weight) # type: ignore[arg-type] loss = ( args.mask_weight * mask_loss + args.boundary_weight * boundary_loss + args.vertex_weight * vertex_loss + args.polygon_weight * poly_loss ) return loss, { "mask_loss": float(mask_loss.detach()), "boundary_loss": float(boundary_loss.detach()), "vertex_loss": float(vertex_loss.detach()), "polygon_loss": float(poly_loss.detach()), } def binary_iou(logits: Tensor, target: Tensor, threshold: float) -> float: pred = logits.sigmoid() >= threshold truth = target >= 0.5 inter = (pred & truth).sum().item() union = (pred | truth).sum().item() return float(inter / union) if union else 1.0 def evaluate(model: nn.Module, loader: DataLoader, device: torch.device, threshold: float) -> PolygonMetrics: model.eval() mask_ious = [] boundary_ious = [] vertex_ious = [] pos_queries = 0 with torch.no_grad(): for batch in loader: image = batch["image"].to(device) out = model(image) mask = batch["mask"].to(device) boundary = batch["boundary"].to(device) vertex = batch["vertex"].to(device) mask_ious.append(binary_iou(out["mask_logits"], mask, threshold)) boundary_ious.append(binary_iou(out["boundary_logits"], boundary, threshold)) vertex_ious.append(binary_iou(out["vertex_logits"], vertex, threshold)) pos_queries += int((out["poly_logits"].sigmoid() >= threshold).sum().item()) return PolygonMetrics( images=len(loader.dataset), mask_iou=sum(mask_ious) / max(len(mask_ious), 1), vertex_iou=sum(vertex_ious) / max(len(vertex_ious), 1), boundary_iou=sum(boundary_ious) / max(len(boundary_ious), 1), polygon_positive_queries=pos_queries, ) def train() -> None: args = parse_args() random.seed(args.seed) torch.manual_seed(args.seed) output_dir = Path(args.output_dir) output_dir.mkdir(parents=True, exist_ok=True) device = torch.device(args.device if torch.cuda.is_available() else "cpu") train_ds = PolygonDataset(args.data_root, "train", args.imgsz, args.vertices_per_polygon) val_ds = PolygonDataset(args.data_root, "val", args.imgsz, args.vertices_per_polygon) test_ds = PolygonDataset(args.data_root, "test", args.imgsz, args.vertices_per_polygon) if len(train_ds) == 0: raise RuntimeError( "No polygon-trainable samples found. Provide masks/{split} or COCO polygon annotations; " "bbox-only datasets are intentionally unsupported for this head." ) train_loader = DataLoader(train_ds, batch_size=args.batch, shuffle=True, num_workers=args.workers, collate_fn=collate) val_loader = DataLoader(val_ds, batch_size=args.batch, shuffle=False, num_workers=args.workers, collate_fn=collate) test_loader = DataLoader(test_ds, batch_size=args.batch, shuffle=False, num_workers=args.workers, collate_fn=collate) model = MarineSAMPolyModel( PolygonModelConfig( vit_weights=args.vit_weights, convnext_weights=args.convnext_weights, pretrained=not args.no_pretrained, num_queries=args.num_queries, vertices_per_polygon=args.vertices_per_polygon, decoder_layers=args.decoder_layers, decoder_heads=args.decoder_heads, ) ).to(device) if args.data_parallel and torch.cuda.device_count() > 1: model = nn.DataParallel(model) raw = model.module if isinstance(model, nn.DataParallel) else model optimizer = torch.optim.AdamW( [ {"params": [p for p in raw.backbone.parameters() if p.requires_grad], "lr": args.backbone_lr}, {"params": raw.head.parameters(), "lr": args.lr}, ], weight_decay=args.weight_decay, ) history_path = output_dir / "history.csv" best_iou = -1.0 with history_path.open("w", newline="", encoding="utf-8") as fp: writer = csv.DictWriter( fp, fieldnames=["epoch", "loss", "mask_loss", "boundary_loss", "vertex_loss", "polygon_loss", "val_mask_iou"], ) writer.writeheader() for epoch in range(1, args.epochs + 1): raw.train() rows = [] for batch in train_loader: image = batch["image"].to(device) optimizer.zero_grad(set_to_none=True) outputs = model(image) loss, items = total_loss(outputs, batch, args) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() rows.append({"loss": float(loss.detach()), **items}) val = evaluate(model, val_loader, device, args.threshold) row = { "epoch": epoch, "loss": sum(r["loss"] for r in rows) / max(len(rows), 1), "mask_loss": sum(r["mask_loss"] for r in rows) / max(len(rows), 1), "boundary_loss": sum(r["boundary_loss"] for r in rows) / max(len(rows), 1), "vertex_loss": sum(r["vertex_loss"] for r in rows) / max(len(rows), 1), "polygon_loss": sum(r["polygon_loss"] for r in rows) / max(len(rows), 1), "val_mask_iou": val.mask_iou, } writer.writerow(row) fp.flush() print(json.dumps(row), flush=True) if val.mask_iou > best_iou: best_iou = val.mask_iou torch.save({"model": raw.state_dict(), "args": vars(args), "val_metrics": asdict(val)}, output_dir / "best.pt") torch.save({"model": raw.state_dict(), "args": vars(args), "val_metrics": asdict(val)}, output_dir / "last.pt") best = torch.load(output_dir / "best.pt", map_location=device) raw.load_state_dict(best["model"]) test = evaluate(model, test_loader, device, args.threshold) (output_dir / "test_metrics.json").write_text(json.dumps(asdict(test), indent=2), encoding="utf-8") print(json.dumps({"test": asdict(test)}, indent=2), flush=True) if __name__ == "__main__": train()