完整训练循环、Loss 监控、Early Stopping、正则化防过拟合。训练不是玄学,是有迹可循的科学。
点击「开始训练」观察 4 阶段循环:Forward 数据前向传播 → Loss 计算误差 → Backward 梯度反向 → Update 参数更新
训练时随机失活神经元,防止协同依赖
损失函数增加权重惩罚项
训练时对样本做随机变换扩充数据
验证集效果连续N轮不降 → 停止训练
import torch; import torch.nn as nn; import torch.optim as optim from torch.cuda.amp import autocast, GradScaler # 混合精度加速 from tqdm import tqdm from utils.metrics import compute_metrics from utils.logging import WandBLogger def train(config, model, train_loader, val_loader): device = torch.device(config.device) model = model.to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.AdamW(model.parameters(), lr=config.lr, weight_decay=config.wd) scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=config.epochs) scaler = GradScaler() # 混合精度 logger = WandBLogger(config) best_val_acc = 0 patience_counter = 0 for epoch in range(config.epochs): # ===== 训练阶段 ===== model.train() train_loss = 0 for images, labels in tqdm(train_loader): images, labels = images.to(device), labels.to(device) optimizer.zero_grad() with autocast(): outputs = model(images) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() train_loss += loss.item() scheduler.step() # ===== 验证阶段 ===== model.eval() val_loss, val_acc = 0, 0 with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) val_loss += criterion(outputs, labels).item() val_acc += (outputs.argmax(1) == labels).sum().item() val_acc /= len(val_loader.dataset) # ===== 记录日志 ===== logger.log({ 'epoch': epoch, 'train_loss': train_loss / len(train_loader), 'val_loss': val_loss / len(val_loader), 'val_acc': val_acc, 'lr': optimizer.param_groups[0]['lr'] }) # ===== Early Stopping ===== if val_acc > best_val_acc: best_val_acc = val_acc torch.save(model.state_dict(), 'models/best.pt') patience_counter = 0 else: patience_counter += 1 if patience_counter > config.patience: print(f'Early stopping at epoch {epoch}') break return best_val_acc