Files
LDETR_V1/engine/train.py
2026-08-18 18:50:32 +05:30

142 lines
6.0 KiB
Python

"""Phase 1 training loop: AdamW (discriminative LR), warmup+cosine, grad clipping,
checkpointing, TuSimple accuracy/FP/FN eval per epoch. AMP/EMA are skipped for now
since this runs on CPU (no usable CUDA on this machine -- see project notes); both
are one-line additions once GPU training is available.
"""
from __future__ import annotations
import math
import os
import time
import torch
from torch.utils.data import DataLoader
from data.tusimple import TuSimpleMaskDataset, collate_fn
from models.laneformer import LaneFormer
from models.losses import compute_loss
from engine.evaluate import evaluate
def build_optimizer(model: LaneFormer, lr_backbone: float, lr_new: float, weight_decay: float) -> torch.optim.Optimizer:
backbone_params = list(model.backbone.parameters())
backbone_ids = {id(p) for p in backbone_params}
new_params = [p for p in model.parameters() if id(p) not in backbone_ids]
return torch.optim.AdamW(
[
{"params": backbone_params, "lr": lr_backbone},
{"params": new_params, "lr": lr_new},
],
weight_decay=weight_decay,
)
def build_scheduler(optimizer: torch.optim.Optimizer, warmup_steps: int, total_steps: int):
def lr_lambda(step: int) -> float:
if step < warmup_steps:
return step / max(1, warmup_steps)
progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)
return 0.5 * (1.0 + math.cos(math.pi * min(progress, 1.0)))
return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
def train(cfg: dict) -> None:
torch.manual_seed(cfg["seed"])
torch.set_num_threads(cfg.get("num_threads", 12))
device = torch.device("cuda" if (cfg["train"].get("use_cuda") and torch.cuda.is_available()) else "cpu")
print(f"Using device: {device}", flush=True)
data_cfg = cfg["data"]
train_ds = TuSimpleMaskDataset(
root=data_cfg["tusimple_root"], split="train",
out_w=data_cfg["input_width"], out_h=data_cfg["input_height"],
max_lanes=data_cfg["max_lanes"], num_sample_ys=data_cfg["num_sample_ys"],
val_fraction=data_cfg["val_fraction"], seed=cfg["seed"],
)
val_ds = TuSimpleMaskDataset(
root=data_cfg["tusimple_root"], split="val",
out_w=data_cfg["input_width"], out_h=data_cfg["input_height"],
max_lanes=data_cfg["max_lanes"], num_sample_ys=data_cfg["num_sample_ys"],
val_fraction=data_cfg["val_fraction"], seed=cfg["seed"],
)
train_loader = DataLoader(
train_ds, batch_size=data_cfg["batch_size"], shuffle=True,
collate_fn=collate_fn, num_workers=data_cfg["num_workers"],
)
val_loader = DataLoader(
val_ds, batch_size=data_cfg["batch_size"], shuffle=False,
collate_fn=collate_fn, num_workers=data_cfg["num_workers"],
)
print(f"train={len(train_ds)} val={len(val_ds)} samples", flush=True)
model_cfg = cfg["model"]
model = LaneFormer(
backbone_name=model_cfg["backbone"], pretrained=model_cfg["pretrained"],
d_model=model_cfg["fusion_channels"], max_lanes=data_cfg["max_lanes"],
encoder_layers=model_cfg["encoder_layers"], decoder_layers=model_cfg["decoder_layers"],
nhead=model_cfg["attn_heads"], ffn_dim=model_cfg["ffn_dim"], dropout=model_cfg["dropout"],
).to(device)
train_cfg = cfg["train"]
optimizer = build_optimizer(model, train_cfg["lr_backbone"], train_cfg["lr_new"], train_cfg["weight_decay"])
total_steps = train_cfg["epochs"] * len(train_loader)
scheduler = build_scheduler(optimizer, train_cfg["warmup_steps"], total_steps)
loss_cfg = cfg["loss"]
ckpt_dir = train_cfg["checkpoint_dir"]
os.makedirs(ckpt_dir, exist_ok=True)
best_acc = -1.0
global_step = 0
for epoch in range(train_cfg["epochs"]):
model.train()
epoch_start = time.time()
running = {"total": 0.0, "cls_loss": 0.0, "reg_loss": 0.0, "endpoint_loss": 0.0}
n_batches = 0
for batch in train_loader:
images = batch["images"].to(device)
targets = {k: batch[k].to(device) for k in
["target_xs", "target_valid_mask", "target_lane_valid", "target_endpoints"]}
sample_ys = batch["sample_ys"].to(device)
out = model(images)
loss_dict = compute_loss(
out, targets, sample_ys,
cls_weight=loss_cfg["cls_weight"], reg_weight=loss_cfg["reg_weight"],
endpoint_weight=loss_cfg["endpoint_weight"], bg_class_weight=loss_cfg["bg_class_weight"],
)
optimizer.zero_grad()
loss_dict["total"].backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), train_cfg["grad_clip_norm"])
optimizer.step()
scheduler.step()
for k in running:
running[k] += float(loss_dict[k].detach())
n_batches += 1
global_step += 1
if global_step % train_cfg["log_every"] == 0:
avg = {k: v / n_batches for k, v in running.items()}
lr = scheduler.get_last_lr()[-1]
print(f"epoch {epoch} step {global_step} lr {lr:.2e} "
f"loss {avg['total']:.4f} (cls {avg['cls_loss']:.4f} "
f"reg {avg['reg_loss']:.4f} ep {avg['endpoint_loss']:.4f})", flush=True)
epoch_time = time.time() - epoch_start
metrics = evaluate(model, val_loader, device, val_ds[0]["sample_ys"] if len(val_ds) else train_ds[0]["sample_ys"])
print(f"[epoch {epoch}] time={epoch_time/60:.1f}min val_acc={metrics['accuracy']:.4f} "
f"fp={metrics['fp']:.4f} fn={metrics['fn']:.4f}", flush=True)
ckpt_path = os.path.join(ckpt_dir, "last.pt")
torch.save({"model": model.state_dict(), "epoch": epoch, "metrics": metrics}, ckpt_path)
if metrics["accuracy"] > best_acc:
best_acc = metrics["accuracy"]
torch.save({"model": model.state_dict(), "epoch": epoch, "metrics": metrics},
os.path.join(ckpt_dir, "best.pt"))
print(f" new best (acc={best_acc:.4f}), saved best.pt", flush=True)