AI模型训练优化实战指南:从损失函数到分布式训练
前言:训练优化是模型成败的关键
模型架构决定了上限,训练优化决定了你能多接近这个上限。同样的模型,好的训练策略能让效果提升 10-20%,差的训练策略可能让模型根本不收敛。
训练优化的核心问题:
- 收敛性:能不能学会
- 速度:学得多快
- 泛化性:学到的东西能不能迁移
- 稳定性:训练过程会不会崩
一、损失函数
1.1 常见损失函数
| 任务 | 损失函数 | 公式 |
|---|---|---|
| 回归 | MSE | ((y - ŷ)²).mean() |
| 回归 | MAE | `( |
| 二分类 | BCE | -[y·log(σ) + (1-y)·log(1-σ)] |
| 多分类 | CrossEntropy | -Σ yᵢ·log(softmax(x)ᵢ) |
1.2 不平衡数据的损失函数
import torch
import torch.nn as nn
import torch.nn.functional as F
# 1. Focal Loss(解决类别不平衡)
class FocalLoss(nn.Module):
def __init__(self, alpha=0.25, gamma=2.0):
super().__init__()
self.alpha = alpha
self.gamma = gamma
def forward(self, pred, target):
ce_loss = F.cross_entropy(pred, target, reduction='none')
pt = torch.exp(-ce_loss)
loss = self.alpha * (1 - pt) ** self.gamma * ce_loss
return loss.mean()
# 2. 类别加权交叉熵
class WeightedCrossEntropy(nn.Module):
def __init__(self, class_weights):
super().__init__()
self.class_weights = torch.tensor(class_weights, dtype=torch.float32)
def forward(self, pred, target):
return F.cross_entropy(pred, target, weight=self.class_weights.to(pred.device))
# 3. Label Smoothing(防止过自信)
def label_smoothing_loss(pred, target, smoothing=0.1):
n_classes = pred.size(-1)
log_preds = F.log_softmax(pred, dim=-1)
nll_loss = F.nll_loss(log_preds, target, reduction='none')
smooth_loss = -log_preds.mean(dim=-1)
return (1 - smoothing) * nll_loss + smoothing * smooth_loss
1.3 损失函数选择建议
- 回归任务:默认 MSE,对异常值敏感用 MAE 或 Huber Loss
- 二分类:BCE With Logits(数值稳定)
- 多分类:CrossEntropy
- 类别不平衡:Focal Loss 或加权 CE
- 目标检测:Focal Loss + Smooth L1
- 图像分割:Dice Loss + CE
- 对比学习:InfoNCE / Triplet Loss
二、学习率调度
2.1 为什么需要学习率调度
2.2 常见调度策略
from torch.optim.lr_scheduler import (
StepLR, CosineAnnealingLR, CosineAnnealingWarmRestarts,
OneCycleLR, ReduceLROnPlateau, LambdaLR
)
# 1. Warmup + Cosine(Transformer 标配)
def get_cosine_schedule_with_warmup(optimizer, num_warmup_steps, num_training_steps):
def lr_lambda(current_step):
if current_step < num_warmup_steps:
return float(current_step) / float(max(1, num_warmup_steps))
progress = float(current_step - num_warmup_steps) / float(max(1, num_training_steps - num_warmup_steps))
return max(0.0, 0.5 * (1.0 + math.cos(math.pi * progress)))
return LambdaLR(optimizer, lr_lambda)
# 2. OneCycleLR(快速收敛)
scheduler = OneCycleLR(
optimizer,
max_lr=1e-3,
total_steps=num_training_steps,
pct_start=0.3, # 30% 时间升学习率
anneal_strategy='cos'
)
# 3. ReduceLROnPlateau(验证集不降就降学习率)
scheduler = ReduceLROnPlateau(
optimizer,
mode='min',
factor=0.5, # 学习率乘以 0.5
patience=3, # 3 个 epoch 不降就触发
min_lr=1e-7
)
2.3 调度策略对比
| 策略 | 适用场景 | 特点 |
|---|---|---|
| StepLR | 简单任务 | 固定间隔衰减 |
| CosineAnnealing | 大部分任务 | 平滑衰减 |
| Warmup + Cosine | Transformer | 先升后降 |
| OneCycleLR | 快速训练 | 超参数收敛 |
| ReduceLROnPlateau | 不确定 LR | 自适应 |
三、正则化技术
3.1 L1 / L2 正则化
# 在优化器中加权重衰减(等价于 L2 正则化)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4)
# 自定义 L1 正则化
def l1_regularization(model, lambda_l1=0.01):
l1_loss = sum(p.abs().sum() for p in model.parameters())
return lambda_l1 * l1_loss
loss = task_loss + l1_regularization(model)
3.2 Dropout
class Model(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(768, 256)
self.dropout = nn.Dropout(0.5)
self.fc2 = nn.Linear(256, 10)
def forward(self, x):
x = F.relu(self.fc1(x))
x = self.dropout(x) # 训练时随机失活 50%
return self.fc2(x)
Dropout 经验:
- CNN:0.1-0.3
- 全连接层:0.3-0.5
- Transformer:0.1
3.3 归一化(Normalization)
归一化要分两层看:数据层(预处理输入特征)解决特征量级差异导致的梯度失衡;网络层(在每一层做归一化)缓解深层网络的内部协变量偏移——每层输入分布不断漂移,训练慢且不稳。
输入数据归一化(预处理层)
最基础的一步是把输入特征拉到合理量级,否则量级悬殊(如一个特征 0-10000,另一个 0-1)会让大数值特征"抢走"梯度,小数值特征几乎不更新,进而梯度爆炸或 NaN。
# Min-Max:线性映射到 [0, 1],适合分布均匀、无明显离群值
def minmax_normalize(X, feature_range=(0, 1)):
min_val, max_val = X.min(axis=0), X.max(axis=0)
normalized = (X - min_val) / (max_val - min_val)
return normalized * (feature_range[1] - feature_range[0]) + feature_range[0]
# Z-Score:减均值除标准差,对离群值更鲁棒(基于分布统计量而非极值)
def zscore_normalize(X):
return (X - X.mean(axis=0)) / X.std(axis=0)
踩坑要点:
- 训练/测试统计量必须一致:测试集要用训练集算出的均值和标准差,否则分布对不上。
- 标准差为 0 要加 epsilon:某特征恒定取值时
std=0会除爆,需(std + 1e-8)。 - 流式数据用滑动窗口或增量统计:实时场景拿不到全量数据,需在线更新统计量。
# 训练时保存统计量,测试时直接复用(不要重新算)
train_mean, train_std = X_train.mean(axis=0), X_train.std(axis=0)
X_train_norm = (X_train - train_mean) / (train_std + 1e-8)
X_test_norm = (X_test - train_mean) / (train_std + 1e-8) # 用训练集统计量
BatchNorm:把归一化搬进网络
数据预处理只解决输入层,深层网络每层输入分布仍在漂移。BatchNorm 在每一层按通道独立做标准化,再用可学习的 γ(缩放)和 β(平移)恢复表达能力——强制零均值、单位方差会破坏表达力,这两个参数让网络"找回"必要的分布。
def batch_norm(x, gamma, beta, running_mean, running_var, eps=1e-5, momentum=0.1, training=True):
if training:
mean = x.mean(dim=(0, 2, 3), keepdim=True) # 跨 (B, H, W) 按通道求均值
var = x.var(dim=(0, 2, 3), keepdim=True)
running_mean = momentum * running_mean + (1 - momentum) * mean.data # 更新全局统计量
running_var = momentum * running_var + (1 - momentum) * var.data
else:
mean, var = running_mean, running_var # 推理用固定的全局统计量
x_norm = (x - mean) / torch.sqrt(var + eps)
return gamma * x_norm + beta
为什么有效: 允许更大的学习率(降低对输入变化的敏感度)、减少对初始化的敏感度、带来轻微正则化效果(batch 统计量的噪声相当于数据增强,但不如 Dropout 明显)。
BatchNorm 踩坑:
- Batch 太小会抖动:batch size < 8 时统计量方差大,效果明显下降,应改用 GroupNorm 或 SyncBatchNorm。
- 推理忘切 eval:训练模式保存、推理时没切
model.eval(),仍会用当前 batch 统计量,导致结果不稳定且不可复现。 - 多 GPU 统计量不同步:DDP 下每卡各算各的,需用
SyncBatchNorm跨卡同步。 - RNN/LSTM 维度对不上:变长序列按时间步无法对齐,这种场景换 LayerNorm。
model.train() # 训练:统计量会更新
model.eval() # 推理:统计量固定
# 多 GPU 训练时把普通 BN 换成跨卡同步版本
from torch.nn import SyncBatchNorm
model = SyncBatchNorm.convert_sync_batchnorm(model)
LayerNorm:序列模型的标配
LayerNorm 不跨样本算统计量,而是在单个样本内部跨特征维度计算,因此不受 batch size 影响,天然适合变长序列、小 batch 场景,是 Transformer/RNN 的标配。一句话区分:BatchNorm 是"横向"的(跨样本对齐同一通道),LayerNorm 是"纵向"的(样本内部把不同特征拉到同量级)。
def layer_norm(x, gamma, beta, eps=1e-5):
mean = x.mean(dim=-1, keepdim=True)
var = x.var(dim=-1, keepdim=True)
return gamma * (x - mean) / torch.sqrt(var + eps) + beta
其他归一化方案
| 方案 | 归一化范围 | 依赖 batch | 典型场景 |
|---|---|---|---|
| InstanceNorm | 单样本、单通道 | 否 | 风格迁移(保留风格、抑制内容) |
| GroupNorm | 通道分组组内 | 否 | 目标检测/分割(batch 小,Mask R-CNN、Detectron2 默认) |
| WeightNorm | 对权重归一化 | 否 | 计算量小,实践中少用(BN/GN 已覆盖大多数场景) |
GroupNorm 是 LayerNorm 和 BatchNorm 的折中:既不依赖 batch size,又能利用通道间信息,是检测/分割等小 batch 任务的主流选择。
# GroupNorm:64 通道分 8 组,每组 8 个通道一起归一化
self.gn = nn.GroupNorm(num_groups=8, num_channels=64)
选型建议
- 图像分类/CNN:batch ≥ 8 用 BatchNorm;batch < 8 用 GroupNorm 或 SyncBatchNorm。
- NLP/Transformer:首选 LayerNorm(几乎默认)。
- 目标检测/分割:batch 通常 2-4,GroupNorm 是主流。
- 多 GPU 训练:用 SyncBatchNorm 替代普通 BatchNorm。
参考论文: BatchNorm · LayerNorm · GroupNorm
3.4 数据增强
import albumentations as A
transform = A.Compose([
A.RandomResizedCrop(224, 224),
A.HorizontalFlip(p=0.5),
A.ColorJitter(brightness=0.2, contrast=0.2),
A.GaussianBlur(p=0.3),
A.Normalize(),
A.CoarseDropout(max_holes=8, max_height=32, max_width=32, p=0.5),
])
3.5 Mixup 和 CutMix
def mixup_data(x, y, alpha=0.2):
lam = np.random.beta(alpha, alpha)
batch_size = x.size(0)
index = torch.randperm(batch_size)
mixed_x = lam * x + (1 - lam) * x[index]
y_a, y_b = y, y[index]
return mixed_x, y_a, y_b, lam
def mixup_criterion(criterion, pred, y_a, y_b, lam):
return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)
3.6 早停(Early Stopping)
class EarlyStopping:
def __init__(self, patience=5, min_delta=0):
self.patience = patience
self.min_delta = min_delta
self.counter = 0
self.best_loss = float('inf')
self.early_stop = False
def __call__(self, val_loss):
if val_loss < self.best_loss - self.min_delta:
self.best_loss = val_loss
self.counter = 0
else:
self.counter += 1
if self.counter >= self.patience:
self.early_stop = True
# 使用
early_stopping = EarlyStopping(patience=5)
for epoch in range(num_epochs):
train_loss = train_one_epoch()
val_loss = validate()
early_stopping(val_loss)
if early_stopping.early_stop:
print(f"Early stopping at epoch {epoch}")
break
四、梯度优化
4.1 梯度累积(小显存训练大模型)
accumulation_steps = 4 # 等效 batch size = batch_size * 4
optimizer.zero_grad()
for i, batch in enumerate(dataloader):
outputs = model(batch)
loss = criterion(outputs, targets) / accumulation_steps # 关键:除以累积步数
loss.backward()
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
4.2 梯度裁剪
# 防止梯度爆炸
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
# 或者按值裁剪
torch.nn.utils.clip_grad_value_(model.parameters(), clip_value=0.5)
4.3 梯度下降算法对比
# SGD(带动量)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9)
# Adam(自适应)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
# AdamW(推荐 Transformer 用)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01)
| 优化器 | 适用场景 | 特点 |
|---|---|---|
| SGD + Momentum | CNN | 收敛慢但泛化好 |
| Adam | 通用 | 自适应学习率 |
| AdamW | Transformer | 解耦权重衰减 |
| LAMB | 大 batch 训练 | Layer-wise 自适应 |
五、混合精度训练
5.1 为什么用混合精度
- FP32:精度高但显存大、速度慢
- FP16:显存减半、速度快,但可能数值溢出
- BF16:动态范围大,数值稳定
5.2 PyTorch AMP 实现
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for batch in dataloader:
optimizer.zero_grad()
# 前向传播用混合精度
with autocast():
outputs = model(batch)
loss = criterion(outputs, targets)
# 反向传播用梯度缩放
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
5.3 BF16 vs FP16
# BF16 更稳定,推荐 A100/H100 用
model = model.to(torch.bfloat16)
# FP16 需要 GradScaler
model = model.to(torch.float16)
六、分布式训练
6.1 数据并行 vs 模型并行
| 方式 | 说明 | 适用 |
|---|---|---|
| 数据并行 (DDP) | 每卡完整模型,不同数据 | 模型能放进单卡 |
| 模型并行 | 模型切分到多卡 | 大模型 |
| 流水线并行 | 按层切分 | 超大模型 |
| 张量并行 | 矩阵内切分 | 超大模型 |
| ZeRO | 优化器状态/梯度/参数切分 | 大模型训练 |
6.2 PyTorch DDP
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
def setup(rank, world_size):
dist.init_process_group("nccl", rank=rank, world_size=world_size)
def train(rank, world_size):
setup(rank, world_size)
torch.cuda.set_device(rank)
model = Model().to(rank)
model = DDP(model, device_ids=[rank])
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
for epoch in range(num_epochs):
for batch in dataloader:
outputs = model(batch)
loss = criterion(outputs, targets)
loss.backward()
optimizer.step()
optimizer.zero_grad()
dist.destroy_process_group()
# 启动
torchrun --nproc_per_node=4 train.py
6.3 ZeRO 优化器(DeepSpeed)
{
"zero_optimization": {
"stage": 2,
"offload_optimizer": {
"device": "cpu"
},
"allgather_partitions": true,
"allgather_bucket_size": 5e8,
"overlap_comm": true,
"contiguous_gradients": true
},
"fp16": {
"enabled": true,
"loss_scale": 0,
"loss_scale_window": 1000
}
}
ZeRO 三个阶段:
- Stage 1:切分优化器状态(省 4x 显存)
- Stage 2:+ 切分梯度(省 8x 显存)
- Stage 3:+ 切分参数(省 N 倍显存,N=GPU 数)
6.4 断点续训
def save_checkpoint(model, optimizer, scheduler, epoch, path):
torch.save({
'epoch': epoch,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'scheduler_state_dict': scheduler.state_dict(),
}, path)
def load_checkpoint(model, optimizer, scheduler, path):
checkpoint = torch.load(path)
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
scheduler.load_state_dict(checkpoint['scheduler_state_dict'])
return checkpoint['epoch']
# 训练循环
start_epoch = 0
if resume:
start_epoch = load_checkpoint(model, optimizer, scheduler, 'checkpoint.pth')
for epoch in range(start_epoch, num_epochs):
train()
if epoch % save_interval == 0:
save_checkpoint(model, optimizer, scheduler, epoch, 'checkpoint.pth')
七、注意力机制
7.1 Self-Attention
class SelfAttention(nn.Module):
def __init__(self, embed_dim, num_heads):
super().__init__()
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
self.qkv = nn.Linear(embed_dim, embed_dim * 3)
self.proj = nn.Linear(embed_dim, embed_dim)
def forward(self, x):
B, N, C = x.shape
qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4)
q, k, v = qkv[0], qkv[1], qkv[2]
attn = (q @ k.transpose(-2, -1)) * (self.head_dim ** -0.5)
attn = attn.softmax(dim=-1)
x = (attn @ v).transpose(1, 2).reshape(B, N, C)
return self.proj(x)
7.2 注意力变体
| 变体 | 特点 |
|---|---|
| Self-Attention | 序列内部交互 |
| Cross-Attention | 两个序列交互 |
| Multi-Head | 多子空间并行 |
| Sparse Attention | 减少计算量 |
| Linear Attention | 线性复杂度 |
| Flash Attention | IO 优化 |
八、特殊训练技术
8.1 自监督学习
# SimCLR(对比学习)
class SimCLR(nn.Module):
def __init__(self, backbone, projection_dim=128):
super().__init__()
self.backbone = backbone
self.projector = nn.Sequential(
nn.Linear(backbone.dim, backbone.dim),
nn.ReLU(),
nn.Linear(backbone.dim, projection_dim)
)
def forward(self, x1, x2):
h1 = self.backbone(x1)
h2 = self.backbone(x2)
z1 = self.projector(h1)
z2 = self.projector(h2)
return z1, z2
def nt_xent_loss(z1, z2, temperature=0.5):
batch_size = z1.shape[0]
z = torch.cat([z1, z2], dim=0)
z = F.normalize(z, dim=-1)
sim = torch.matmul(z, z.T) / temperature
sim_i_j = torch.diag(sim, batch_size)
sim_j_i = torch.diag(sim, -batch_size)
positive_samples = torch.cat([sim_i_j, sim_j_i], dim=0)
mask = (~torch.eye(2 * batch_size, dtype=torch.bool, device=z.device)).float()
negative_samples = sim * mask
logits = torch.cat([positive_samples.unsqueeze(1), negative_samples], dim=1)
labels = torch.zeros(2 * batch_size, dtype=torch.long, device=z.device)
return F.cross_entropy(logits, labels)
8.2 对抗训练
# FGSM(Fast Gradient Sign Method)
def fgsm_attack(model, x, y, epsilon):
x.requires_grad = True
output = model(x)
loss = F.cross_entropy(output, y)
loss.backward()
perturbation = epsilon * x.grad.sign()
adversarial_x = x + perturbation
return adversarial_x.detach()
# PGD(Project Gradient Descent,更强的攻击)
def pgd_attack(model, x, y, epsilon, alpha, num_steps):
x_adv = x.clone().detach()
for _ in range(num_steps):
x_adv.requires_grad = True
output = model(x_adv)
loss = F.cross_entropy(output, y)
loss.backward()
perturbation = alpha * x_adv.grad.sign()
x_adv = x_adv + perturbation
x_adv = torch.clamp(x_adv, x - epsilon, x + epsilon)
x_adv = x_adv.detach()
return x_adv
# 对抗训练
for x, y in dataloader:
x_adv = pgd_attack(model, x, y, epsilon=0.03, alpha=0.01, num_steps=7)
# 用对抗样本训练
outputs = model(x_adv)
loss = criterion(outputs, y)
8.3 知识蒸馏
def distillation_loss(student_logits, teacher_logits, temperature=4.0):
soft_targets = F.softmax(teacher_logits / temperature, dim=-1)
soft_prob = F.log_softmax(student_logits / temperature, dim=-1)
return F.kl_div(soft_prob, soft_targets, reduction='batchmean') * (temperature ** 2)
# 训练
for x, y in dataloader:
with torch.no_grad():
teacher_logits = teacher_model(x)
student_logits = student_model(x)
# 硬标签 + 软标签
hard_loss = F.cross_entropy(student_logits, y)
soft_loss = distillation_loss(student_logits, teacher_logits)
loss = 0.5 * hard_loss + 0.5 * soft_loss
8.4 多任务学习
class MultiTaskModel(nn.Module):
def __init__(self, backbone):
super().__init__()
self.backbone = backbone
self.classification_head = nn.Linear(backbone.dim, num_classes)
self.regression_head = nn.Linear(backbone.dim, 1)
def forward(self, x):
features = self.backbone(x)
return {
'classification': self.classification_head(features),
'regression': self.regression_head(features)
}
# 动态权重
class UncertaintyWeighting(nn.Module):
def __init__(self, num_tasks):
super().__init__()
self.log_vars = nn.Parameter(torch.zeros(num_tasks))
def forward(self, losses):
total = 0
for i, loss in enumerate(losses):
precision = torch.exp(-self.log_vars[i])
total += precision * loss + self.log_vars[i]
return total
8.5 课程学习
class CurriculumSampler:
"""从简单到困难的课程学习"""
def __init__(self, dataset, difficulty_fn):
self.difficulties = [difficulty_fn(item) for item in dataset]
self.indices = list(range(len(dataset)))
def get_subset(self, epoch, total_epochs):
"""根据 epoch 返回不同难度的子集"""
progress = epoch / total_epochs
max_difficulty = self.difficulties.quantile(progress)
return [i for i, d in zip(self.indices, self.difficulties) if d <= max_difficulty]
8.6 主动学习
def uncertainty_sampling(model, unlabeled_pool, n_samples):
"""选择模型最不确定的样本标注"""
model.eval()
uncertainties = []
with torch.no_grad():
for x in unlabeled_pool:
output = model(x)
prob = F.softmax(output, dim=-1)
entropy = -(prob * torch.log(prob + 1e-8)).sum()
uncertainties.append(entropy)
# 选不确定性最高的 n 个
top_indices = torch.topk(torch.tensor(uncertainties), n_samples).indices
return [unlabeled_pool[i] for i in top_indices]
九、图神经网络(GNN)
9.1 GCN(图卷积网络)
import torch_geometric.nn as gnn
class GCN(nn.Module):
def __init__(self, in_channels, hidden_channels, num_classes):
super().__init__()
self.conv1 = gnn.GCNConv(in_channels, hidden_channels)
self.conv2 = gnn.GCNConv(hidden_channels, num_classes)
def forward(self, x, edge_index):
x = F.relu(self.conv1(x, edge_index))
x = F.dropout(x, p=0.5, training=self.training)
x = self.conv2(x, edge_index)
return x
9.2 GAT(图注意力网络)
class GAT(nn.Module):
def __init__(self, in_channels, hidden_channels, num_classes, heads=8):
super().__init__()
self.conv1 = gnn.GATConv(in_channels, hidden_channels, heads=heads)
self.conv2 = gnn.GATConv(hidden_channels * heads, num_classes, heads=1)
def forward(self, x, edge_index):
x = F.elu(self.conv1(x, edge_index))
x = self.conv2(x, edge_index)
return x
十、Word2Vec 嵌入
import gensim
from gensim.models import Word2Vec
# 训练 Word2Vec
sentences = [['hello', 'world'], ['foo', 'bar']]
model = Word2Vec(
sentences,
vector_size=100,
window=5,
min_count=1,
workers=4,
sg=0 # 0=CBOW, 1=Skip-gram
)
# 使用
vector = model.wv['hello']
similar = model.wv.most_similar('hello')
十一、混合专家(MoE)
class MoELayer(nn.Module):
def __init__(self, experts, gate):
super().__init__()
self.experts = nn.ModuleList(experts)
self.gate = gate # 路由网络
def forward(self, x):
gate_scores = self.gate(x) # [batch, num_experts]
top_k_scores, top_k_indices = torch.topk(gate_scores, k=2, dim=-1)
output = torch.zeros_like(x)
for i, expert in enumerate(self.experts):
mask = (top_k_indices == i).any(dim=-1)
if mask.any():
expert_input = x[mask]
expert_output = expert(expert_input)
weights = top_k_scores[mask].max(dim=-1).values.unsqueeze(-1)
output[mask] += expert_output * weights
return output
十二、训练踩坑总结
坑一:loss 不下降
排查顺序:
- 学习率太大/太小 → 调到 1e-4 试试
- 数据问题 → 检查标签是否对齐
- 梯度消失/爆炸 → 加 BatchNorm、梯度裁剪
- 损失函数错误 → 检查是否匹配任务
坑二:训练集好测试集差(过拟合)
解决:
- 加 Dropout / L2
- 数据增强
- 早停
- 减小模型
坑三:训练集都不好(欠拟合)
解决:
- 增大模型
- 训练更久
- 增大学习率
- 检查数据质量
坑四:训练不稳定(loss 振荡)
解决:
- 降低学习率
- 加 warmup
- 用 AdamW
- 梯度裁剪
坑五:OOM(显存不足)
解决:
- 减小 batch size
- 梯度累积
- 混合精度
- 梯度检查点
- ZeRO
坑六:训练速度慢
解决:
- 用 SSD
- 增加 num_workers
- pin_memory=True
- 混合精度
- 分布式训练
十三、训练配置模板
# 完整训练配置模板
config = {
# 模型
'model_name': 'resnet50',
'pretrained': True,
'num_classes': 1000,
# 数据
'batch_size': 64,
'num_workers': 8,
'pin_memory': True,
# 训练
'epochs': 100,
'learning_rate': 1e-4,
'weight_decay': 0.01,
'warmup_steps': 1000,
'max_grad_norm': 1.0,
# 优化
'optimizer': 'AdamW',
'scheduler': 'cosine_with_warmup',
'mixed_precision': True,
'gradient_accumulation_steps': 1,
# 正则化
'dropout': 0.1,
'label_smoothing': 0.1,
# 早停
'early_stopping_patience': 5,
# 保存
'save_interval': 5,
'save_top_k': 3,
}
十四、写在最后
训练优化没有银弹,每个任务都需要根据数据、模型、资源来调整。
几条核心原则:
- 从简单开始:SGD + 固定学习率 → 逐步加复杂策略
- 监控是基础:loss 曲线、梯度范数、学习率
- 超参数敏感:学习率 > batch size > 其他
- 小数据用强正则:Dropout、数据增强、早停
- 大数据用弱正则:让模型充分学习
- 混合精度是标配:免费的速度和显存提升
- 断点续训必须有:训练中断不丢失进度
训练是一门经验科学,多实验、多记录、多对比。论文里的最佳实践不一定适合你的数据,要自己跑实验验证。
本文整合了 20+ 篇模型训练优化相关文章,涵盖损失函数、学习率调度、正则化、梯度优化、混合精度、分布式训练、注意力机制、自监督学习、对抗训练、知识蒸馏、多任务学习、课程学习、主动学习、图神经网络、Word2Vec、混合专家等核心技术。
版权声明: 本文首发于 指尖魔法屋-AI模型训练优化实战指南:从损失函数到分布式训练(https://blog.thinkmoon.cn/post/ai-training-optimization-comprehensive-guide/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。