本模块包含交互式超参数模拟器、学习率调度器可视化、正则化策略选择器。调参技巧一网打尽。
调整下面的参数,观察对训练效果的模拟影响
训练时随机失活神经元,防止协同依赖
损失函数增加权重惩罚项
训练时对样本做随机变换扩充数据
验证集效果连续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
遍历所有组合,简单但组合爆炸
param_grid = {
'lr': [1e-3, 1e-4, 1e-5],
'wd': [0, 1e-4, 1e-2]
}
随机采样,高概率找到好区域
params = {
'lr': uniform(1e-5, 1e-2),
'wd': loguniform(1e-5, 1e-1)
}
用概率模型指导搜索,高效