Spaces:
Running
Running
File size: 15,937 Bytes
d70361b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 | """
Day 4 (Uday): materialize 70/15/15 train/val/test splits under data/delhi_cd/.
Writes lightweight manifests that point at Priyanka's Delhi pairs (no GeoTIFF
copies). Same seed / fractions as scripts/finetune_adaptformer.py so smoke
training and the on-disk split stay aligned.
Usage:
python scripts/build_delhi_cd_splits.py
python scripts/build_delhi_cd_splits.py --manifest docs/delhi_eval/manifest.json \\
--out data/delhi_cd --seed 0
"""
from __future__ import annotations
import argparse
import json
import random
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT))
from app.evaluation.delhi_eval import DelhiEvalNotReady, load_manifest # noqa: E402
def _training_use_ok(pair_id: str) -> bool:
"""False only when the pack's meta.json explicitly opts out (``training_use: false``).
Evidence 2026-07-31: the 9 un-orthorectified drone pairs (NCC 0.07-0.61)
measurably degraded a held-out fine-tune (frozen test F1 0.589 without
them vs 0.487 with) — see docs/ACCURACY_IMPROVE_STATUS.md. Their meta.json
now records ``training_use: false`` so a plain rerun of this script can't
silently re-include them; pass ``--include-low-quality`` to override.
"""
meta_path = ROOT / "docs" / "delhi_eval" / "dda_labeling" / pair_id / "meta.json"
if not meta_path.is_file():
return True
try:
meta = json.loads(meta_path.read_text(encoding="utf-8"))
except Exception:
return True
return meta.get("training_use", True) is not False
def _labeled_pairs(manifest: Path, min_change_frac: float = 0.0,
include_low_quality: bool = False) -> list[dict]:
import numpy as np
from PIL import Image
data = load_manifest(manifest, required=True)
pairs = []
dropped = []
excluded_low_quality = []
for pair in data.get("pairs", []):
pair_id = pair.get("pair_id") or pair.get("id")
before = pair.get("before_path") or pair.get("before")
after = pair.get("after_path") or pair.get("after")
gt = pair.get("gt_mask")
if not (pair_id and before and after):
continue
if not include_low_quality and not _training_use_ok(pair_id):
excluded_low_quality.append(pair_id)
continue
if not gt:
auto = ROOT / "docs" / "delhi_eval" / "labels" / f"{pair_id}.png"
gt = str(auto.relative_to(ROOT)).replace("\\", "/") if auto.is_file() else None
if not gt:
continue
if not (ROOT / before).is_file() or not (ROOT / after).is_file():
continue
if not (ROOT / gt).is_file():
continue
arr = np.array(Image.open(ROOT / gt).convert("L"))
frac = float((arr > 127).mean())
if min_change_frac > 0 and frac < min_change_frac:
dropped.append(pair_id)
continue
pairs.append({
"pair_id": pair_id,
"before_path": before.replace("\\", "/"),
"after_path": after.replace("\\", "/"),
"gt_mask": gt.replace("\\", "/"),
"change_types": pair.get("change_types") or [],
"notes": pair.get("notes") or "",
"change_frac": frac,
})
if excluded_low_quality:
print(f"Auto-excluded {len(excluded_low_quality)} training_use=false pair(s): "
f"{excluded_low_quality}")
if dropped:
print(f"Dropped {len(dropped)} empty/near-empty GT pairs: {dropped}")
return pairs
def split_pairs(pairs: list[dict], seed: int = 0,
train_frac: float = 0.70, val_frac: float = 0.15,
stratify: bool = False, frozen_test_ids: list[str] | None = None):
"""70/15/15 split. With stratify=True, balance change density across splits.
``frozen_test_ids``, when given, pins those pair_ids to the test split every
time regardless of what else is added to/removed from the labeled pool, so a
frozen model/threshold's test F1 stays comparable across days. Any frozen id
not currently in ``pairs`` (label missing this run) is skipped with a
warning rather than silently producing a smaller/different test set.
"""
n = len(pairs)
if n < 3:
raise SystemExit(f"Need at least 3 labeled pairs for 70/15/15; got {n}")
if frozen_test_ids:
by_id = {p["pair_id"]: p for p in pairs}
forced_test = [by_id[pid] for pid in frozen_test_ids if pid in by_id]
missing = [pid for pid in frozen_test_ids if pid not in by_id]
if missing:
print(f"WARNING: frozen test id(s) not in labeled pool this run: {missing}")
remaining = [p for p in pairs if p["pair_id"] not in set(frozen_test_ids)]
# Two-way train/val split of the remaining pool (no test carve-out here —
# the frozen ids above are the entire test set, always, every run).
rng = random.Random(seed)
if stratify:
ordered = sorted(remaining, key=lambda p: p.get("change_frac", 0.0))
else:
ordered = list(remaining)
rng.shuffle(ordered)
n_val = max(1, int(round(len(remaining) * val_frac / (train_frac + val_frac)))) if remaining else 0
if stratify:
# Take every Nth pair across the density-sorted order for val, so val
# still spans easy/medium/hard rather than clustering at one end.
step = max(1, len(ordered) // max(1, n_val))
val_idx = set(range(0, len(ordered), step)[:n_val])
val = [p for i, p in enumerate(ordered) if i in val_idx]
train = [p for i, p in enumerate(ordered) if i not in val_idx]
else:
val = ordered[:n_val]
train = ordered[n_val:]
return train, val, forced_test
rng = random.Random(seed)
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)
if not stratify:
idx = list(range(n))
rng.shuffle(idx)
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
# Stratified by change_frac tertiles so val/test are not all "easy dense" scenes
ordered = sorted(pairs, key=lambda p: p.get("change_frac", 0.0))
buckets: list[list[dict]] = [[], [], []]
for i, p in enumerate(ordered):
buckets[min(2, (i * 3) // max(n, 1))].append(p)
train, val, test = [], [], []
n_train_target = n - n_test - n_val
for bucket in buckets:
rng.shuffle(bucket)
nb = len(bucket)
# Proportional take from each density band
bt = max(0, int(round(nb * n_test / n)))
bv = max(0, int(round(nb * n_val / n)))
# Ensure at least one val/test from a band when the band is large enough
if nb >= 3:
bt = max(1, bt)
bv = max(1, bv)
if bt + bv > nb:
bt = min(bt, max(0, nb - 1))
bv = min(bv, max(0, nb - bt))
test.extend(bucket[:bt])
val.extend(bucket[bt:bt + bv])
train.extend(bucket[bt + bv:])
# Repair global counts without fully reshuffling (preserve density mix)
def _steal(src: list, dst: list, k: int):
for _ in range(k):
if not src:
break
dst.append(src.pop())
_steal(train, test, n_test - len(test))
_steal(test, train, len(test) - n_test)
_steal(train, val, n_val - len(val))
_steal(val, train, len(val) - n_val)
# Final size clamp
while len(test) > n_test and test:
train.append(test.pop())
while len(val) > n_val and val:
train.append(val.pop())
while len(train) > n_train_target and train:
if len(val) < n_val:
val.append(train.pop())
elif len(test) < n_test:
test.append(train.pop())
else:
break
print(
"Stratified by change density: "
f"train_mean%={100 * sum(p['change_frac'] for p in train) / max(len(train), 1):.2f} "
f"val_mean%={100 * sum(p['change_frac'] for p in val) / max(len(val), 1):.2f} "
f"test_mean%={100 * sum(p['change_frac'] for p in test) / max(len(test), 1):.2f}"
)
return train, val, test
def _existing_hard_negatives(out_dir: Path, name: str) -> list[dict]:
"""Hard-negative entries already in ``<out_dir>/<name>/manifest.json``, if any.
``mine_hard_negatives.py`` appends tiles directly to a split's
``train/manifest.json`` — outside the ``docs/delhi_eval/manifest.json``
source this script regenerates from. Without this, every rerun of
``build_delhi_cd_splits.py`` silently wipes them (bug found 2026-07-31:
cost wed_retrain's 5 mined tiles from both the official and an ablation
split, understating the ablation's F1 vs wed_retrain until restored).
"""
path = out_dir / name / "manifest.json"
if not path.is_file():
return []
try:
data = json.loads(path.read_text(encoding="utf-8"))
except Exception:
return []
return [p for p in data.get("pairs", [])
if "hard_negative" in (p.get("change_types") or [])]
def _write_split_dir(out_dir: Path, name: str, pairs: list[dict],
preserve_hard_negatives: bool = False) -> Path:
split_dir = out_dir / name
split_dir.mkdir(parents=True, exist_ok=True)
if preserve_hard_negatives:
existing_ids = {p["pair_id"] for p in pairs}
carried = [hn for hn in _existing_hard_negatives(out_dir, name)
if hn.get("pair_id") not in existing_ids]
if carried:
print(f"Preserved {len(carried)} existing hard-negative tile(s) in {name}: "
f"{[p['pair_id'] for p in carried]}")
pairs = pairs + carried
manifest = {
"version": 1,
"split": name,
"n_pairs": len(pairs),
"pairs": pairs,
}
path = split_dir / "manifest.json"
path.write_text(json.dumps(manifest, indent=2), encoding="utf-8")
ids_path = split_dir / "pair_ids.txt"
ids_path.write_text("\n".join(p["pair_id"] for p in pairs) + "\n", encoding="utf-8")
return path
def main():
parser = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--manifest", default="docs/delhi_eval/manifest.json")
parser.add_argument("--out", default="data/delhi_cd")
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--min-change-frac", type=float, default=0.001,
help="drop GT masks with change fraction below this (empty labels)")
parser.add_argument("--stratify", action="store_true",
help="balance change density across train/val/test (fixes Val>>Test gap)")
parser.add_argument(
"--freeze-test-ids", default="data/delhi_cd/frozen_test_ids.json",
help=(
"JSON file listing pair_ids to always assign to test, so a frozen "
"model/threshold's test F1 stays comparable across days even as "
"labeled pairs are added/removed. Pass '' to disable and let the "
"splitter pick test freely (NOT recommended once a model/threshold "
"has been calibrated against a specific test set)."
),
)
parser.add_argument(
"--exclude-prefix", action="append", default=[],
help=(
"Drop any pair whose pair_id starts with this prefix (repeatable). "
"Use for ablation splits, e.g. --exclude-prefix dda_before --exclude-prefix "
"dda_after --exclude-prefix dda_1_2 to exclude the un-orthorectified drone "
"pairs and isolate their effect on F1. Combine with a distinct --out so this "
"never overwrites the official split."
),
)
parser.add_argument(
"--include-low-quality", action="store_true",
help=(
"Include pairs whose meta.json sets training_use=false (currently the 9 "
"un-orthorectified drone pairs, confirmed to hurt F1 — see "
"docs/ACCURACY_IMPROVE_STATUS.md). Off by default; only pass this "
"deliberately, e.g. to re-test after orthomosaics are available."
),
)
args = parser.parse_args()
manifest = (ROOT / args.manifest).resolve() if not Path(args.manifest).is_absolute() else Path(args.manifest)
try:
pairs = _labeled_pairs(manifest, min_change_frac=args.min_change_frac,
include_low_quality=args.include_low_quality)
except DelhiEvalNotReady as exc:
raise SystemExit(str(exc)) from exc
if args.exclude_prefix:
before_n = len(pairs)
pairs = [p for p in pairs if not any(
(p.get("pair_id") or "").startswith(pfx) for pfx in args.exclude_prefix)]
print(f"Excluded {before_n - len(pairs)} pair(s) matching prefixes {args.exclude_prefix}")
if not pairs:
raise SystemExit("No labeled Delhi pairs found — cannot build splits.")
frozen_test_ids = None
if args.freeze_test_ids:
freeze_path = ROOT / args.freeze_test_ids
if freeze_path.is_file():
frozen_test_ids = json.loads(freeze_path.read_text(encoding="utf-8"))["test_pair_ids"]
print(f"Freezing test split to {len(frozen_test_ids)} pair(s) from {freeze_path.name}")
else:
print(f"NOTE: {freeze_path} not found — test split will NOT be frozen this run.")
train, val, test = split_pairs(
pairs, seed=args.seed, stratify=args.stratify, frozen_test_ids=frozen_test_ids)
out_dir = ROOT / args.out
out_dir.mkdir(parents=True, exist_ok=True)
_write_split_dir(out_dir, "train", train, preserve_hard_negatives=True)
_write_split_dir(out_dir, "val", val)
_write_split_dir(out_dir, "test", test)
def _mean_frac(ps):
return round(sum(p.get("change_frac", 0) for p in ps) / max(len(ps), 1), 6)
summary = {
"version": 1,
"source_manifest": str(manifest.relative_to(ROOT)).replace("\\", "/"),
"seed": args.seed,
"split": "70/15/15",
"stratified": bool(args.stratify),
"min_change_frac": args.min_change_frac,
"n_labeled": len(pairs),
"n_train": len(train),
"n_val": len(val),
"n_test": len(test),
"change_frac_mean": {
"train": _mean_frac(train),
"val": _mean_frac(val),
"test": _mean_frac(test),
},
"train": [p["pair_id"] for p in train],
"val": [p["pair_id"] for p in val],
"test": [p["pair_id"] for p in test],
}
(out_dir / "split.json").write_text(json.dumps(summary, indent=2), encoding="utf-8")
(out_dir / "README.md").write_text(
"# Delhi CD splits (Day 4)\n\n"
"70/15/15 train/val/test over labeled pairs from `docs/delhi_eval/`.\n"
"Manifests store repo-relative paths (no GeoTIFF copies).\n\n"
f"- seed={args.seed}\n"
f"- train={len(train)} val={len(val)} test={len(test)}\n"
"- Built by `scripts/build_delhi_cd_splits.py`\n"
"- Consumed by `scripts/finetune_adaptformer.py --delhi-cd data/delhi_cd`\n",
encoding="utf-8",
)
print(f"Labeled pairs: {len(pairs)}")
print(f"train={len(train)} val={len(val)} test={len(test)} (seed={args.seed})")
print(f"Wrote {out_dir / 'split.json'}")
for name in ("train", "val", "test"):
print(f" {out_dir / name / 'manifest.json'}")
if __name__ == "__main__":
main()
|