ckwolfe's picture
Publish complete fixed-wing ViT runs, dataset provenance, and evaluation artifacts
728cafc verified
Raw History Blame Contribute Delete
7.67 kB
#!/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()