Buckets:

KaisResearch's picture
download
raw
6.21 kB
"""Training entry point for the lip-reading word-classification baseline.
Quick smoke test (no data needed) — verifies the whole pipeline runs:
python -m src.train --dummy --epochs 2 --batch_size 8 --num_workers 0
Real training once data is prepared (see scripts/prepare_data.py):
python -m src.train --lrw_root data/lrw --custom_root data/custom --epochs 40
"""
from __future__ import annotations
import math
import os
import sys
import time
if __package__ in (None, ""): # allow `python src/train.py` per repo README
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from src.config import Config
from src.dataset import build_datasets, save_label_map
from src.model import build_model
from src.utils import (AverageMeter, accuracy, resolve_device,
save_checkpoint, set_seed)
def make_loader(ds, cfg, shuffle):
if ds is None or len(ds) == 0:
return None
return DataLoader(ds, batch_size=cfg.batch_size, shuffle=shuffle,
num_workers=cfg.num_workers, pin_memory=True,
drop_last=shuffle)
def lr_at(step, total_steps, warmup_steps, base_lr):
"""Linear warmup then cosine decay."""
if step < warmup_steps:
return base_lr * (step + 1) / max(warmup_steps, 1)
progress = (step - warmup_steps) / max(total_steps - warmup_steps, 1)
return 0.5 * base_lr * (1 + math.cos(math.pi * progress))
def run_epoch(model, loader, criterion, device, optimizer=None,
scaler=None, cfg=None, epoch=0, sched=None):
train = optimizer is not None
model.train(train)
loss_m, acc_m = AverageMeter(), AverageMeter()
torch.set_grad_enabled(train)
for step, (clips, targets) in enumerate(loader):
clips = clips.to(device, non_blocking=True)
targets = targets.to(device, non_blocking=True)
if train and sched is not None:
for g in optimizer.param_groups:
g["lr"] = sched(epoch * len(loader) + step)
with torch.autocast(device_type=device.type,
enabled=(cfg.amp and device.type == "cuda")):
logits = model(clips)
loss = criterion(logits, targets)
if train:
optimizer.zero_grad(set_to_none=True)
if scaler is not None:
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_clip)
scaler.step(optimizer)
scaler.update()
else:
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_clip)
optimizer.step()
(top1,) = accuracy(logits, targets, topk=(1,))
loss_m.update(loss.item(), clips.size(0))
acc_m.update(top1, clips.size(0))
if train and step % cfg.log_every == 0:
lr = optimizer.param_groups[0]["lr"]
print(f" epoch {epoch} step {step}/{len(loader)} "
f"loss {loss_m.avg:.4f} acc {acc_m.avg:.3f} lr {lr:.2e}")
return loss_m.avg, acc_m.avg
def main():
cfg = Config.from_args()
set_seed(cfg.seed)
device = resolve_device(cfg.device)
print(f"Device: {device}")
train_ds, val_ds, test_ds, label_map = build_datasets(cfg)
cfg.num_classes = len(label_map)
save_label_map(label_map, cfg.out_dir)
print(f"Classes: {cfg.num_classes} | train {len(train_ds)} "
f"val {len(val_ds)} test {len(test_ds)}")
train_loader = make_loader(train_ds, cfg, shuffle=True)
val_loader = make_loader(val_ds, cfg, shuffle=False)
test_loader = make_loader(test_ds, cfg, shuffle=False)
if train_loader is None:
raise SystemExit("Training set is empty — check your data roots.")
model = build_model(cfg).to(device)
n_params = sum(p.numel() for p in model.parameters()) / 1e6
print(f"Model parameters: {n_params:.1f}M")
criterion = nn.CrossEntropyLoss(label_smoothing=cfg.label_smoothing)
optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr,
weight_decay=cfg.weight_decay)
scaler = (torch.amp.GradScaler("cuda")
if cfg.amp and device.type == "cuda" else None)
total_steps = cfg.epochs * len(train_loader)
warmup_steps = cfg.warmup_epochs * len(train_loader)
sched = lambda s: lr_at(s, total_steps, warmup_steps, cfg.lr)
start_epoch, best_acc = 0, 0.0
if cfg.resume and os.path.isfile(cfg.resume):
ckpt = torch.load(cfg.resume, map_location=device)
model.load_state_dict(ckpt["model"])
optimizer.load_state_dict(ckpt["optimizer"])
start_epoch = ckpt["epoch"] + 1
best_acc = ckpt.get("best_acc", 0.0)
print(f"Resumed from {cfg.resume} at epoch {start_epoch}")
for epoch in range(start_epoch, cfg.epochs):
t0 = time.time()
tr_loss, tr_acc = run_epoch(model, train_loader, criterion, device,
optimizer, scaler, cfg, epoch, sched)
if val_loader is not None:
va_loss, va_acc = run_epoch(model, val_loader, criterion, device,
cfg=cfg, epoch=epoch)
else:
va_loss, va_acc = tr_loss, tr_acc
is_best = va_acc > best_acc
best_acc = max(best_acc, va_acc)
save_checkpoint({
"epoch": epoch, "model": model.state_dict(),
"optimizer": optimizer.state_dict(), "best_acc": best_acc,
"label_map": label_map, "config": vars(cfg),
}, cfg.out_dir, is_best)
print(f"[epoch {epoch}] train loss {tr_loss:.4f} acc {tr_acc:.3f} | "
f"val loss {va_loss:.4f} acc {va_acc:.3f} | "
f"best {best_acc:.3f} | {time.time()-t0:.0f}s"
+ (" *" if is_best else ""))
if test_loader is not None:
te_loss, te_acc = run_epoch(model, test_loader, criterion, device,
cfg=cfg)
print(f"[test] loss {te_loss:.4f} acc {te_acc:.3f}")
if __name__ == "__main__":
main()

Xet Storage Details

Size:
6.21 kB
·
Xet hash:
14da241e6088babf3e2f590d2026399db4f57b107e63d69d1a07d1fec023dcfc

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.