#!/usr/bin/env python3 """Fine-tune the landing-value ViT on HJ-labeled tiles. .venv/bin/python scripts/train.py --data data/tiles --epochs 8 """ from __future__ import annotations import argparse import sys import time from pathlib import Path import numpy as np import torch import torch.nn.functional as F from torch.utils.data import DataLoader, Subset sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from reachdown.paths import data_path from reachdown.data import TileDataset from reachdown.model import LandingValueViT def soft_dice_loss(logits: torch.Tensor, target: torch.Tensor, eps: float = 1.0) -> torch.Tensor: p = torch.sigmoid(logits) num = 2.0 * (p * target).sum(dim=(1, 2, 3)) + eps den = (p + target).sum(dim=(1, 2, 3)) + eps return 1.0 - (num / den).mean() def auroc(scores: np.ndarray, labels: np.ndarray) -> float: """Rank-based AUROC (labels thresholded at 0.5); subsampled for speed.""" idx = np.random.default_rng(0).choice(scores.size, min(scores.size, 200_000), replace=False) s, y = scores.ravel()[idx], labels.ravel()[idx] > 0.5 n_pos, n_neg = int(y.sum()), int((~y).sum()) if not n_pos or not n_neg: return float("nan") ranks = s.argsort().argsort().astype(np.float64) + 1 return float((ranks[y].sum() - n_pos * (n_pos + 1) / 2) / (n_pos * n_neg)) @torch.no_grad() def evaluate(model, loader, device) -> tuple[float, float]: model.eval() scores, labels, maes = [], [], [] for x, y in loader: x, y = x.to(device), y.to(device) p = torch.sigmoid(model(x)) scores.append(p.cpu().numpy()) labels.append(y.cpu().numpy()) maes.append(float((p - y).abs().mean())) return auroc(np.concatenate(scores), np.concatenate(labels)), float(np.mean(maes)) def main() -> None: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("--data", type=Path, default=data_path("data", "tiles")) ap.add_argument("--label-key", default="label_value", choices=("label_value", "label_run")) ap.add_argument("--safety-buffer", type=float, default=0.0, help="re-compose conservative labels from stored margin_hj (m)") ap.add_argument("--input-size", type=int, default=364, help="multiple of 14") ap.add_argument("--epochs", type=int, default=8) ap.add_argument("--batch-size", type=int, default=8) ap.add_argument("--lr", type=float, default=3e-4, help="decoder LR; backbone gets 0.1x") ap.add_argument("--false-safe-weight", type=float, default=1.0, help=">1 penalizes predicting-safe-where-unsafe (conservative bias)") ap.add_argument("--val-frac", type=float, default=0.2) ap.add_argument("--seed", type=int, default=0, help="torch/numpy seed for multi-seed runs") ap.add_argument("--from-scratch", action="store_true", help="FW-3: random-init backbone (no depth pretraining)") ap.add_argument("--unfreeze-blocks", type=int, default=4, help="FW-4: number of last encoder blocks to unfreeze") ap.add_argument("--zero-channels", type=int, nargs="*", default=[], help="FW-1: input channels to zero (e.g. 4 = analytic margin)") ap.add_argument("--out", type=Path, default=Path("runs/vit")) args = ap.parse_args() torch.manual_seed(args.seed) np.random.seed(args.seed) device = "cuda" if torch.cuda.is_available() else "cpu" dataset = TileDataset(args.data, input_size=args.input_size, label_key=args.label_key, safety_buffer_m=args.safety_buffer, zero_channels=tuple(args.zero_channels)) # split by terrain, not tile: tiles differing only in aircraft state must # not straddle the split terrains = [dataset.terrain_seed(i) for i in range(len(dataset))] unique = sorted(set(terrains)) if len(unique) < 2: raise SystemExit("need >= 2 distinct terrains for a terrain-split validation") val_terrains = set(unique[max(1, int(len(unique) * (1 - args.val_frac))):]) val_idx = [i for i, t in enumerate(terrains) if t in val_terrains] train_idx = [i for i, t in enumerate(terrains) if t not in val_terrains] train_set, val_set = Subset(dataset, train_idx), Subset(dataset, val_idx) train_loader = DataLoader( train_set, batch_size=args.batch_size, shuffle=True, num_workers=4, persistent_workers=True, pin_memory=True, ) val_loader = DataLoader( val_set, batch_size=args.batch_size, num_workers=2, persistent_workers=True, pin_memory=True, ) model = LandingValueViT(unfreeze_last_blocks=args.unfreeze_blocks, pretrained=not args.from_scratch).to(device) backbone_params, decoder_params = model.trainable_parameters() opt = torch.optim.AdamW( [{"params": decoder_params, "lr": args.lr}, {"params": backbone_params, "lr": 0.1 * args.lr}], weight_decay=1e-4, ) sched = torch.optim.lr_scheduler.CosineAnnealingLR( opt, T_max=args.epochs * len(train_loader) ) args.out.mkdir(parents=True, exist_ok=True) # persist the exact run configuration: reproducibility surface for the paper import json (args.out / "args.json").write_text(json.dumps( {k: str(v) if isinstance(v, Path) else v for k, v in vars(args).items()}, indent=2, sort_keys=True)) history = args.out / "history.csv" history.write_text("epoch,train_loss,val_auroc,val_mae\n") best = -float("inf") print(f"{len(train_set)} train / {len(val_set)} val tiles on {device}; " f"trainable params: {sum(p.numel() for p in decoder_params + backbone_params) / 1e6:.1f}M") for epoch in range(args.epochs): model.train() t0, losses = time.time(), [] for x, y in train_loader: x, y = x.to(device), y.to(device) with torch.autocast(device, dtype=torch.bfloat16, enabled=device == "cuda"): logits = model(x) # asymmetric BCE: a false-safe (predict safe where label unsafe) # costs more than a false-unsafe, biasing the map conservative w = 1.0 + (args.false_safe_weight - 1.0) * (y < 0.5).float() bce = F.binary_cross_entropy_with_logits(logits, y, weight=w) loss = bce + 0.5 * soft_dice_loss(logits, y) opt.zero_grad(set_to_none=True) loss.backward() opt.step() sched.step() losses.append(loss.item()) val_auroc, val_mae = evaluate(model, val_loader, device) print(f"epoch {epoch + 1:02d}/{args.epochs} loss {np.mean(losses):.4f} " f"val AUROC {val_auroc:.4f} val MAE {val_mae:.4f} ({time.time() - t0:.0f}s)") with history.open("a") as f: f.write(f"{epoch + 1},{np.mean(losses):.6f},{val_auroc:.6f},{val_mae:.6f}\n") # single-class validation yields NaN AUROC; fall back to MAE so a # checkpoint is always produced score = -val_mae if np.isnan(val_auroc) else val_auroc if score > best: best = score torch.save( {"state_dict": model.state_dict(), "input_size": args.input_size, "label_key": args.label_key, "val_auroc": val_auroc, "zero_channels": list(args.zero_channels), "unfreeze_blocks": args.unfreeze_blocks, "pretrained": not args.from_scratch}, args.out / "best.pt", ) print(f"best val score {best:.4f} -> {args.out / 'best.pt'}") if __name__ == "__main__": main()