142 lines
6.0 KiB
Python
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)
|