VLAlert / training /Nexar /train_resnet18_bigru.py
AsianPlayer's picture
Add VLAlert code
1e05592 verified
Raw
History Blame Contribute Delete
10.7 kB
#!/usr/bin/env python3
"""ResNet18-BiGRU Track-A trainer.
Protocol (CORRECTED 2026-04-28):
- TRAIN = all 1500 nexar train clips (data/rgb_clip_cache/train.pt)
- VAL = test-public subset (~672 clips) of data/rgb_clip_cache/test.pt,
filtered by NEXAR_COLLISION/solution.csv `Usage == "Public"`.
- TEST = test-private subset (~672 clips, `Usage == "Private"`).
This replaces the earlier 1280/220 internal-val split.
Cache (`tools/extract_rgb_clip_cache.py`) contains T=64 [N, T, 512] feats
(frozen-ResNet18 ImageNet pooled features).
Model: BiGRU(d=256) → attentive pool → MLP → 1 logit.
Usage:
python -m training.Nexar.train_resnet18_bigru --seed 0 \\
--train_cache data/rgb_clip_cache/train.pt \\
--test_cache data/rgb_clip_cache/test.pt \\
--solution_csv NEXAR_COLLISION/solution.csv \\
--label_csv nexar-collision-prediction/train.csv \\
--output checkpoints/Nexar/resnet18_bigru_seed0
"""
from __future__ import annotations
import argparse
import csv
import json
import random
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
from sklearn.metrics import average_precision_score, roc_auc_score
from torch.utils.data import DataLoader, Dataset
def set_seed(s: int):
random.seed(s); np.random.seed(s); torch.manual_seed(s)
torch.cuda.manual_seed_all(s)
class FeatDataset(Dataset):
def __init__(self, cache_path: Path, labels: dict[str, int],
badas_soft: dict | None = None,
keep_ids: set[str] | None = None):
c = torch.load(cache_path, weights_only=False, map_location="cpu")
self.ids: list[str] = c["ids"]
self.feat = c["feat"]
self.labels = labels
self.badas_soft = badas_soft or {}
kept = []
for i, vid in enumerate(self.ids):
if vid not in labels:
continue
if keep_ids is not None and vid not in keep_ids:
continue
kept.append(i)
self.kept = kept
def __len__(self):
return len(self.kept)
def __getitem__(self, i):
idx = self.kept[i]
vid = self.ids[idx]
x = self.feat[idx].float()
y = float(self.labels[vid])
soft = float(self.badas_soft.get(vid, y))
return x, torch.tensor(y), torch.tensor(soft), vid
class BiGRUHead(nn.Module):
def __init__(self, in_dim: int = 512, hidden: int = 256,
dropout: float = 0.3):
super().__init__()
self.gru = nn.GRU(in_dim, hidden, num_layers=1, bidirectional=True,
batch_first=True)
self.attn = nn.Sequential(nn.Linear(2 * hidden, 128),
nn.Tanh(), nn.Linear(128, 1))
self.cls = nn.Sequential(nn.Dropout(dropout),
nn.Linear(2 * hidden, 1))
def forward(self, x: torch.Tensor) -> torch.Tensor:
h, _ = self.gru(x)
a = torch.softmax(self.attn(h).squeeze(-1), dim=1)
pooled = (h * a.unsqueeze(-1)).sum(dim=1)
return self.cls(pooled).squeeze(-1)
def load_labels(csv_path: Path) -> dict[str, int]:
return {r["id"]: int(r["target"] or 0)
for r in csv.DictReader(open(csv_path))
if r.get("target") is not None}
def load_solution_split(csv_path: Path):
rows = list(csv.DictReader(open(csv_path)))
pub = {r["id"] for r in rows if r["Usage"] == "Public"}
priv = {r["id"] for r in rows if r["Usage"] == "Private"}
targets = {r["id"]: int(r["target"]) for r in rows}
return pub, priv, targets
def load_badas_soft(p: Path) -> dict[str, float]:
if not p.exists():
return {}
j = json.loads(p.read_text())
return {cid: float(rec["score_last4s"])
for cid, rec in j.items()
if rec.get("score_last4s") is not None
and not np.isnan(rec["score_last4s"])}
@torch.no_grad()
def eval_split(model, loader, device) -> tuple[float, float]:
model.eval()
ys, ps = [], []
for x, y, _soft, _vid in loader:
logit = model(x.to(device)).cpu().numpy()
ys.extend(y.numpy().tolist()); ps.extend(logit.tolist())
ys, ps = np.asarray(ys), np.asarray(ps)
if len(np.unique(ys)) < 2:
return float("nan"), float("nan")
return (float(average_precision_score(ys, ps)),
float(roc_auc_score(ys, ps)))
@torch.no_grad()
def eval_kaggle_mAP(model, ds: FeatDataset, device, solution: Path,
batch: int = 64) -> tuple[float, float]:
rows = list(csv.DictReader(open(solution)))
group = {r["id"]: int(r["group"]) for r in rows}
usage = {r["id"]: r["Usage"] for r in rows}
model.eval()
score: dict[str, float] = {}
target: dict[str, int] = {}
dl = DataLoader(ds, batch_size=batch, shuffle=False, num_workers=4)
for x, y, _soft, vids in dl:
logit = model(x.to(device)).cpu().numpy()
for v, p, t in zip(vids, logit.tolist(), y.numpy().tolist()):
score[v] = float(p); target[v] = int(t)
common = sorted(set(score) & set(group))
if not common:
return float("nan"), float("nan")
pub_aps, priv_aps = [], []
for g in (0, 1, 2):
for u, sink in (("Public", pub_aps), ("Private", priv_aps)):
ids = [v for v in common if usage.get(v) == u and group.get(v) == g]
if len(ids) < 2: continue
y = np.array([target[v] for v in ids])
s = np.array([score[v] for v in ids])
if len(np.unique(y)) < 2: continue
sink.append(float(average_precision_score(y, s)))
pub_mAP = float(np.mean(pub_aps)) if pub_aps else float("nan")
priv_mAP = float(np.mean(priv_aps)) if priv_aps else float("nan")
return pub_mAP, priv_mAP
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--train_cache", default="data/rgb_clip_cache/train.pt")
ap.add_argument("--test_cache", default="data/rgb_clip_cache/test.pt")
ap.add_argument("--solution_csv", default="NEXAR_COLLISION/solution.csv")
ap.add_argument("--label_csv", default="nexar-collision-prediction/train.csv")
ap.add_argument("--output", default="checkpoints/Nexar/resnet18_bigru_seed0")
ap.add_argument("--epochs", type=int, default=25)
ap.add_argument("--batch", type=int, default=64)
ap.add_argument("--lr", type=float, default=3e-4)
ap.add_argument("--wd", type=float, default=1e-4)
ap.add_argument("--hidden", type=int, default=256)
ap.add_argument("--dropout", type=float, default=0.3)
ap.add_argument("--badas_soft_label", default=None)
ap.add_argument("--alpha_soft", type=float, default=0.3)
ap.add_argument("--report_private", action="store_true")
args = ap.parse_args()
set_seed(args.seed)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
out = Path(args.output); out.mkdir(parents=True, exist_ok=True)
badas_soft = (load_badas_soft(Path(args.badas_soft_label))
if args.badas_soft_label else {})
if badas_soft:
print(f"[init] loaded {len(badas_soft)} BADAS soft labels")
train_labels = load_labels(Path(args.label_csv))
print(f"[init] train labels: {len(train_labels)}")
pub_ids, priv_ids, test_targets = load_solution_split(Path(args.solution_csv))
print(f"[init] solution split: public={len(pub_ids)} private={len(priv_ids)}")
train_ds = FeatDataset(Path(args.train_cache), train_labels, badas_soft)
val_ds = FeatDataset(Path(args.test_cache), test_targets, keep_ids=pub_ids)
priv_ds = FeatDataset(Path(args.test_cache), test_targets, keep_ids=priv_ids)
print(f"[init] datasets: train_n={len(train_ds)} val_pub_n={len(val_ds)} "
f"priv_n={len(priv_ds)}")
train_dl = DataLoader(train_ds, batch_size=args.batch, shuffle=True,
num_workers=4, pin_memory=True)
val_dl = DataLoader(val_ds, batch_size=args.batch, shuffle=False,
num_workers=4, pin_memory=True)
model = BiGRUHead(in_dim=512, hidden=args.hidden,
dropout=args.dropout).to(device)
opt = torch.optim.AdamW(model.parameters(), lr=args.lr,
weight_decay=args.wd)
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=args.epochs)
bce = nn.BCEWithLogitsLoss()
best_pub_mAP = -1.0; best_path = out / "best.pt"
history = []
for ep in range(args.epochs):
model.train(); running = 0.0; n_batch = 0
for x, y, soft, _vid in train_dl:
x = x.to(device); y = y.to(device); soft = soft.to(device)
logit = model(x)
loss_gt = bce(logit, y)
loss = ((1 - args.alpha_soft) * loss_gt +
args.alpha_soft * bce(logit, soft)) if badas_soft else loss_gt
opt.zero_grad(); loss.backward(); opt.step()
running += float(loss.item()); n_batch += 1
sched.step()
pub_ap, pub_auc = eval_split(model, val_dl, device)
pub_mAP, _ = eval_kaggle_mAP(model, val_ds, device,
Path(args.solution_csv),
batch=args.batch)
priv_str = ""
if args.report_private:
priv_dl = DataLoader(priv_ds, batch_size=args.batch, shuffle=False,
num_workers=4, pin_memory=True)
priv_ap, _ = eval_split(model, priv_dl, device)
_, priv_mAP = eval_kaggle_mAP(model, priv_ds, device,
Path(args.solution_csv),
batch=args.batch)
priv_str = f" [info] priv_mAP={priv_mAP:.4f}"
avg = running / max(n_batch, 1)
print(f"epoch {ep+1:02d}/{args.epochs} loss={avg:.4f} "
f"pub_AP={pub_ap:.4f} pub_AUC={pub_auc:.4f} "
f"pub_mAP={pub_mAP:.4f}{priv_str}", flush=True)
history.append({"epoch": ep + 1, "loss": avg,
"pub_AP": pub_ap, "pub_AUC": pub_auc,
"pub_mAP": pub_mAP})
if pub_mAP > best_pub_mAP:
best_pub_mAP = pub_mAP
torch.save({"head_state": model.state_dict(),
"args": vars(args), "epoch": ep + 1,
"pub_AP": pub_ap, "pub_AUC": pub_auc,
"pub_mAP": pub_mAP}, best_path)
(out / "history.json").write_text(json.dumps(history, indent=2))
print(f"[done] best pub_mAP={best_pub_mAP:.4f} ckpt={best_path}")
if __name__ == "__main__":
main()