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 为什么需要学习率调度

graph LR A[高学习率] --> B[快速接近最优] B --> C[降低学习率] C --> D[精细收敛]

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 + CosineTransformer先升后降
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)

踩坑要点:

  1. 训练/测试统计量必须一致:测试集要用训练集算出的均值和标准差,否则分布对不上。
  2. 标准差为 0 要加 epsilon:某特征恒定取值时 std=0 会除爆,需 (std + 1e-8)
  3. 流式数据用滑动窗口或增量统计:实时场景拿不到全量数据,需在线更新统计量。
# 训练时保存统计量,测试时直接复用(不要重新算)
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 踩坑:

  1. Batch 太小会抖动:batch size < 8 时统计量方差大,效果明显下降,应改用 GroupNorm 或 SyncBatchNorm。
  2. 推理忘切 eval:训练模式保存、推理时没切 model.eval(),仍会用当前 batch 统计量,导致结果不稳定且不可复现。
  3. 多 GPU 统计量不同步:DDP 下每卡各算各的,需用 SyncBatchNorm 跨卡同步。
  4. 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。
flowchart TD A[选择归一化方案] --> B{数据类型?} B -->|图像/CNN| C{Batch Size?} B -->|序列/Transformer| D[LayerNorm] B -->|目标检测/分割| E[GroupNorm] C -->|≥ 8| F[BatchNorm] C -->|< 8| G{多 GPU?} G -->|是| H[SyncBatchNorm] G -->|否| E

参考论文: 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 + MomentumCNN收敛慢但泛化好
Adam通用自适应学习率
AdamWTransformer解耦权重衰减
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 AttentionIO 优化

八、特殊训练技术

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 不下降

排查顺序:

  1. 学习率太大/太小 → 调到 1e-4 试试
  2. 数据问题 → 检查标签是否对齐
  3. 梯度消失/爆炸 → 加 BatchNorm、梯度裁剪
  4. 损失函数错误 → 检查是否匹配任务

坑二:训练集好测试集差(过拟合)

解决:

  • 加 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,
}

十四、写在最后

训练优化没有银弹,每个任务都需要根据数据、模型、资源来调整。

几条核心原则:

  1. 从简单开始:SGD + 固定学习率 → 逐步加复杂策略
  2. 监控是基础:loss 曲线、梯度范数、学习率
  3. 超参数敏感:学习率 > batch size > 其他
  4. 小数据用强正则:Dropout、数据增强、早停
  5. 大数据用弱正则:让模型充分学习
  6. 混合精度是标配:免费的速度和显存提升
  7. 断点续训必须有:训练中断不丢失进度

训练是一门经验科学,多实验、多记录、多对比。论文里的最佳实践不一定适合你的数据,要自己跑实验验证。


本文整合了 20+ 篇模型训练优化相关文章,涵盖损失函数、学习率调度、正则化、梯度优化、混合精度、分布式训练、注意力机制、自监督学习、对抗训练、知识蒸馏、多任务学习、课程学习、主动学习、图神经网络、Word2Vec、混合专家等核心技术。

版权声明: 本文首发于 指尖魔法屋-AI模型训练优化实战指南:从损失函数到分布式训练https://blog.thinkmoon.cn/post/ai-training-optimization-comprehensive-guide/) 转载或引用必须申明原指尖魔法屋来源及源地址!