task1 / src /pmdm /train.py
siddhant20's picture
Add files using upload-large-folder tool
d667566 verified
Raw History Blame Contribute Delete
6.19 kB
"""Stage-1 training loop."""
from __future__ import annotations
import json
import time
from pathlib import Path
import numpy as np
import torch
from torch.utils.data import DataLoader
from .config import BATCH, CKPT, EPOCHS, LR, PREP, SYNTH, WEIGHT_DECAY
from .dataset import GroupedSampler, PairTileDataset, load_gt
from .folds import split as fold_split
from .losses import total_loss
from .metric import sweep_threshold
from .model import SiamCenterNet
def build_gt(items: list[tuple[str, int]], synth_gt: dict | None = None) -> dict:
"""Ground truth keyed by (split, idx) for both real and synthetic pairs."""
real = load_gt()
gt: dict = {}
for split, idx in items:
if split == "synth":
gt[(split, idx)] = synth_gt[idx] if synth_gt else np.zeros((0, 4), np.float32)
else:
gt[(split, idx)] = real.get(idx, np.zeros((0, 4), np.float32))
return gt
def load_synth_gt(root: Path | None = None) -> dict:
root = root or SYNTH
path = root / "synth" / "boxes.json"
if not path.exists():
return {}
raw = json.loads(path.read_text())
return {int(k): np.asarray(v, np.float32) for k, v in raw.items()}
def evaluate(model, val_idx: list[int], device: str, gt: dict, **kwargs) -> dict:
from .infer import predict_pair
model.eval()
candidates, gt_eval = {}, {}
for idx in val_idx:
boxes, scores = predict_pair(model, "train", idx, device=device, **kwargs)
key = f"train_{idx:03d}"
candidates[key] = (boxes, scores)
gt_eval[key] = gt[("train", idx)]
thr, best = sweep_threshold(candidates, gt_eval)
best["threshold"] = thr
model.train()
return best
def train_fold(fold: int, epochs: int = EPOCHS, backbone: str = "convnext_tiny",
batch: int = BATCH, lr: float = LR, use_synth: bool = True,
n_synth: int = 0, samples_per_pair: int = 8, device: str = "cuda",
out_dir: Path | None = None, num_workers: int = 8,
eval_every: int = 5, resume: bool = True, on_epoch_end=None) -> dict:
out_dir = Path(out_dir or CKPT) / f"stage1_fold{fold}_{backbone}"
out_dir.mkdir(parents=True, exist_ok=True)
train_idx, val_idx = fold_split(fold)
items = [("train", i) for i in train_idx]
synth_gt = load_synth_gt() if use_synth else {}
if use_synth and synth_gt:
avail = sorted(synth_gt)
chosen = avail if n_synth <= 0 else avail[:n_synth]
items += [("synth", i) for i in chosen]
gt = build_gt(items + [("train", i) for i in val_idx], synth_gt)
ds = PairTileDataset(items, gt, samples_per_pair=samples_per_pair, prep_root=PREP)
sampler = GroupedSampler(len(items), samples_per_pair)
dl = DataLoader(ds, batch_size=batch, sampler=sampler, num_workers=num_workers,
pin_memory=True, drop_last=True, persistent_workers=num_workers > 0)
model = SiamCenterNet(backbone=backbone, pretrained=True).to(device)
opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=WEIGHT_DECAY)
sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=lr, total_steps=epochs * len(dl),
pct_start=0.1)
scaler = torch.amp.GradScaler(enabled="cuda" in device)
start_epoch, history = 0, []
ckpt_path = out_dir / "last.pt"
if resume and ckpt_path.exists():
state = torch.load(ckpt_path, map_location="cpu")
model.load_state_dict(state["model"])
opt.load_state_dict(state["opt"])
sched.load_state_dict(state["sched"])
start_epoch = state["epoch"] + 1
history = state.get("history", [])
print(f"resumed from epoch {start_epoch}", flush=True)
best_f1 = max([h.get("f1", 0.0) for h in history], default=0.0)
for epoch in range(start_epoch, epochs):
model.train()
sampler.set_epoch(epoch)
t0, running = time.time(), {}
for step, batch_data in enumerate(dl):
a = batch_data["a"].to(device, non_blocking=True)
b = batch_data["b"].to(device, non_blocking=True)
targets = {k: v.to(device, non_blocking=True) for k, v in batch_data.items()
if k not in ("a", "b")}
with torch.autocast(device_type="cuda" if "cuda" in device else "cpu",
dtype=torch.bfloat16, enabled="cuda" in device):
out = model(a, b)
losses = total_loss(out, targets)
opt.zero_grad(set_to_none=True)
scaler.scale(losses["total"]).backward()
scaler.unscale_(opt)
torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)
scaler.step(opt)
scaler.update()
sched.step()
for k, v in losses.items():
running[k] = running.get(k, 0.0) + float(v.detach())
if step % 50 == 0:
msg = " ".join(f"{k}={running[k] / (step + 1):.3f}" for k in sorted(running))
print(f"[fold {fold}] epoch {epoch} step {step}/{len(dl)} {msg}", flush=True)
record = {"epoch": epoch, "seconds": round(time.time() - t0, 1),
**{k: running[k] / max(1, len(dl)) for k in running}}
if (epoch + 1) % eval_every == 0 or epoch == epochs - 1:
record.update(evaluate(model, val_idx, device, gt))
print(f"[fold {fold}] epoch {epoch} val {record}", flush=True)
if record.get("f1", 0.0) >= best_f1:
best_f1 = record["f1"]
torch.save({"model": model.state_dict(), "epoch": epoch, "metric": record},
out_dir / "best.pt")
history.append(record)
torch.save({"model": model.state_dict(), "opt": opt.state_dict(),
"sched": sched.state_dict(), "epoch": epoch, "history": history},
ckpt_path)
(out_dir / "history.json").write_text(json.dumps(history, indent=2))
if on_epoch_end is not None:
on_epoch_end() # commit the volume so a killed run resumes from here
return {"fold": fold, "best_f1": best_f1, "history": history[-1] if history else {}}