""" Fine-tune AdaptFormer on Delhi change-detection tiles. Priority improvements (v2): 1. Validate logit→prob (2-class softmax, change = channel 1) + print stats 2. Auto threshold search (fixed grid + score quantiles); freeze thr with best ckpt 3. Positive-tile oversampling + Focal+Dice / CE+Dice losses 4. Lazy tile index (pair_idx, x, y) — no duplicated arrays in RAM 5. Track Precision / Recall / IoU (not just F1) 6. Save Before/After/GT/Pred/Prob panels each epoch 7. Stronger aug + ReduceLROnPlateau Run: python scripts/build_delhi_cd_splits.py --min-change-frac 0.001 --stratify python scripts/finetune_adaptformer.py --delhi-cd data/delhi_cd --preset v2 """ from __future__ import annotations import argparse import json import random import sys import time from pathlib import Path import numpy as np from PIL import Image ROOT = Path(__file__).resolve().parent.parent sys.path.insert(0, str(ROOT)) from app.evaluation.delhi_eval import DelhiEvalNotReady, dummy_delhi_pairs, iter_delhi_pairs # noqa: E402 from app.evaluation.metrics import binary_metrics # noqa: E402 _MODEL_ID = "deepang/adaptformer-LEVIR-CD" _TILE = 256 _DAY5_PRESET = { "epochs": 30, "lr": 3e-5, "batch_size": 2, "augment": True, "stride": 128, "early_stop_patience": 8, "loss": "bce_dice", "exclude_empty": False, "full_resize": False, "pos_oversample": 1, "visualize": False, } _FIX_PRESET = { "epochs": 25, "lr": 1e-4, "batch_size": 2, "augment": True, "stride": 64, "early_stop_patience": 6, "loss": "ce", "exclude_empty": True, "full_resize": True, "min_change_frac": 0.001, "pos_oversample": 1, "visualize": False, } # Claude priority plan — target Val F1 0.5+ and closer Val/Test gap _V2_PRESET = { "epochs": 20, "lr": 5e-5, "batch_size": 2, "augment": True, "stride": 64, "early_stop_patience": 7, "loss": "focal_dice", "exclude_empty": True, "full_resize": True, "min_change_frac": 0.001, "pos_oversample": 3, "min_tile_change": 0.005, "visualize": True, "scheduler": True, } # Recall / FN-focused follow-up (Test F1>0.50, R>0.45 target) _V3_PRESET = { "epochs": 20, "lr": 3e-5, "batch_size": 2, "augment": True, "stride": 64, "early_stop_patience": 6, "loss": "tversky", "exclude_empty": True, "full_resize": True, "min_change_frac": 0.001, "pos_oversample": 4, "min_tile_change": 0.01, "change_centered": True, "visualize": False, # enable with --visualize; keeps CPU train faster "scheduler": True, "thr_min": 0.2, "thr_max": 0.7, "thr_objective": "fbeta", # F_beta=1.5 favors recall on val thr pick "warm_start": "runs/finetune_v2/20260716_210208/best", } # v4: ONE change vs frozen v3 — stronger positive-tile sampling only (loss/aug unchanged) _V4_PRESET = { "epochs": 12, "lr": 2e-5, "batch_size": 2, "augment": True, "stride": 64, "early_stop_patience": 5, "loss": "tversky", # unchanged from v3 "exclude_empty": True, "full_resize": True, "min_change_frac": 0.001, "pos_oversample": 6, # ↑ from 4 "min_tile_change": 0.02, # ↑ from 0.01 — stricter positive tiles "change_centered": True, "pos_only": True, # NEW: train batches from change tiles only "visualize": False, "scheduler": True, "thr_min": 0.10, # match recommended ops sweep window "thr_max": 0.40, "thr_objective": "fbeta", "warm_start": "models/adaptformer_delhi/v3_frozen", } # Wednesday plan: training_failure_diagnosis fixes + hard-neg retention # Target: test F1 > 0.60 _WED_PRESET = { "epochs": 20, "lr": 5e-5, "batch_size": 2, "augment": True, "stride": 64, "early_stop_patience": 7, "loss": "ce", # CE + pos_weight (diagnosis #2) "exclude_empty": True, # drop empty real GT (diagnosis #1) "keep_hard_neg": True, # but keep mined hn_* empty tiles "full_resize": True, # full-image 256 resize (diagnosis #4) "min_change_frac": 0.001, "pos_oversample": 4, # oversample change tiles "min_tile_change": 0.005, "change_centered": True, "visualize": False, "scheduler": True, "thr_min": 0.05, "thr_max": 0.50, "thr_objective": "f1", # val-calibrate + freeze threshold "warm_start": "models/adaptformer_delhi/v3_frozen", } def _try_torch(): try: import torch from torch.utils.data import DataLoader, Dataset, WeightedRandomSampler from transformers import AutoImageProcessor, AutoModel return torch, DataLoader, Dataset, WeightedRandomSampler, AutoImageProcessor, AutoModel except ImportError as exc: raise SystemExit( "PyTorch + transformers required for fine-tuning. " f"Import error: {exc}" ) from exc def _augment_triplet(b: np.ndarray, a: np.ndarray, g: np.ndarray, rng: random.Random): """H/V flips, 90° rotations, mild brightness/contrast (same transform on both dates).""" if rng.random() < 0.5: b = np.ascontiguousarray(np.flip(b, axis=1)) a = np.ascontiguousarray(np.flip(a, axis=1)) g = np.ascontiguousarray(np.flip(g, axis=1)) if rng.random() < 0.5: b = np.ascontiguousarray(np.flip(b, axis=0)) a = np.ascontiguousarray(np.flip(a, axis=0)) g = np.ascontiguousarray(np.flip(g, axis=0)) if rng.random() < 0.5: k = rng.randint(1, 3) b = np.ascontiguousarray(np.rot90(b, k)) a = np.ascontiguousarray(np.rot90(a, k)) g = np.ascontiguousarray(np.rot90(g, k)) if rng.random() < 0.5: # Shared photometric jitter so relative change is preserved alpha = 1.0 + rng.uniform(-0.15, 0.15) # contrast beta = rng.uniform(-12.0, 12.0) # brightness b = np.clip(b.astype(np.float32) * alpha + beta, 0, 255).astype(np.uint8) a = np.clip(a.astype(np.float32) * alpha + beta, 0, 255).astype(np.uint8) return b, a, g class DelhiTileDataset: """Lazy tile dataset: stores (pair_index, kind, x, y) instead of cropped arrays.""" KIND_FULL = 0 KIND_CROP = 1 KIND_CENTER = 2 def __init__( self, pairs: list[tuple], crop_size: int = _TILE, train: bool = True, stride: int | None = None, augment: bool = False, seed: int = 0, full_resize: bool = True, min_tile_change: float = 0.0, pos_oversample: int = 1, change_centered: bool = False, drop_empty_tiles: bool = False, pos_only: bool = False, max_train_tiles: int | None = 160, ): _torch, _dl, Dataset, _wrs, _proc, _model = _try_torch() self.pairs = pairs # keep source images once self.crop_size = crop_size self.train = train self.augment = bool(augment and train) self._rng = random.Random(seed) self.min_tile_change = float(min_tile_change) self.pos_oversample = max(1, int(pos_oversample)) self.change_centered = bool(change_centered and train) self.pos_only = bool(pos_only and train) # index entries: (pair_i, kind, x, y, is_positive) self.index: list[tuple[int, int, int, int, bool]] = [] use_stride = stride if stride is not None else (crop_size // 2 if train else crop_size) use_stride = max(16, min(crop_size, int(use_stride))) n_pos = n_neg = n_center = 0 for pi, (before, after, gt, pair_id) in enumerate(pairs): h, w = before.shape[:2] is_hard_neg = str(pair_id).startswith("hn_") if full_resize or h < crop_size or w < crop_size: gr = np.array(Image.fromarray(gt).resize( (crop_size, crop_size), resample=Image.NEAREST)) frac = float((gr > 127).mean()) # Empty GT (hard negatives) must count as neg even when min_tile_change=0 is_pos = frac > 0.0 and frac >= self.min_tile_change # Hard-neg empty tiles are kept even when drop_empty_tiles is on if not (train and drop_empty_tiles and not is_pos and not is_hard_neg): self.index.append((pi, self.KIND_FULL, 0, 0, is_pos)) n_pos += int(is_pos) n_neg += int(not is_pos) if h >= crop_size and w >= crop_size: coords = set() for y in range(0, h - crop_size + 1, use_stride): for x in range(0, w - crop_size + 1, use_stride): coords.add((x, y)) coords.add((max(0, w - crop_size), max(0, h - crop_size))) for x, y in coords: tile_gt = gt[y:y + crop_size, x:x + crop_size] frac = float((tile_gt > 127).mean()) is_pos = frac > 0.0 and frac >= self.min_tile_change # Drop empty / near-empty crops when exclude_empty path is active # (but keep hard-negative empty tiles so FP patterns are learned) if train and drop_empty_tiles and not is_pos and not is_hard_neg: continue self.index.append((pi, self.KIND_CROP, x, y, is_pos)) n_pos += int(is_pos) n_neg += int(not is_pos) # Change-centered crops (priority: more tiles on actual buildings) if self.change_centered: for cx, cy in _gt_change_centers(gt, max_centers=16 if self.pos_only else 10): x0 = int(np.clip(cx - crop_size // 2, 0, max(0, w - crop_size))) y0 = int(np.clip(cy - crop_size // 2, 0, max(0, h - crop_size))) # Small jitter for diversity if train: x0 = int(np.clip(x0 + self._rng.randint(-24, 24), 0, max(0, w - crop_size))) y0 = int(np.clip(y0 + self._rng.randint(-24, 24), 0, max(0, h - crop_size))) tile_gt = gt[y0:y0 + crop_size, x0:x0 + crop_size] frac = float((tile_gt > 127).mean()) if frac < max(self.min_tile_change, 0.002): continue self.index.append((pi, self.KIND_CENTER, x0, y0, True)) n_pos += 1 n_center += 1 # Drop negatives entirely when pos_only (no-change tiles cannot dominate) if train and self.pos_only: before_n = len(self.index) self.index = [e for e in self.index if e[4]] print(f" pos_only: kept {len(self.index)}/{before_n} positive tiles", flush=True) # Expand positive indices for oversampling (simple list multiply) if train and self.pos_oversample > 1: extras = [e for e in self.index if e[4]] for _ in range(self.pos_oversample - 1): self.index.extend(extras) # Soft cap so CPU/6GB GPU runs stay tractable. Prefer real Delhi tiles # over synthetic so the 2000-tile set cannot dominate the cap. max_tiles = int(max_train_tiles) if (train and max_train_tiles) else None if max_tiles and len(self.index) > max_tiles: delhi = [e for e in self.index if not str(self.pairs[e[0]][3]).startswith("synth_")] synth = [e for e in self.index if str(self.pairs[e[0]][3]).startswith("synth_")] self._rng.shuffle(delhi) self._rng.shuffle(synth) reserve_synth = min(len(synth), int(max_tiles * 0.35)) if synth else 0 n_delhi = min(len(delhi), max_tiles - reserve_synth) n_synth = min(len(synth), max_tiles - n_delhi) keep = delhi[:n_delhi] + synth[:n_synth] self._rng.shuffle(keep) self.index = keep print(f" Capped train tiles to {len(self.index)} " f"(delhi={n_delhi}, synth={n_synth})", flush=True) self.n_pos_unique = n_pos self.n_neg_unique = n_neg print(f" Dataset({'train' if train else 'eval'}): index={len(self.index)} " f"(pos~{n_pos}, neg~{n_neg}, centered~{n_center}, " f"oversamplex{self.pos_oversample})", flush=True) outer = self class _Inner(Dataset): def __len__(inner_self): return len(outer.index) def __getitem__(inner_self, idx): pi, kind, x, y, _is_pos = outer.index[idx] before, after, gt, _ = outer.pairs[pi] cs = outer.crop_size if kind == outer.KIND_FULL: b = np.array(Image.fromarray(before).resize((cs, cs))) a = np.array(Image.fromarray(after).resize((cs, cs))) g = np.array(Image.fromarray(gt).resize((cs, cs), resample=Image.NEAREST)) else: b = before[y:y + cs, x:x + cs].copy() a = after[y:y + cs, x:x + cs].copy() g = gt[y:y + cs, x:x + cs].copy() if outer.augment: b, a, g = _augment_triplet(b, a, g, outer._rng) return b, a, (g > 127).astype(np.float32) self._dataset = _Inner() @property def samples(self): """Backward-compat length alias.""" return self.index def torch_dataset(self): return self._dataset def sampler_weights(self) -> list[float]: """Per-index weights: positives/centered heavier, Delhi vs synthetic balanced.""" domains = [] n_delhi = n_synth = 0 for pi, _kind, _x, _y, _is_pos in self.index: is_synth = str(self.pairs[pi][3]).startswith("synth_") domains.append(is_synth) n_synth += int(is_synth) n_delhi += int(not is_synth) w = [] for (_pi, kind, _x, _y, is_pos), is_synth in zip(self.index, domains): base = float(self.pos_oversample) if is_pos else 1.0 if kind == self.KIND_CENTER: base *= 2.0 if n_delhi and n_synth: # Equal domain mass so 2000 synthetic tiles cannot drown Delhi. base *= (0.5 / n_synth) if is_synth else (0.5 / n_delhi) w.append(base) return w def _gt_change_centers(gt: np.ndarray, max_centers: int = 10) -> list[tuple[int, int]]: """Centroids of GT change blobs — used for positive-focused crops.""" import cv2 binary = (gt > 127).astype(np.uint8) n, _lab, stats, centroids = cv2.connectedComponentsWithStats(binary, connectivity=8) centers = [] for i in range(1, n): area = int(stats[i, cv2.CC_STAT_AREA]) if area < 8: continue cx, cy = int(centroids[i][0]), int(centroids[i][1]) centers.append((area, cx, cy)) centers.sort(key=lambda t: -t[0]) return [(cx, cy) for _a, cx, cy in centers[:max_centers]] def _dice_loss(prob, target, eps: float = 1e-6): p = prob.reshape(-1) t = target.reshape(-1) inter = (p * t).sum() return 1.0 - (2.0 * inter + eps) / (p.sum() + t.sum() + eps) def _focal_loss(prob, target, gamma: float = 2.0, alpha: float = 0.75, eps: float = 1e-6): p = prob.clamp(eps, 1.0 - eps) pt = p * target + (1.0 - p) * (1.0 - target) w = alpha * target + (1.0 - alpha) * (1.0 - target) return (-(w * (1.0 - pt).pow(gamma) * pt.log())).mean() def _tversky_loss(prob, target, alpha: float = 0.3, beta: float = 0.7, eps: float = 1e-6): """Tversky: beta>alpha penalizes false negatives more (recall-oriented).""" p = prob.reshape(-1) t = target.reshape(-1) tp = (p * t).sum() fp = (p * (1.0 - t)).sum() fn = ((1.0 - p) * t).sum() return 1.0 - (tp + eps) / (tp + alpha * fp + beta * fn + eps) def _change_prob_from_logits(logits, torch): """Convert AdaptFormer logits → change probability. Confirmed on deepang/adaptformer-LEVIR-CD: logits are (N, 2, H, W) where channel 0 = no-change, channel 1 = change. Softmax last channel is correct. """ from app.model_inference import _logits_to_change_prob return _logits_to_change_prob(logits, torch) def _probe_output_scale(model, processor, device, pairs: list[tuple], n: int = 2) -> dict: """Print / return logit→prob stats to validate conversion (priority #1).""" torch, *_ = _try_torch() from PIL import Image as PILImage rows = [] for before, after, gt, pair_id in pairs[:n]: if before.shape[0] != _TILE or before.shape[1] != _TILE: before_r = np.array(Image.fromarray(before).resize((_TILE, _TILE))) after_r = np.array(Image.fromarray(after).resize((_TILE, _TILE))) gt_r = np.array(Image.fromarray(gt).resize((_TILE, _TILE), Image.NEAREST)) else: before_r, after_r, gt_r = before, after, gt inputs = processor( images=(PILImage.fromarray(before_r), PILImage.fromarray(after_r)), return_tensors="pt", ) inputs = {k: v.to(device) for k, v in inputs.items()} with torch.no_grad(): logits = model(**inputs).logits if logits.dim() == 3: logits = logits.unsqueeze(0) g = (gt_r > 127) gt_pos = float(g.mean()) row = { "pair_id": pair_id, "logits_shape": list(logits.shape), "logits_min": float(logits.min()), "logits_max": float(logits.max()), "logits_mean": float(logits.mean()), "gt_pos_frac": gt_pos, } if logits.shape[1] >= 2: sm = torch.softmax(logits, dim=1) p0 = sm[0, 0].cpu().numpy() p1 = sm[0, 1].cpu().numpy() row.update({ "ch0_mean": float(p0.mean()), "ch1_mean": float(p1.mean()), "ch1_min": float(p1.min()), "ch1_max": float(p1.max()), "ch1_on_gt": float(p1[g].mean()) if g.any() else None, "ch1_on_bg": float(p1[~g].mean()) if (~g).any() else None, "ch0_on_gt": float(p0[g].mean()) if g.any() else None, "change_channel": 1 if ( (p1[g].mean() if g.any() else 0) >= (p0[g].mean() if g.any() else 0) ) else 0, }) prob = _change_prob_from_logits(logits, torch).cpu().numpy() row.update({ "prob_min": float(prob.min()), "prob_max": float(prob.max()), "prob_mean": float(prob.mean()), "prob_on_gt": float(prob[g].mean()) if g.any() else None, "prob_on_bg": float(prob[~g].mean()) if (~g).any() else None, "pred_pos_at_0.5": float((prob >= 0.5).mean()), }) rows.append(row) print( f" [probe] {pair_id}: logits[{row['logits_min']:.2f},{row['logits_max']:.2f}] " f"prob[{row['prob_min']:.2e},{row['prob_max']:.4f}] mean={row['prob_mean']:.4f} " f"GT%={gt_pos:.3f} p@GT={row['prob_on_gt']} p@bg={row['prob_on_bg']} " f"pred%@0.5={row['pred_pos_at_0.5']:.3f}", flush=True, ) return {"n": len(rows), "pairs": rows} def _filter_empty( pairs: list[tuple], min_change_frac: float, *, keep_hard_neg: bool = True, ) -> list[tuple]: kept, dropped = [], [] for before, after, gt, pair_id in pairs: frac = float((gt > 127).mean()) if gt is not None else 0.0 is_hn = keep_hard_neg and str(pair_id).startswith("hn_") if frac >= min_change_frac or is_hn: kept.append((before, after, gt, pair_id)) else: dropped.append(pair_id) if dropped: print(f" Excluded {len(dropped)} empty/near-empty GT pairs: {dropped}") hn_kept = sum(1 for *_, pid in kept if str(pid).startswith("hn_")) if hn_kept: print(f" Kept {hn_kept} hard-negative (empty-GT) tiles for FP suppression") return kept def _class_balance_report(pairs: list[tuple], name: str) -> dict: fracs = [float((g > 127).mean()) for _, _, g, _ in pairs] report = { "split": name, "n_pairs": len(pairs), "change_frac_mean": round(float(np.mean(fracs)), 6) if fracs else 0.0, "change_frac_min": round(float(np.min(fracs)), 6) if fracs else 0.0, "change_frac_max": round(float(np.max(fracs)), 6) if fracs else 0.0, "bg_to_change_ratio": round( (1.0 - float(np.mean(fracs))) / max(float(np.mean(fracs)), 1e-6), 2 ) if fracs else None, } print( f" Imbalance[{name}]: change%={report['change_frac_mean']*100:.2f} " f"(min={report['change_frac_min']*100:.2f} max={report['change_frac_max']*100:.2f}) " f"bg:change~{report['bg_to_change_ratio']}:1", flush=True, ) return report def _resolve_pair_file(rel: str, pair_id: str, kind: str) -> Path | None: """Resolve a train path, including Priyanka absolute paths and labeling-pack PNGs.""" raw = Path(rel) candidates = [] if raw.is_absolute(): candidates.append(raw) candidates.append(Path(r"C:\Users\udayb\Downloads") / raw.name) # before5 (1).tif → before5.tif on this machine candidates.append(Path(r"C:\Users\udayb\Downloads") / raw.name.replace(" (1)", "")) else: candidates.append(ROOT / raw) pack = ROOT / "docs" / "delhi_eval" / "dda_labeling" / pair_id kind_name = {"before": "before.png", "after": "after.png", "gt": "gt_mask.png"}[kind] candidates.append(pack / kind_name) if kind == "gt": candidates.append(ROOT / "docs" / "delhi_eval" / "labels" / f"{pair_id}.png") for path in candidates: if path.is_file(): return path return None def _load_rgb_pair(before_rel: str, after_rel: str, gt_rel: str, pair_id: str) -> tuple: from app.evaluation.delhi_eval import _load_label, _load_rgb before_p = _resolve_pair_file(before_rel, pair_id, "before") after_p = _resolve_pair_file(after_rel, pair_id, "after") gt_p = _resolve_pair_file(gt_rel, pair_id, "gt") if not before_p or not after_p or not gt_p: raise FileNotFoundError( f"{pair_id}: missing files before={before_p} after={after_p} gt={gt_p}" ) before = _load_rgb(before_p) after = _load_rgb(after_p) gt = _load_label(gt_p) return before, after, gt, pair_id _SYNTHETIC_DIR_CANDIDATES = ( Path(r"C:\Users\Priyanka\Downloads\Synthetic_CD_dataset"), Path(r"C:\Users\udayb\Downloads\Synthetic_CD_dataset"), ROOT / "data" / "synthetic_cd", ROOT / "data" / "Synthetic_CD_dataset", ) def discover_synthetic_dir(explicit: str = "") -> Path | None: """Priyanka's gen_synthetic.py layout: before/*.png, after/*.png, mask/*.png.""" ordered = [] if explicit: ordered.append(Path(explicit)) ordered.extend(_SYNTHETIC_DIR_CANDIDATES) seen: set[str] = set() for path in ordered: key = str(path.resolve()) if path.exists() else str(path) if key in seen: continue seen.add(key) if (path / "before").is_dir() and (path / "mask").is_dir() and (path / "after").is_dir(): return path return None def _load_pairs_from_synthetic(dataset_dir: Path) -> list[tuple]: before_dir = dataset_dir / "before" after_dir = dataset_dir / "after" mask_dir = dataset_dir / "mask" pairs = [] for before_p in sorted(before_dir.glob("*.png")): name = before_p.name after_p = after_dir / name mask_p = mask_dir / name if not after_p.is_file() or not mask_p.is_file(): continue before = np.array(Image.open(before_p).convert("RGB")) after = np.array(Image.open(after_p).convert("RGB")) if after.shape[:2] != before.shape[:2]: after = np.array(Image.fromarray(after).resize((before.shape[1], before.shape[0]), Image.Resampling.LANCZOS)) gt = np.array(Image.open(mask_p).convert("L")) if gt.shape[:2] != before.shape[:2]: gt = np.array(Image.fromarray(gt).resize((before.shape[1], before.shape[0]), Image.Resampling.NEAREST)) pairs.append((before, after, gt, before_p.stem)) print(f" Synthetic GT: {len(pairs)} triplets from {dataset_dir}", flush=True) return pairs def _load_pairs_from_delhi_cd(delhi_cd: Path) -> tuple[list[tuple], list[tuple], list[tuple], dict]: split_path = delhi_cd / "split.json" if not split_path.is_file(): raise SystemExit( f"Missing {split_path}. Run: python scripts/build_delhi_cd_splits.py" ) summary = json.loads(split_path.read_text(encoding="utf-8")) loaded = {} for name in ("train", "val", "test"): man = delhi_cd / name / "manifest.json" if not man.is_file(): raise SystemExit(f"Missing {man}") rows = json.loads(man.read_text(encoding="utf-8")).get("pairs", []) loaded[name] = [] for p in rows: try: loaded[name].append( _load_rgb_pair(p["before_path"], p["after_path"], p["gt_mask"], p["pair_id"]) ) except Exception as exc: print(f" SKIP {p.get('pair_id')}: {exc}", flush=True) if not loaded[name]: raise SystemExit(f"No loadable pairs in {man}") split_info = { "train": [p[3] for p in loaded["train"]], "val": [p[3] for p in loaded["val"]], "test": [p[3] for p in loaded["test"]], "split": summary.get("split", "70/15/15"), "seed": summary.get("seed", 0), "source": str(delhi_cd), "stratified": summary.get("stratified", False), } return loaded["train"], loaded["val"], loaded["test"], split_info def _load_pairs(manifest: Path | None, dummy: bool) -> list[tuple]: if dummy: return [(b, a, g, pid) for b, a, g, pid, _, _ in dummy_delhi_pairs()] try: loaded = list(iter_delhi_pairs(manifest)) except DelhiEvalNotReady as exc: raise SystemExit(str(exc)) from exc labeled = [(b, a, g, pid) for b, a, g, pid, _, _ in loaded if g is not None] if labeled: return labeled raise SystemExit("No Delhi pairs with GT masks. Use --dummy for scaffold runs.") def _split_pairs(pairs: list[tuple], seed: int = 0, train_frac: float = 0.70, val_frac: float = 0.15): n = len(pairs) if n < 3: return pairs[: max(1, n - 1)], pairs[-1:], [] rng = random.Random(seed) idx = list(range(n)) rng.shuffle(idx) n_test = max(1, int(round(n * (1.0 - train_frac - val_frac)))) n_val = max(1, int(round(n * val_frac))) if n_test + n_val >= n: n_test = max(1, n // 5) n_val = max(1, n // 5) test_idx = set(idx[:n_test]) val_idx = set(idx[n_test:n_test + n_val]) train = [pairs[i] for i in range(n) if i not in test_idx and i not in val_idx] val = [pairs[i] for i in range(n) if i in val_idx] test = [pairs[i] for i in range(n) if i in test_idx] return train, val, test def _predict_mask(model, processor, device, before, after, threshold=0.5): torch, *_ = _try_torch() from PIL import Image as PILImage if before.shape[0] != _TILE or before.shape[1] != _TILE: before = np.array(Image.fromarray(before).resize((_TILE, _TILE))) after = np.array(Image.fromarray(after).resize((_TILE, _TILE))) inputs = processor( images=(PILImage.fromarray(before), PILImage.fromarray(after)), return_tensors="pt", ) inputs = {k: v.to(device) for k, v in inputs.items()} with torch.no_grad(): outputs = model(**inputs) score = _change_prob_from_logits(outputs.logits, torch).cpu().numpy().astype(np.float32) mask = (score >= threshold).astype(np.uint8) * 255 return mask, score def _resize_to_gt(arr: np.ndarray, gt: np.ndarray, nearest: bool = False) -> np.ndarray: if arr.shape[:2] == gt.shape[:2]: return arr from cv2 import resize, INTER_NEAREST, INTER_LINEAR return resize( arr, (gt.shape[1], gt.shape[0]), interpolation=INTER_NEAREST if nearest else INTER_LINEAR, ) def _eval_pairs(model, processor, device, pairs: list[tuple], threshold: float = 0.5) -> dict: f1s, precs, recs, ious, accs = [], [], [], [], [] scores_all = [] for before, after, gt, _pair_id in pairs: mask, score = _predict_mask(model, processor, device, before, after, threshold) score = _resize_to_gt(score, gt, nearest=False) mask = _resize_to_gt(mask, gt, nearest=True) scores_all.append(score) m = binary_metrics(mask, gt) f1s.append(m["f1"]) precs.append(m["precision"]) recs.append(m["recall"]) ious.append(m["iou"]) accs.append(m["pixelAccuracy"]) return { "mean_f1": round(float(np.mean(f1s)), 4) if f1s else 0.0, "mean_precision": round(float(np.mean(precs)), 4) if precs else 0.0, "mean_recall": round(float(np.mean(recs)), 4) if recs else 0.0, "mean_iou": round(float(np.mean(ious)), 4) if ious else 0.0, "mean_pixel_acc": round(float(np.mean(accs)), 4) if accs else 0.0, "n": len(f1s), "threshold": threshold, "mean_prob": round(float(np.mean([s.mean() for s in scores_all])), 6) if scores_all else 0.0, "max_prob": round(float(np.max([s.max() for s in scores_all])), 6) if scores_all else 0.0, } def _calibrate_threshold( model, processor, device, pairs: list[tuple], thr_min: float = 0.2, thr_max: float = 0.7, objective: str = "f1", ) -> tuple[float, float, dict]: """Sweep thresholds on val only; return (best_thr, best_f1, detail). Default search window 0.2–0.7 (Claude recall plan). objective: - f1: maximize mean F1 - fbeta: maximize F_1.5 (recall-oriented) then report F1 at that thr """ if not pairs: return 0.5, 0.0, {} scores, gts = [], [] for before, after, gt, _ in pairs: _mask, score = _predict_mask(model, processor, device, before, after, 0.5) score = _resize_to_gt(score, gt, nearest=False) scores.append(score.astype(np.float32)) gts.append(gt > 127) # Dense grid inside [thr_min, thr_max] plus baseline 0.5 grid = list(np.linspace(thr_min, thr_max, 26)) if 0.5 not in grid: grid.append(0.5) candidates = sorted({round(float(t), 6) for t in grid if thr_min - 1e-9 <= t <= thr_max + 1e-9}) best_thr, best_score, best_f1 = 0.5, -1.0, -1.0 best_row = {} beta = 1.5 sweep = [] for thr in candidates: f1s, precs, recs = [], [], [] for score, gt in zip(scores, gts): if not gt.any(): continue m = score >= thr tp = int((m & gt).sum()) fp = int((m & ~gt).sum()) fn = int((~m & gt).sum()) p = 0.0 if (tp + fp) == 0 else tp / (tp + fp) r = 0.0 if (tp + fn) == 0 else tp / (tp + fn) f1 = 0.0 if (p + r) == 0 else 2 * p * r / (p + r) f1s.append(f1) precs.append(p) recs.append(r) if not f1s: continue mean_f1 = float(np.mean(f1s)) mean_p = float(np.mean(precs)) mean_r = float(np.mean(recs)) if objective == "fbeta": # F_beta with beta>1 weights recall higher b2 = beta * beta mean_obj = ( 0.0 if (mean_p + mean_r) == 0 else (1 + b2) * mean_p * mean_r / (b2 * mean_p + mean_r) ) else: mean_obj = mean_f1 row = {"thr": thr, "f1": mean_f1, "precision": mean_p, "recall": mean_r, "obj": mean_obj} sweep.append(row) # Prefer higher obj; tie-break toward higher recall, then higher F1 better = ( mean_obj > best_score + 1e-9 or (abs(mean_obj - best_score) < 1e-9 and mean_r > best_row.get("recall", -1) + 1e-9) ) if better: best_score, best_thr, best_f1 = mean_obj, thr, mean_f1 best_row = row print( f" Calibrated threshold={best_thr:.4f} (val F1={best_f1:.4f} " f"P={best_row.get('precision', 0):.3f} R={best_row.get('recall', 0):.3f} " f"obj={objective}, window=[{thr_min},{thr_max}], {len(candidates)} candidates)", flush=True, ) return best_thr, best_f1, {"sweep": sweep, "selected": best_row, "objective": objective} def _save_epoch_visuals(model, processor, device, pairs: list[tuple], thr: float, out_dir: Path, epoch: int, max_pairs: int = 4): """Save Before | After | GT | Pred | Prob panels (priority #6).""" out_dir.mkdir(parents=True, exist_ok=True) for before, after, gt, pair_id in pairs[:max_pairs]: mask, score = _predict_mask(model, processor, device, before, after, thr) score = _resize_to_gt(score, gt, nearest=False) mask = _resize_to_gt(mask, gt, nearest=True) h, w = gt.shape[:2] b = np.array(Image.fromarray(before).resize((w, h))) a = np.array(Image.fromarray(after).resize((w, h))) gt_rgb = np.stack([gt, gt, gt], axis=-1) pred_rgb = np.stack([mask, mask, mask], axis=-1) # Probability heatmap (grayscale → red tint) p_u8 = (np.clip(score, 0, 1) * 255).astype(np.uint8) prob_rgb = np.stack([p_u8, (p_u8 * 0.3).astype(np.uint8), (p_u8 * 0.3).astype(np.uint8)], -1) # Labels strip panel = np.concatenate([b, a, gt_rgb, pred_rgb, prob_rgb], axis=1) Image.fromarray(panel).save(out_dir / f"ep{epoch:02d}_{pair_id}.png") def _compute_loss(logits, labels_t, loss_mode: str, ce_weight, torch, F, bce): if logits.dim() == 3: logits = logits.unsqueeze(0) target = labels_t if labels_t.dim() == 3 else labels_t.unsqueeze(0) if logits.shape[-2:] != target.shape[-2:]: target = F.interpolate( target.unsqueeze(1).float(), size=logits.shape[-2:], mode="nearest").squeeze(1) if loss_mode == "ce": return F.cross_entropy(logits, target.long(), weight=ce_weight) prob = _change_prob_from_logits(logits, torch).unsqueeze(0) if prob.shape[-2:] != target.shape[-2:]: prob = F.interpolate( prob.unsqueeze(1), size=target.shape[-2:], mode="bilinear", align_corners=False).squeeze(1) labels_b = target.float() if loss_mode == "bce": return bce(prob, labels_b) if loss_mode == "bce_dice": return 0.5 * bce(prob, labels_b) + 0.5 * _dice_loss(prob, labels_b) if loss_mode == "focal_dice": return 0.5 * _focal_loss(prob, labels_b) + 0.5 * _dice_loss(prob, labels_b) if loss_mode == "tversky": # beta=0.7 > alpha=0.3 → penalize FN (recall-oriented) return 0.6 * _tversky_loss(prob, labels_b, alpha=0.3, beta=0.7) + 0.4 * _focal_loss( prob, labels_b, alpha=0.75) if loss_mode == "tversky_dice": return 0.5 * _tversky_loss(prob, labels_b, alpha=0.3, beta=0.7) + 0.5 * _dice_loss( prob, labels_b) if loss_mode == "ce_dice": ce = F.cross_entropy(logits, target.long(), weight=ce_weight) return 0.5 * ce + 0.5 * _dice_loss(prob, labels_b) # default return 0.5 * _focal_loss(prob, labels_b) + 0.5 * _dice_loss(prob, labels_b) def train( eval_dir: Path | None, dummy: bool, epochs: int, batch_size: int, lr: float, out_root: Path, delhi_cd: Path | None = None, augment: bool = False, stride: int = 128, early_stop_patience: int = 0, loss_mode: str = "focal_dice", exclude_empty: bool = True, full_resize: bool = True, min_change_frac: float = 0.001, pos_oversample: int = 3, min_tile_change: float = 0.005, visualize: bool = True, use_scheduler: bool = True, preset_name: str | None = None, change_centered: bool = False, thr_min: float = 0.2, thr_max: float = 0.7, thr_objective: str = "f1", warm_start: str | None = None, pos_only: bool = False, keep_hard_neg: bool = True, synthetic_dir: Path | None = None, max_train_tiles: int | None = 160, synth_train_cap: int = 256, ) -> Path: torch, DataLoader, _Dataset, WeightedRandomSampler, AutoImageProcessor, AutoModel = _try_torch() import torch.nn.functional as F if delhi_cd is not None and not dummy: train_pairs, val_pairs, test_pairs, split_info = _load_pairs_from_delhi_cd(delhi_cd) else: pairs = _load_pairs(eval_dir, dummy) train_pairs, val_pairs, test_pairs = _split_pairs(pairs) split_info = { "train": [p[3] for p in train_pairs], "val": [p[3] for p in val_pairs], "test": [p[3] for p in test_pairs], "split": "70/15/15", } synth_holdout_pairs: list[tuple] = [] if synthetic_dir is not None and not dummy: synth_pairs = _load_pairs_from_synthetic(synthetic_dir) if synth_pairs: s_train, s_val, s_test = _split_pairs(synth_pairs) pool = list(s_train) + list(s_test) + list(s_val) rng = random.Random(0) rng.shuffle(pool) cap = max(0, int(synth_train_cap)) extra_train = pool[:cap] synth_holdout_pairs = pool[cap:cap + 64] train_pairs = list(train_pairs) + extra_train # Primary val stays Delhi-only. The 94% F1 last run was synthetic val. split_info["synthetic_dir"] = str(synthetic_dir) split_info["synthetic_train"] = [p[3] for p in extra_train] split_info["synthetic_holdout"] = [p[3] for p in synth_holdout_pairs] split_info["synthetic_train_cap"] = cap split_info["primary_val"] = "delhi_only" split_info["primary_score"] = "frozen_delhi_test_f1" split_info["train"] = [p[3] for p in train_pairs] split_info["val"] = [p[3] for p in val_pairs] print( f" Synthetic mix: train+={len(extra_train)} holdout={len(synth_holdout_pairs)} " f"(val remains {len(val_pairs)} Delhi pairs)", flush=True, ) if exclude_empty and not dummy: train_pairs = _filter_empty( train_pairs, min_change_frac, keep_hard_neg=keep_hard_neg) val_pairs = _filter_empty( val_pairs, min_change_frac, keep_hard_neg=False) test_change = _filter_empty( test_pairs, min_change_frac, keep_hard_neg=False) if not train_pairs: raise SystemExit("No train pairs left after excluding empty GT.") if not val_pairs: val_pairs = train_pairs[-1:] test_pairs = test_change or test_pairs split_info["excluded_empty"] = True split_info["min_change_frac"] = min_change_frac split_info["train"] = [p[3] for p in train_pairs] split_info["val"] = [p[3] for p in val_pairs] split_info["test"] = [p[3] for p in test_pairs] balance = { "train": _class_balance_report(train_pairs, "train"), "val": _class_balance_report(val_pairs, "val"), "test": _class_balance_report(test_pairs, "test"), } # keep-empty / exclude_empty=False must retain hard-negative (all-zero GT) tiles drop_empty_tiles = bool(exclude_empty) and not pos_only train_ds = DelhiTileDataset( train_pairs, train=True, stride=stride, augment=augment, seed=0, full_resize=full_resize, min_tile_change=min_tile_change, pos_oversample=pos_oversample, change_centered=change_centered, drop_empty_tiles=drop_empty_tiles, pos_only=pos_only, max_train_tiles=max_train_tiles) val_ds = DelhiTileDataset( val_pairs, train=False, stride=_TILE, augment=False, full_resize=full_resize, min_tile_change=0.0, pos_oversample=1, change_centered=False) pos_frac = float(np.mean([(g > 127).mean() for _, _, g, _ in train_pairs])) pos_frac = max(pos_frac, 1e-3) pos_weight = (1.0 - pos_frac) / pos_frac pos_weight = float(min(50.0, max(2.0, pos_weight))) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Device: {device} | split {split_info['split']} | " f"train={len(train_pairs)} val={len(val_pairs)} test={len(test_pairs)} | " f"tiles train={len(train_ds.index)} val={len(val_ds.index)} | " f"lr={lr} aug={augment} loss={loss_mode} pos_w={pos_weight:.1f} " f"pos_oversamplex{pos_oversample} scheduler={use_scheduler}", flush=True) processor = AutoImageProcessor.from_pretrained(_MODEL_ID, trust_remote_code=True) warm_path = Path(warm_start).resolve() if warm_start else None if warm_path and warm_path.is_dir(): print(f"Warm-start from {warm_path}", flush=True) model = AutoModel.from_pretrained(warm_path, trust_remote_code=True) try: processor = AutoImageProcessor.from_pretrained(warm_path, trust_remote_code=True) except Exception: pass else: if warm_start: print(f"Warm-start path missing ({warm_start}); loading hub weights", flush=True) model = AutoModel.from_pretrained(_MODEL_ID, trust_remote_code=True) model.to(device) model.eval() print("Validating logit->prob conversion...", flush=True) probe = _probe_output_scale(model, processor, device, train_pairs, n=2) model.train() weights = train_ds.sampler_weights() sampler = WeightedRandomSampler( weights=weights, num_samples=len(weights), replacement=True) train_loader = DataLoader( train_ds.torch_dataset(), batch_size=batch_size, sampler=sampler, num_workers=0) optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4) scheduler = None if use_scheduler: scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, mode="max", factor=0.5, patience=2, min_lr=1e-6) bce = torch.nn.BCELoss() ce_weight = torch.tensor([1.0, pos_weight], dtype=torch.float32, device=device) run_id = time.strftime("%Y%m%d_%H%M%S") run_dir = out_root / run_id run_dir.mkdir(parents=True, exist_ok=True) vis_dir = run_dir / "visuals" (run_dir / "split.json").write_text(json.dumps(split_info, indent=2), encoding="utf-8") (run_dir / "class_balance.json").write_text(json.dumps(balance, indent=2), encoding="utf-8") (run_dir / "output_probe.json").write_text(json.dumps(probe, indent=2), encoding="utf-8") (run_dir / "config.json").write_text(json.dumps({ "model_id": _MODEL_ID, "epochs": epochs, "lr": lr, "batch_size": batch_size, "augment": augment, "stride": stride, "early_stop_patience": early_stop_patience, "loss": loss_mode, "pos_weight": pos_weight, "pos_oversample": pos_oversample, "min_tile_change": min_tile_change, "change_centered": change_centered, "pos_only": pos_only, "exclude_empty": exclude_empty, "full_resize": full_resize, "min_change_frac": min_change_frac, "scheduler": use_scheduler, "visualize": visualize, "thr_min": thr_min, "thr_max": thr_max, "thr_objective": thr_objective, "warm_start": warm_start, "delhi_cd": str(delhi_cd) if delhi_cd else None, "preset": preset_name, "change_channel": "softmax_last (ch1)", }, indent=2), encoding="utf-8") history = [] best_f1 = -1.0 best_path = run_dir / "best" best_thr = 0.5 stale = 0 for epoch in range(1, epochs + 1): model.train() total_loss = 0.0 n_batches = 0 from PIL import Image as PILImage for batch in train_loader: before_np, after_np, gt_np = batch optimizer.zero_grad() batch_loss = 0.0 for i in range(before_np.shape[0]): b = before_np[i].numpy().astype(np.uint8) a = after_np[i].numpy().astype(np.uint8) label = gt_np[i].numpy() inputs = processor( images=(PILImage.fromarray(b), PILImage.fromarray(a)), return_tensors="pt", ) inputs = {k: v.to(device) for k, v in inputs.items()} labels_t = torch.from_numpy(label).to(device) outputs = model(**inputs) sample_loss = _compute_loss( outputs.logits, labels_t, loss_mode, ce_weight, torch, F, bce) batch_loss = batch_loss + sample_loss batch_loss = batch_loss / max(before_np.shape[0], 1) batch_loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() total_loss += float(batch_loss.item()) n_batches += 1 model.eval() thr, thr_f1, thr_detail = _calibrate_threshold( model, processor, device, val_pairs, thr_min=thr_min, thr_max=thr_max, objective=thr_objective) val_metrics = _eval_pairs(model, processor, device, val_pairs, threshold=thr) avg_loss = total_loss / max(n_batches, 1) cur_lr = float(optimizer.param_groups[0]["lr"]) row = { "epoch": epoch, "train_loss": round(avg_loss, 4), "val_mean_f1": val_metrics["mean_f1"], "val_precision": val_metrics["mean_precision"], "val_recall": val_metrics["mean_recall"], "val_iou": val_metrics["mean_iou"], "threshold": thr, "calibrate_f1": round(thr_f1, 4), "calibrate_selected": thr_detail.get("selected"), "val_mean_prob": val_metrics["mean_prob"], "val_max_prob": val_metrics["max_prob"], "lr": cur_lr, } history.append(row) print( f" epoch {epoch}/{epochs} loss={avg_loss:.4f} " f"val_F1={val_metrics['mean_f1']:.4f} P={val_metrics['mean_precision']:.3f} " f"R={val_metrics['mean_recall']:.3f} IoU={val_metrics['mean_iou']:.3f} " f"thr={thr:.6g} max_p={val_metrics['max_prob']:.4f} lr={cur_lr:.2e}", flush=True, ) (run_dir / "history.json").write_text(json.dumps(history, indent=2), encoding="utf-8") if visualize: _save_epoch_visuals( model, processor, device, val_pairs, thr, vis_dir, epoch) if scheduler is not None: scheduler.step(val_metrics["mean_f1"]) improved = val_metrics["mean_f1"] > best_f1 + 1e-6 if improved or best_f1 < 0: best_f1 = val_metrics["mean_f1"] best_thr = thr stale = 0 best_path.mkdir(parents=True, exist_ok=True) model.save_pretrained(best_path) processor.save_pretrained(best_path) (best_path / "threshold.json").write_text( json.dumps({ "threshold": best_thr, "val_f1": best_f1, "val_precision": val_metrics["mean_precision"], "val_recall": val_metrics["mean_recall"], "val_iou": val_metrics["mean_iou"], "epoch": epoch, }, indent=2), encoding="utf-8") else: stale += 1 if early_stop_patience and stale >= early_stop_patience: print(f"Early stop at epoch {epoch}. Best val F1={best_f1:.4f} thr={best_thr}", flush=True) break # Frozen threshold from best checkpoint for methodologically sound test F1 test_metrics = { "mean_f1": 0.0, "mean_precision": 0.0, "mean_recall": 0.0, "mean_iou": 0.0, "n": 0, "threshold": best_thr, } if test_pairs and best_path.is_dir(): best_model = AutoModel.from_pretrained(best_path, trust_remote_code=True) best_model.to(device) best_processor = AutoImageProcessor.from_pretrained(best_path, trust_remote_code=True) best_model.eval() thr_path = best_path / "threshold.json" if thr_path.is_file(): best_thr = float(json.loads(thr_path.read_text()).get("threshold", best_thr)) print(f"Test eval with FROZEN threshold={best_thr} from best checkpoint", flush=True) test_metrics = _eval_pairs( best_model, best_processor, device, test_pairs, threshold=best_thr) _save_epoch_visuals( best_model, best_processor, device, test_pairs, best_thr, run_dir / "visuals_test", epoch=0, max_pairs=len(test_pairs)) synth_holdout_metrics = None if synth_holdout_pairs and best_path.is_dir(): if "best_model" not in locals(): best_model = AutoModel.from_pretrained(best_path, trust_remote_code=True) best_model.to(device) best_processor = AutoImageProcessor.from_pretrained(best_path, trust_remote_code=True) best_model.eval() synth_holdout_metrics = _eval_pairs( best_model, best_processor, device, synth_holdout_pairs, threshold=best_thr) print( f"Synthetic holdout (secondary) F1={synth_holdout_metrics['mean_f1']:.4f} " f"n={synth_holdout_metrics['n']}", flush=True, ) meta = { "model_id": _MODEL_ID, "dummy": dummy, "epochs": epochs, "epochs_ran": len(history), "lr": lr, "augment": augment, "stride": stride, "loss": loss_mode, "pos_weight": pos_weight, "pos_oversample": pos_oversample, "exclude_empty": exclude_empty, "full_resize": full_resize, "threshold": best_thr, "device": str(device), "split": split_info, "class_balance": balance, "train_pairs": len(train_pairs), "val_pairs": len(val_pairs), "test_pairs": len(test_pairs), "train_tiles": len(train_ds.index), "best_val_f1": best_f1 if best_f1 >= 0 else 0.0, "test_mean_f1": test_metrics["mean_f1"], "test_precision": test_metrics.get("mean_precision", 0.0), "test_recall": test_metrics.get("mean_recall", 0.0), "test_iou": test_metrics.get("mean_iou", 0.0), "test_pixel_acc": test_metrics.get("mean_pixel_acc", 0.0), "primary_score": { "name": "frozen_delhi_test_f1", "f1": test_metrics["mean_f1"], "precision": test_metrics.get("mean_precision", 0.0), "recall": test_metrics.get("mean_recall", 0.0), "iou": test_metrics.get("mean_iou", 0.0), "pixel_acc": test_metrics.get("mean_pixel_acc", 0.0), "n": test_metrics.get("n", 0), "threshold": best_thr, }, "synthetic_holdout": synth_holdout_metrics, "history": history, "preset": preset_name, "output_probe": probe, } (run_dir / "metrics.json").write_text(json.dumps(meta, indent=2), encoding="utf-8") (run_dir / "primary_score.json").write_text( json.dumps(meta.get("primary_score"), indent=2), encoding="utf-8") print( f"PRIMARY SCORE (frozen Delhi test): F1={test_metrics['mean_f1']:.4f} " f"P={test_metrics.get('mean_precision', 0):.3f} " f"R={test_metrics.get('mean_recall', 0):.3f} " f"IoU={test_metrics.get('mean_iou', 0):.3f} " f"acc={test_metrics.get('mean_pixel_acc', 0):.3f} " f"@ thr={best_thr} n={test_metrics['n']}", flush=True, ) print(f"Run complete. Artifacts: {run_dir}", flush=True) return run_dir def finalize_run( run_dir: Path, eval_dir: Path | None, dummy: bool, history: list[dict] | None = None, ) -> Path: torch, *_rest = _try_torch() AutoImageProcessor = _rest[-2] AutoModel = _rest[-1] split_path = run_dir / "split.json" best_path = run_dir / "best" if not split_path.is_file(): raise SystemExit(f"Missing {split_path}") if not best_path.is_dir(): raise SystemExit(f"Missing checkpoint at {best_path}") split_info = json.loads(split_path.read_text(encoding="utf-8")) pairs = _load_pairs(eval_dir, dummy) by_id = {p[3]: p for p in pairs} test_pairs = [by_id[pid] for pid in split_info.get("test", []) if pid in by_id] thr = 0.5 thr_path = best_path / "threshold.json" if thr_path.is_file(): thr = float(json.loads(thr_path.read_text()).get("threshold", 0.5)) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") best_model = AutoModel.from_pretrained(best_path, trust_remote_code=True) best_model.to(device) best_processor = AutoImageProcessor.from_pretrained(best_path, trust_remote_code=True) best_model.eval() test_metrics = _eval_pairs(best_model, best_processor, device, test_pairs, threshold=thr) hist = history or [] hist_path = run_dir / "history.json" if not hist and hist_path.is_file(): hist = json.loads(hist_path.read_text(encoding="utf-8")) best_val_f1 = max((row.get("val_mean_f1", 0.0) for row in hist), default=0.0) meta = { "model_id": _MODEL_ID, "dummy": dummy, "epochs": len(hist) or None, "device": str(device), "threshold": thr, "split": split_info, "train_pairs": len(split_info.get("train", [])), "val_pairs": len(split_info.get("val", [])), "test_pairs": len(test_pairs), "best_val_f1": best_val_f1, "test_mean_f1": test_metrics["mean_f1"], "test_precision": test_metrics.get("mean_precision", 0.0), "test_recall": test_metrics.get("mean_recall", 0.0), "test_iou": test_metrics.get("mean_iou", 0.0), "history": hist, "finalized": True, } (run_dir / "metrics.json").write_text(json.dumps(meta, indent=2), encoding="utf-8") print(f"Test F1={test_metrics['mean_f1']:.4f} ({test_metrics['n']} pairs) thr={thr}") return run_dir def main(): parser = argparse.ArgumentParser(description="Fine-tune AdaptFormer on Delhi tiles") parser.add_argument("--manifest", type=str, default="docs/delhi_eval/manifest.json") parser.add_argument("--dummy", action="store_true") parser.add_argument("--preset", choices=["", "day5", "fix", "v2", "v3", "v4", "wed"], default="", help="wed = diagnosis CE+pos_weight + hard-neg retention (target test F1>0.60)") parser.add_argument("--epochs", type=int, default=None) parser.add_argument("--batch-size", type=int, default=None) parser.add_argument("--lr", type=float, default=None) parser.add_argument("--out", type=str, default="runs/finetune_adaptformer") parser.add_argument("--delhi-cd", type=str, default="") parser.add_argument("--augment", action="store_true") parser.add_argument("--stride", type=int, default=None) parser.add_argument("--early-stop", type=int, default=None) parser.add_argument( "--loss", choices=["bce", "bce_dice", "ce", "focal_dice", "ce_dice", "tversky", "tversky_dice"], default=None, ) parser.add_argument("--exclude-empty", action="store_true", default=None) parser.add_argument("--keep-empty", action="store_true", help="do not exclude empty GT (overrides preset)") parser.add_argument("--no-hard-neg", action="store_true", help="drop mined hn_* hard-negatives even if preset keeps them") parser.add_argument("--full-resize", action="store_true", default=None) parser.add_argument("--min-change-frac", type=float, default=None) parser.add_argument("--pos-oversample", type=int, default=None) parser.add_argument("--min-tile-change", type=float, default=None) parser.add_argument("--change-centered", action="store_true") parser.add_argument("--no-change-centered", action="store_true") parser.add_argument("--pos-only", action="store_true", help="train only on tiles that contain change pixels") parser.add_argument("--thr-min", type=float, default=None) parser.add_argument("--thr-max", type=float, default=None) parser.add_argument("--thr-objective", choices=["f1", "fbeta"], default=None) parser.add_argument("--warm-start", type=str, default="") parser.add_argument("--visualize", action="store_true") parser.add_argument("--no-visualize", action="store_true") parser.add_argument("--scheduler", action="store_true") parser.add_argument("--no-scheduler", action="store_true") parser.add_argument("--eval-run-dir", type=str, default="") parser.add_argument( "--synthetic-dir", type=str, default="", help="Priyanka synthetic GT folder (before/after/mask PNG triplets). " "If omitted, common local paths are auto-detected.", ) parser.add_argument( "--no-synthetic", action="store_true", help="do not mix auto-detected synthetic GT into training", ) parser.add_argument( "--max-train-tiles", type=int, default=None, help="Cap training tiles after oversample (default 160; 768 when synthetic is used).", ) parser.add_argument( "--synth-train-cap", type=int, default=256, help="Max synthetic triplets mixed into train (val stays Delhi-only).", ) args = parser.parse_args() if args.preset == "wed": preset = _WED_PRESET preset_name = "wed" elif args.preset == "v4": preset = _V4_PRESET preset_name = "v4" elif args.preset == "v3": preset = _V3_PRESET preset_name = "v3" elif args.preset == "v2": preset = _V2_PRESET preset_name = "v2" elif args.preset == "fix": preset = _FIX_PRESET preset_name = "fix" elif args.preset == "day5": preset = _DAY5_PRESET preset_name = "day5" else: preset = {} preset_name = None epochs = args.epochs if args.epochs is not None else preset.get("epochs", 12) batch_size = args.batch_size if args.batch_size is not None else preset.get("batch_size", 2) lr = args.lr if args.lr is not None else preset.get("lr", 1e-5) augment = True if args.augment or preset.get("augment") else False stride = args.stride if args.stride is not None else preset.get("stride", _TILE) early_stop = (args.early_stop if args.early_stop is not None else preset.get("early_stop_patience", 0)) loss_mode = args.loss if args.loss is not None else preset.get("loss", "focal_dice") exclude_empty = preset.get("exclude_empty", True) if args.keep_empty: exclude_empty = False elif args.exclude_empty: exclude_empty = True full_resize = preset.get("full_resize", True) if args.full_resize: full_resize = True min_change_frac = (args.min_change_frac if args.min_change_frac is not None else preset.get("min_change_frac", 0.001)) pos_oversample = (args.pos_oversample if args.pos_oversample is not None else preset.get("pos_oversample", 1)) min_tile_change = (args.min_tile_change if args.min_tile_change is not None else preset.get("min_tile_change", 0.0)) change_centered = bool(preset.get("change_centered", False)) if args.change_centered: change_centered = True if args.no_change_centered: change_centered = False pos_only = bool(preset.get("pos_only", False) or args.pos_only) thr_min = args.thr_min if args.thr_min is not None else float(preset.get("thr_min", 0.2)) thr_max = args.thr_max if args.thr_max is not None else float(preset.get("thr_max", 0.7)) thr_objective = args.thr_objective or preset.get("thr_objective", "f1") warm_start = args.warm_start or preset.get("warm_start") or None keep_hard_neg = bool(preset.get("keep_hard_neg", True)) and not args.no_hard_neg visualize = preset.get("visualize", False) if args.visualize: visualize = True if args.no_visualize: visualize = False use_scheduler = preset.get("scheduler", False) if args.scheduler: use_scheduler = True if args.no_scheduler: use_scheduler = False delhi_cd_arg = args.delhi_cd if not delhi_cd_arg and args.preset in ("day5", "fix", "v2", "v3", "v4", "wed") and not args.dummy: delhi_cd_arg = "data/delhi_cd" manifest = Path(args.manifest).resolve() if not args.dummy else None delhi_cd = Path(delhi_cd_arg).resolve() if delhi_cd_arg else None if args.eval_run_dir: finalize_run(Path(args.eval_run_dir).resolve(), manifest, args.dummy) return synthetic_dir = None if not args.no_synthetic: synthetic_dir = discover_synthetic_dir(args.synthetic_dir) if synthetic_dir: print(f"Using synthetic GT dataset: {synthetic_dir}", flush=True) elif args.synthetic_dir: raise SystemExit(f"Synthetic GT folder not found: {args.synthetic_dir}") else: print( "No synthetic GT folder on this machine " f"(looked for {_SYNTHETIC_DIR_CANDIDATES[0]}). Training on Delhi labeled GT only.", flush=True, ) max_train_tiles = args.max_train_tiles if max_train_tiles is None: max_train_tiles = 768 if synthetic_dir else 160 if synthetic_dir and args.thr_max is None: thr_max = max(thr_max, 0.85) print(f" Delhi-primary threshold window [{thr_min}, {thr_max}]", flush=True) train( eval_dir=manifest, dummy=args.dummy, epochs=epochs, batch_size=batch_size, lr=lr, out_root=Path(args.out).resolve(), delhi_cd=delhi_cd, synthetic_dir=synthetic_dir, augment=augment, stride=stride, early_stop_patience=early_stop, loss_mode=loss_mode, exclude_empty=exclude_empty, full_resize=full_resize, min_change_frac=min_change_frac, pos_oversample=pos_oversample, min_tile_change=min_tile_change, visualize=visualize, use_scheduler=use_scheduler, preset_name=preset_name, change_centered=change_centered, thr_min=thr_min, thr_max=thr_max, thr_objective=thr_objective, warm_start=warm_start, pos_only=pos_only, keep_hard_neg=keep_hard_neg, max_train_tiles=max_train_tiles, synth_train_cap=args.synth_train_cap, ) if __name__ == "__main__": main()