cuibinge commited on
Commit
61308d7
·
verified ·
1 Parent(s): da5bc3e

Sync strict polygon dataset importer

Browse files
Files changed (1) hide show
  1. scripts/train_sampoly_polygon.py +385 -0
scripts/train_sampoly_polygon.py ADDED
@@ -0,0 +1,385 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Train the SAMPolyBuild-style polygon head for marine ecological features."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import csv
7
+ import json
8
+ import math
9
+ import random
10
+ import sys
11
+ from dataclasses import asdict, dataclass
12
+ from pathlib import Path
13
+
14
+ import torch
15
+ from PIL import Image, ImageDraw
16
+ from torch import Tensor, nn
17
+ import torch.nn.functional as F
18
+ from torch.utils.data import DataLoader, Dataset
19
+ from torchvision.transforms import functional as TF
20
+
21
+ ROOT = Path(__file__).resolve().parents[1]
22
+ if str(ROOT) not in sys.path:
23
+ sys.path.append(str(ROOT))
24
+
25
+ from marine_sampoly_polygon_model import MarineSAMPolyModel, PolygonModelConfig, cyclic_l1_distance # noqa: E402
26
+
27
+
28
+ IMAGE_SUFFIXES = {".jpg", ".jpeg", ".png", ".tif", ".tiff"}
29
+
30
+
31
+ @dataclass
32
+ class PolygonMetrics:
33
+ images: int
34
+ mask_iou: float
35
+ vertex_iou: float
36
+ boundary_iou: float
37
+ polygon_positive_queries: int
38
+
39
+
40
+ def parse_args() -> argparse.Namespace:
41
+ parser = argparse.ArgumentParser(description=__doc__)
42
+ parser.add_argument("--data-root", required=True)
43
+ parser.add_argument("--vit-weights", required=True)
44
+ parser.add_argument("--convnext-weights", required=True)
45
+ parser.add_argument("--output-dir", required=True)
46
+ parser.add_argument("--epochs", type=int, default=20)
47
+ parser.add_argument("--imgsz", type=int, default=512)
48
+ parser.add_argument("--batch", type=int, default=1)
49
+ parser.add_argument("--workers", type=int, default=0)
50
+ parser.add_argument("--device", default="cuda")
51
+ parser.add_argument("--lr", type=float, default=1e-4)
52
+ parser.add_argument("--backbone-lr", type=float, default=1e-5)
53
+ parser.add_argument("--weight-decay", type=float, default=1e-4)
54
+ parser.add_argument("--num-queries", type=int, default=100)
55
+ parser.add_argument("--vertices-per-polygon", type=int, default=32)
56
+ parser.add_argument("--decoder-layers", type=int, default=4)
57
+ parser.add_argument("--decoder-heads", type=int, default=8)
58
+ parser.add_argument("--mask-weight", type=float, default=2.0)
59
+ parser.add_argument("--boundary-weight", type=float, default=1.0)
60
+ parser.add_argument("--vertex-weight", type=float, default=1.0)
61
+ parser.add_argument("--polygon-weight", type=float, default=2.0)
62
+ parser.add_argument("--no-object-weight", type=float, default=0.1)
63
+ parser.add_argument("--threshold", type=float, default=0.5)
64
+ parser.add_argument("--seed", type=int, default=0)
65
+ parser.add_argument("--no-pretrained", action="store_true")
66
+ parser.add_argument("--data-parallel", action="store_true")
67
+ return parser.parse_args()
68
+
69
+
70
+ def image_paths_for_split(root: Path, split: str) -> list[Path]:
71
+ image_dir = root / "images" / split
72
+ return sorted(path for path in image_dir.iterdir() if path.suffix.lower() in IMAGE_SUFFIXES)
73
+
74
+
75
+ def mask_path_for_image(root: Path, image_path: Path, split: str) -> Path:
76
+ mask_dir = root / "masks" / split
77
+ for suffix in (".png", ".tif", ".tiff", ".jpg", ".jpeg"):
78
+ path = mask_dir / f"{image_path.stem}{suffix}"
79
+ if path.exists():
80
+ return path
81
+ return mask_dir / f"{image_path.stem}.png"
82
+
83
+
84
+ def coco_ann_path(root: Path, split: str) -> Path:
85
+ for name in (f"{split}.json", "ann.json", "annotations.json"):
86
+ path = root / "annotations" / name
87
+ if path.exists():
88
+ return path
89
+ return root / "annotations" / f"{split}.json"
90
+
91
+
92
+ def resample_polygon(points: list[tuple[float, float]], n: int) -> list[tuple[float, float]]:
93
+ if len(points) < 3:
94
+ return [(0.0, 0.0)] * n
95
+ closed = points + [points[0]]
96
+ lengths = []
97
+ total = 0.0
98
+ for a, b in zip(closed[:-1], closed[1:]):
99
+ seg = math.hypot(b[0] - a[0], b[1] - a[1])
100
+ lengths.append(seg)
101
+ total += seg
102
+ if total <= 0:
103
+ return [points[0]] * n
104
+ samples = []
105
+ cursor = 0.0
106
+ seg_idx = 0
107
+ seg_start = 0.0
108
+ for k in range(n):
109
+ target = total * k / n
110
+ while seg_idx < len(lengths) - 1 and seg_start + lengths[seg_idx] < target:
111
+ seg_start += lengths[seg_idx]
112
+ seg_idx += 1
113
+ a = closed[seg_idx]
114
+ b = closed[seg_idx + 1]
115
+ t = (target - seg_start) / max(lengths[seg_idx], 1e-8)
116
+ samples.append((a[0] + (b[0] - a[0]) * t, a[1] + (b[1] - a[1]) * t))
117
+ cursor = target
118
+ return samples
119
+
120
+
121
+ def draw_targets(polygons: list[Tensor], size: int) -> tuple[Tensor, Tensor, Tensor]:
122
+ mask_img = Image.new("L", (size, size), 0)
123
+ boundary_img = Image.new("L", (size, size), 0)
124
+ vertex_img = Image.new("L", (size, size), 0)
125
+ mask_draw = ImageDraw.Draw(mask_img)
126
+ boundary_draw = ImageDraw.Draw(boundary_img)
127
+ vertex_draw = ImageDraw.Draw(vertex_img)
128
+ for poly in polygons:
129
+ pts = [(float(x * size), float(y * size)) for x, y in poly.tolist()]
130
+ if len(pts) < 3:
131
+ continue
132
+ mask_draw.polygon(pts, fill=255)
133
+ boundary_draw.line(pts + [pts[0]], fill=255, width=max(2, size // 128))
134
+ radius = max(1, size // 192)
135
+ for x, y in pts:
136
+ vertex_draw.ellipse((x - radius, y - radius, x + radius, y + radius), fill=255)
137
+ mask = TF.to_tensor(mask_img)
138
+ boundary = TF.to_tensor(boundary_img)
139
+ vertex = TF.to_tensor(vertex_img)
140
+ return mask, boundary, vertex
141
+
142
+
143
+ class PolygonDataset(Dataset):
144
+ def __init__(self, root: str | Path, split: str, image_size: int, vertices_per_polygon: int) -> None:
145
+ self.root = Path(root)
146
+ self.split = split
147
+ self.image_size = image_size
148
+ self.vertices_per_polygon = vertices_per_polygon
149
+ self.images = image_paths_for_split(self.root, split)
150
+ self.coco_by_file = self._load_coco_polygons()
151
+
152
+ def _load_coco_polygons(self) -> dict[str, list[list[tuple[float, float]]]]:
153
+ path = coco_ann_path(self.root, self.split)
154
+ if not path.exists():
155
+ return {}
156
+ data = json.loads(path.read_text(encoding="utf-8"))
157
+ image_by_id = {item["id"]: item for item in data.get("images", [])}
158
+ grouped: dict[str, list[list[tuple[float, float]]]] = {}
159
+ for ann in data.get("annotations", []):
160
+ image = image_by_id.get(ann.get("image_id"))
161
+ if not image:
162
+ continue
163
+ width = float(image.get("width", 1))
164
+ height = float(image.get("height", 1))
165
+ for seg in ann.get("segmentation", []):
166
+ if not isinstance(seg, list) or len(seg) < 6:
167
+ continue
168
+ pts = [(seg[i] / width, seg[i + 1] / height) for i in range(0, len(seg), 2)]
169
+ grouped.setdefault(Path(image["file_name"]).name, []).append(pts)
170
+ return grouped
171
+
172
+ def __len__(self) -> int:
173
+ return len(self.images)
174
+
175
+ def __getitem__(self, idx: int) -> dict[str, object]:
176
+ image_path = self.images[idx]
177
+ image = Image.open(image_path).convert("RGB")
178
+ image = image.resize((self.image_size, self.image_size), Image.BILINEAR)
179
+ tensor = TF.to_tensor(image)
180
+
181
+ raw_polygons = self.coco_by_file.get(image_path.name, [])
182
+ polygons = [
183
+ torch.tensor(resample_polygon(poly, self.vertices_per_polygon), dtype=torch.float32).clamp(0, 1)
184
+ for poly in raw_polygons
185
+ ]
186
+ mask_path = mask_path_for_image(self.root, image_path, self.split)
187
+ if not polygons and mask_path.exists():
188
+ mask = Image.open(mask_path).convert("L").resize((self.image_size, self.image_size), Image.NEAREST)
189
+ mask_tensor = (TF.to_tensor(mask) > 0.5).float()
190
+ boundary = torch.zeros_like(mask_tensor)
191
+ vertex = torch.zeros_like(mask_tensor)
192
+ else:
193
+ mask_tensor, boundary, vertex = draw_targets(polygons, self.image_size)
194
+ return {
195
+ "image": tensor,
196
+ "mask": mask_tensor,
197
+ "boundary": boundary,
198
+ "vertex": vertex,
199
+ "polygons": polygons,
200
+ "path": str(image_path),
201
+ }
202
+
203
+
204
+ def collate(batch: list[dict[str, object]]) -> dict[str, object]:
205
+ return {
206
+ "image": torch.stack([item["image"] for item in batch]), # type: ignore[index]
207
+ "mask": torch.stack([item["mask"] for item in batch]), # type: ignore[index]
208
+ "boundary": torch.stack([item["boundary"] for item in batch]), # type: ignore[index]
209
+ "vertex": torch.stack([item["vertex"] for item in batch]), # type: ignore[index]
210
+ "polygons": [item["polygons"] for item in batch],
211
+ "path": [item["path"] for item in batch],
212
+ }
213
+
214
+
215
+ def dice_loss(logits: Tensor, target: Tensor) -> Tensor:
216
+ prob = logits.sigmoid()
217
+ inter = (prob * target).sum(dim=(1, 2, 3))
218
+ denom = prob.sum(dim=(1, 2, 3)) + target.sum(dim=(1, 2, 3))
219
+ return (1 - (2 * inter + 1) / (denom + 1)).mean()
220
+
221
+
222
+ def polygon_loss(poly_logits: Tensor, polygons: Tensor, targets: list[list[Tensor]], no_object_weight: float) -> Tensor:
223
+ object_target = torch.zeros_like(poly_logits)
224
+ losses = []
225
+ for b, target_list in enumerate(targets):
226
+ n = min(len(target_list), polygons.shape[1])
227
+ if n == 0:
228
+ continue
229
+ target = torch.stack(target_list[:n]).to(polygons.device)
230
+ object_target[b, :n] = 1.0
231
+ losses.append(cyclic_l1_distance(polygons[b, :n], target).mean())
232
+ weight = torch.where(object_target > 0, torch.ones_like(object_target), torch.full_like(object_target, no_object_weight))
233
+ objectness = F.binary_cross_entropy_with_logits(poly_logits, object_target, weight=weight)
234
+ if losses:
235
+ return objectness + torch.stack(losses).mean()
236
+ return objectness
237
+
238
+
239
+ def total_loss(outputs: dict[str, Tensor], batch: dict[str, object], args: argparse.Namespace) -> tuple[Tensor, dict[str, float]]:
240
+ mask = batch["mask"].to(outputs["mask_logits"].device) # type: ignore[union-attr]
241
+ boundary = batch["boundary"].to(outputs["mask_logits"].device) # type: ignore[union-attr]
242
+ vertex = batch["vertex"].to(outputs["mask_logits"].device) # type: ignore[union-attr]
243
+ mask_loss = F.binary_cross_entropy_with_logits(outputs["mask_logits"], mask) + dice_loss(outputs["mask_logits"], mask)
244
+ boundary_loss = F.binary_cross_entropy_with_logits(outputs["boundary_logits"], boundary) + dice_loss(
245
+ outputs["boundary_logits"], boundary
246
+ )
247
+ vertex_loss = F.binary_cross_entropy_with_logits(outputs["vertex_logits"], vertex) + dice_loss(
248
+ outputs["vertex_logits"], vertex
249
+ )
250
+ poly_loss = polygon_loss(outputs["poly_logits"], outputs["polygons"], batch["polygons"], args.no_object_weight) # type: ignore[arg-type]
251
+ loss = (
252
+ args.mask_weight * mask_loss
253
+ + args.boundary_weight * boundary_loss
254
+ + args.vertex_weight * vertex_loss
255
+ + args.polygon_weight * poly_loss
256
+ )
257
+ return loss, {
258
+ "mask_loss": float(mask_loss.detach()),
259
+ "boundary_loss": float(boundary_loss.detach()),
260
+ "vertex_loss": float(vertex_loss.detach()),
261
+ "polygon_loss": float(poly_loss.detach()),
262
+ }
263
+
264
+
265
+ def binary_iou(logits: Tensor, target: Tensor, threshold: float) -> float:
266
+ pred = logits.sigmoid() >= threshold
267
+ truth = target >= 0.5
268
+ inter = (pred & truth).sum().item()
269
+ union = (pred | truth).sum().item()
270
+ return float(inter / union) if union else 1.0
271
+
272
+
273
+ def evaluate(model: nn.Module, loader: DataLoader, device: torch.device, threshold: float) -> PolygonMetrics:
274
+ model.eval()
275
+ mask_ious = []
276
+ boundary_ious = []
277
+ vertex_ious = []
278
+ pos_queries = 0
279
+ with torch.no_grad():
280
+ for batch in loader:
281
+ image = batch["image"].to(device)
282
+ out = model(image)
283
+ mask = batch["mask"].to(device)
284
+ boundary = batch["boundary"].to(device)
285
+ vertex = batch["vertex"].to(device)
286
+ mask_ious.append(binary_iou(out["mask_logits"], mask, threshold))
287
+ boundary_ious.append(binary_iou(out["boundary_logits"], boundary, threshold))
288
+ vertex_ious.append(binary_iou(out["vertex_logits"], vertex, threshold))
289
+ pos_queries += int((out["poly_logits"].sigmoid() >= threshold).sum().item())
290
+ return PolygonMetrics(
291
+ images=len(loader.dataset),
292
+ mask_iou=sum(mask_ious) / max(len(mask_ious), 1),
293
+ vertex_iou=sum(vertex_ious) / max(len(vertex_ious), 1),
294
+ boundary_iou=sum(boundary_ious) / max(len(boundary_ious), 1),
295
+ polygon_positive_queries=pos_queries,
296
+ )
297
+
298
+
299
+ def train() -> None:
300
+ args = parse_args()
301
+ random.seed(args.seed)
302
+ torch.manual_seed(args.seed)
303
+ output_dir = Path(args.output_dir)
304
+ output_dir.mkdir(parents=True, exist_ok=True)
305
+ device = torch.device(args.device if torch.cuda.is_available() else "cpu")
306
+ train_ds = PolygonDataset(args.data_root, "train", args.imgsz, args.vertices_per_polygon)
307
+ val_ds = PolygonDataset(args.data_root, "val", args.imgsz, args.vertices_per_polygon)
308
+ test_ds = PolygonDataset(args.data_root, "test", args.imgsz, args.vertices_per_polygon)
309
+ if len(train_ds) == 0:
310
+ raise RuntimeError(
311
+ "No polygon-trainable samples found. Provide masks/{split} or COCO polygon annotations; "
312
+ "bbox-only datasets are intentionally unsupported for this head."
313
+ )
314
+ train_loader = DataLoader(train_ds, batch_size=args.batch, shuffle=True, num_workers=args.workers, collate_fn=collate)
315
+ val_loader = DataLoader(val_ds, batch_size=args.batch, shuffle=False, num_workers=args.workers, collate_fn=collate)
316
+ test_loader = DataLoader(test_ds, batch_size=args.batch, shuffle=False, num_workers=args.workers, collate_fn=collate)
317
+
318
+ model = MarineSAMPolyModel(
319
+ PolygonModelConfig(
320
+ vit_weights=args.vit_weights,
321
+ convnext_weights=args.convnext_weights,
322
+ pretrained=not args.no_pretrained,
323
+ num_queries=args.num_queries,
324
+ vertices_per_polygon=args.vertices_per_polygon,
325
+ decoder_layers=args.decoder_layers,
326
+ decoder_heads=args.decoder_heads,
327
+ )
328
+ ).to(device)
329
+ if args.data_parallel and torch.cuda.device_count() > 1:
330
+ model = nn.DataParallel(model)
331
+ raw = model.module if isinstance(model, nn.DataParallel) else model
332
+ optimizer = torch.optim.AdamW(
333
+ [
334
+ {"params": [p for p in raw.backbone.parameters() if p.requires_grad], "lr": args.backbone_lr},
335
+ {"params": raw.head.parameters(), "lr": args.lr},
336
+ ],
337
+ weight_decay=args.weight_decay,
338
+ )
339
+
340
+ history_path = output_dir / "history.csv"
341
+ best_iou = -1.0
342
+ with history_path.open("w", newline="", encoding="utf-8") as fp:
343
+ writer = csv.DictWriter(
344
+ fp,
345
+ fieldnames=["epoch", "loss", "mask_loss", "boundary_loss", "vertex_loss", "polygon_loss", "val_mask_iou"],
346
+ )
347
+ writer.writeheader()
348
+ for epoch in range(1, args.epochs + 1):
349
+ raw.train()
350
+ rows = []
351
+ for batch in train_loader:
352
+ image = batch["image"].to(device)
353
+ optimizer.zero_grad(set_to_none=True)
354
+ outputs = model(image)
355
+ loss, items = total_loss(outputs, batch, args)
356
+ loss.backward()
357
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
358
+ optimizer.step()
359
+ rows.append({"loss": float(loss.detach()), **items})
360
+ val = evaluate(model, val_loader, device, args.threshold)
361
+ row = {
362
+ "epoch": epoch,
363
+ "loss": sum(r["loss"] for r in rows) / max(len(rows), 1),
364
+ "mask_loss": sum(r["mask_loss"] for r in rows) / max(len(rows), 1),
365
+ "boundary_loss": sum(r["boundary_loss"] for r in rows) / max(len(rows), 1),
366
+ "vertex_loss": sum(r["vertex_loss"] for r in rows) / max(len(rows), 1),
367
+ "polygon_loss": sum(r["polygon_loss"] for r in rows) / max(len(rows), 1),
368
+ "val_mask_iou": val.mask_iou,
369
+ }
370
+ writer.writerow(row)
371
+ fp.flush()
372
+ print(json.dumps(row), flush=True)
373
+ if val.mask_iou > best_iou:
374
+ best_iou = val.mask_iou
375
+ torch.save({"model": raw.state_dict(), "args": vars(args), "val_metrics": asdict(val)}, output_dir / "best.pt")
376
+ torch.save({"model": raw.state_dict(), "args": vars(args), "val_metrics": asdict(val)}, output_dir / "last.pt")
377
+ best = torch.load(output_dir / "best.pt", map_location=device)
378
+ raw.load_state_dict(best["model"])
379
+ test = evaluate(model, test_loader, device, args.threshold)
380
+ (output_dir / "test_metrics.json").write_text(json.dumps(asdict(test), indent=2), encoding="utf-8")
381
+ print(json.dumps({"test": asdict(test)}, indent=2), flush=True)
382
+
383
+
384
+ if __name__ == "__main__":
385
+ train()