Buckets:
| """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.