🔥 Stage 05 · 模型训练

反复迭代,收敛 Loss

完整训练循环、Loss 监控、Early Stopping、正则化防过拟合。训练不是玄学,是有迹可循的科学。

🎯 训练四要素

① 训练循环
Epoch → Forward → Loss → Backward → Optimizer.step()
② Loss 监控
Train Loss ↓ + Val Loss 对比,判断过拟合
③ Early Stopping
Val Loss 连续 N 轮不降就停
④ 正则化
Dropout / L1-L2 / 数据增强

训练循环 · 交互式动画

点击「开始训练」观察 4 阶段循环:Forward 数据前向传播 → Loss 计算误差 → Backward 梯度反向 → Update 参数更新

0.01
① Forward 前向 ② Loss 误差 ③ Backward 反向 ④ Update 更新
等待开始训练...
点击「▶ 开始训练」或「⏭ 单步执行」启动动画。神经网络将展示一个 3→4→3→1 的全连接网络在每个 batch 中的一次完整训练循环。
🧠 3-4-3-1 全连接网络 等待开始...
x₁ x₂ x₃ h⁽¹⁾₁ h⁽¹⁾₂ h⁽¹⁾₃ h⁽¹⁾₄ h⁽²⁾₁ h⁽²⁾₂ h⁽²⁾₃ ŷ 0.0 0.0 0.0 输入层 0.0 0.0 0.0 0.0 隐藏层 1 0.0 0.0 0.0 隐藏层 2 0.0 输出层
📉 Loss 曲线(实时) Train Val
Train: -- Val: -- Epoch: 0/50
📋 训练日志 空闲
等待训练开始...

🛡️ 防止过拟合的正则化策略

Dropout

训练时随机失活神经元,防止协同依赖

nn.Dropout(0.5)
常用值:0.3 ~ 0.7
⚖️

L1/L2 正则

损失函数增加权重惩罚项

weight_decay=1e-4
L2最常用(权重衰减)

数据增强

训练时对样本做随机变换扩充数据

CV:翻转/旋转/色彩
NLP:同义词替换/回译
🛑

Early Stopping

验证集效果连续N轮不降 → 停止训练

patience=5~10
最省心的防过拟合手段

📝 完整训练循环模板(PyTorch)

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
← 返回模型选择 下一章:评估与迭代 →