Python-pytorch-训练流程

PyTorch 训练流程

🎯 训练循环全貌

PyTorch 的训练流程虽需手写,但极其灵活。以下是标准模板:

python 复制代码
import torch
import torch.nn as nn
import torch.optim as optim

# ========== 准备 ==========
model = MyModel().to(device)
loss_fn = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100)

# ========== 训练循环 ==========
for epoch in range(num_epochs):
    # 训练阶段
    model.train()
    train_loss = 0.0
    for X, y in train_loader:
        X, y = X.to(device), y.to(device)

        optimizer.zero_grad()          # 1. 清零梯度
        pred = model(X)                # 2. 前向传播
        loss = loss_fn(pred, y)        # 3. 计算损失
        loss.backward()                # 4. 反向传播
        optimizer.step()               # 5. 更新参数

        train_loss += loss.item() * X.size(0)

    scheduler.step()                   # 6. 更新学习率

    # 验证阶段
    model.eval()
    correct, total = 0, 0
    with torch.no_grad():
        for X, y in val_loader:
            X, y = X.to(device), y.to(device)
            pred = model(X)
            _, predicted = pred.max(1)
            correct += (predicted == y).sum().item()
            total += y.size(0)

    print(f'Epoch {epoch+1}: '
          f'Train Loss: {train_loss/total:.4f}, '
          f'Val Acc: {100*correct/total:.2f}%')

⚡ 优化器(Optimizer)

SGD 家族

python 复制代码
# 基础 SGD
optimizer = optim.SGD(model.parameters(), lr=0.01)

# SGD + 动量
optimizer = optim.SGD(model.parameters(), lr=0.01, momentum=0.9)

# SGD + 动量 + 权重衰减(L2 正则)
optimizer = optim.SGD(
    model.parameters(),
    lr=0.01,
    momentum=0.9,
    weight_decay=5e-4,         # L2 正则化系数
    nesterov=True              # Nesterov 加速梯度
)

Adam 家族 ⭐

python 复制代码
# Adam(最常用)
optimizer = optim.Adam(model.parameters(), lr=0.001)

# Adam 完整参数
optimizer = optim.Adam(
    model.parameters(),
    lr=1e-3,                   # 学习率
    betas=(0.9, 0.999),        # 一阶/二阶动量系数
    eps=1e-8,                  # 数值稳定性
    weight_decay=0.01,         # 解耦权重衰减(AdamW 推荐)
    amsgrad=False
)

# AdamW(权重衰减更正确的版本,推荐)
optimizer = optim.AdamW(model.parameters(), lr=0.001, weight_decay=0.01)

# Adamax(稀疏梯度场景)
optimizer = optim.Adamax(model.parameters(), lr=0.002)

其他优化器

python 复制代码
# RMSprop(RNN / 强化学习用)
optimizer = optim.RMSprop(model.parameters(), lr=0.001, alpha=0.99)

# Adagrad(稀疏特征)
optimizer = optim.Adagrad(model.parameters(), lr=0.01)

# LAMB(大批量训练)
# from torch_optimizer import Lamb

优化器选择建议

场景 推荐
CV 分类/检测 SGD(momentum=0.9)AdamW
NLP / Transformer AdamW
GAN Adam (betas=(0.5, 0.999))
强化学习 RMSpropAdam
快速实验 Adam

📐 学习率调度器(LR Scheduler)

常用调度器

python 复制代码
import torch.optim.lr_scheduler as lr_scheduler

# ===== 衰减型 =====
# StepLR: 每 step_size 个 epoch 乘 gamma
scheduler = lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)

# MultiStepLR: 在特定 epoch 衰减
scheduler = lr_scheduler.MultiStepLR(optimizer, milestones=[30, 80], gamma=0.1)

# ExponentialLR: 指数衰减
scheduler = lr_scheduler.ExponentialLR(optimizer, gamma=0.99)

# ===== 余弦退火(推荐)=====
# 余弦退火到 0
scheduler = lr_scheduler.CosineAnnealingLR(optimizer, T_max=100)

# 余弦退火 + 热重启
scheduler = lr_scheduler.CosineAnnealingWarmRestarts(
    optimizer, T_0=10, T_mult=2, eta_min=1e-6
)

# ===== 自适应型 =====
# ReduceLROnPlateau: 监控指标停滞时衰减
scheduler = lr_scheduler.ReduceLROnPlateau(
    optimizer, mode='min', factor=0.1, patience=10, min_lr=1e-6
)

# ===== 预热型 =====
# OneCycleLR: 先升后降
scheduler = lr_scheduler.OneCycleLR(
    optimizer, max_lr=0.01, steps_per_epoch=len(train_loader), epochs=100
)

# LinearLR: 线性衰减
scheduler = lr_scheduler.LinearLR(optimizer, start_factor=1.0, end_factor=0.0)

# ===== LambdaLR: 自定义 =====
scheduler = lr_scheduler.LambdaLR(optimizer, lr_lambda=lambda epoch: 0.95 ** epoch)

使用方式

python 复制代码
for epoch in range(num_epochs):
    train_one_epoch(model, optimizer, train_loader)
    val_loss = validate(model, val_loader)

    # 大多数 scheduler 在每个 epoch 后 step
    scheduler.step()

    # ReduceLROnPlateau 需要传入监控指标
    # scheduler.step(val_loss)

    current_lr = scheduler.get_last_lr()[0]
    print(f'LR: {current_lr:.6f}')

学习率预热(Warmup)

python 复制代码
# 手动实现线性预热
def warmup_lr(optimizer, step, warmup_steps, target_lr):
    """在前 warmup_steps 步中线性增长到 target_lr"""
    if step < warmup_steps:
        lr = target_lr * (step + 1) / warmup_steps
        for param_group in optimizer.param_groups:
            param_group['lr'] = lr
    return optimizer

# 配合 CosineAnnealingLR 使用
for epoch in range(num_epochs):
    for batch_idx, (X, y) in enumerate(train_loader):
        global_step = epoch * len(train_loader) + batch_idx
        if global_step < 1000:
            warmup_lr(optimizer, global_step, 1000, 0.001)
        # ... training step ...
    scheduler.step()   # epoch 级调度

📊 训练技巧与最佳实践

1. 梯度裁剪

防止梯度爆炸,尤其对 RNN/LSTM 至关重要。

python 复制代码
# 按范数裁剪(推荐)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

# 按值裁剪
torch.nn.utils.clip_grad_value_(model.parameters(), clip_value=0.5)

# 在训练循环中的位置
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)  # ← 在 step 前
optimizer.step()

2. 参数分组 --- 不同层不同学习率

python 复制代码
# 骨干网络用更小的学习率,分类头用更大的
optimizer = optim.SGD([
    {'params': model.backbone.parameters(), 'lr': 1e-4},  # 小学习率
    {'params': model.classifier.parameters(), 'lr': 1e-3}, # 大学习率
], momentum=0.9, weight_decay=5e-4)

3. 权重初始化

python 复制代码
def init_weights(m):
    if isinstance(m, nn.Linear):
        nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
        if m.bias is not None:
            nn.init.constant_(m.bias, 0)
    elif isinstance(m, nn.Conv2d):
        nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
    elif isinstance(m, nn.BatchNorm2d):
        nn.init.constant_(m.weight, 1)
        nn.init.constant_(m.bias, 0)

model.apply(init_weights)

4. 完整的训练模板

python 复制代码
def train_epoch(model, loader, loss_fn, optimizer, device):
    model.train()
    running_loss = 0.0
    correct, total = 0, 0

    for X, y in loader:
        X, y = X.to(device), y.to(device)

        optimizer.zero_grad()
        outputs = model(X)
        loss = loss_fn(outputs, y)
        loss.backward()
        optimizer.step()

        running_loss += loss.item() * X.size(0)
        _, preds = outputs.max(1)
        correct += (preds == y).sum().item()
        total += y.size(0)

    return running_loss / total, correct / total


@torch.no_grad()
def validate(model, loader, loss_fn, device):
    model.eval()
    running_loss = 0.0
    correct, total = 0, 0

    for X, y in loader:
        X, y = X.to(device), y.to(device)
        outputs = model(X)
        loss = loss_fn(outputs, y)

        running_loss += loss.item() * X.size(0)
        _, preds = outputs.max(1)
        correct += (preds == y).sum().item()
        total += y.size(0)

    return running_loss / total, correct / total


# 主训练
def main():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    model = MyModel().to(device)
    loss_fn = nn.CrossEntropyLoss()
    optimizer = optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100)

    best_acc = 0.0
    for epoch in range(100):
        train_loss, train_acc = train_epoch(model, train_loader, loss_fn, optimizer, device)
        val_loss, val_acc = validate(model, val_loader, loss_fn, device)
        scheduler.step()

        # 保存最佳模型
        if val_acc > best_acc:
            best_acc = val_acc
            torch.save(model.state_dict(), 'best_model.pth')

        print(f'Epoch {epoch+1:3d}: '
              f'Train Loss {train_loss:.4f}, Train Acc {train_acc:.3f} | '
              f'Val Loss {val_loss:.4f}, Val Acc {val_acc:.3f}')

    print(f'Best Val Acc: {best_acc:.4f}')

🔁 自定义训练逻辑

对抗训练(GAN)

python 复制代码
# GAN 交替训练
for epoch in range(num_epochs):
    for real_imgs in loader:
        # ----- 训练判别器 -----
        optimizer_D.zero_grad()
        real_pred = discriminator(real_imgs)
        fake_imgs = generator(torch.randn(batch, latent_dim))
        fake_pred = discriminator(fake_imgs.detach())
        d_loss = -torch.mean(torch.log(real_pred) + torch.log(1 - fake_pred))
        d_loss.backward()
        optimizer_D.step()

        # ----- 训练生成器 -----
        optimizer_G.zero_grad()
        fake_imgs = generator(torch.randn(batch, latent_dim))
        fake_pred = discriminator(fake_imgs)
        g_loss = -torch.mean(torch.log(fake_pred))
        g_loss.backward()
        optimizer_G.step()

梯度累积(模拟大 Batch)

python 复制代码
accumulation_steps = 4  # 梯度累积步数
optimizer.zero_grad()

for i, (X, y) in enumerate(train_loader):
    loss = loss_fn(model(X), y)
    loss = loss / accumulation_steps   # 归一化
    loss.backward()                     # 梯度累积

    if (i + 1) % accumulation_steps == 0:
        optimizer.step()
        optimizer.zero_grad()

📝 Metrics 计算

python 复制代码
# 分类指标
from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score

all_preds, all_labels = [], []
with torch.no_grad():
    for X, y in test_loader:
        outputs = model(X)
        _, preds = outputs.max(1)
        all_preds.extend(preds.cpu().numpy())
        all_labels.extend(y.cpu().numpy())

print(f'Accuracy:  {accuracy_score(all_labels, all_preds):.4f}')
print(f'Precision: {precision_score(all_labels, all_preds, average="macro"):.4f}')
print(f'Recall:    {recall_score(all_labels, all_preds, average="macro"):.4f}')
print(f'F1:        {f1_score(all_labels, all_preds, average="macro"):.4f}')

📝 速查表

需求 代码
SGD optim.SGD(params, lr, momentum=0.9, weight_decay=5e-4)
Adam optim.Adam(params, lr=1e-3)
AdamW optim.AdamW(params, lr=1e-3, weight_decay=0.01)
清零梯度 optimizer.zero_grad()
更新参数 optimizer.step()
梯度裁剪 nn.utils.clip_grad_norm_(params, max_norm=1.0)
StepLR lr_scheduler.StepLR(opt, step_size, gamma)
余弦退火 lr_scheduler.CosineAnnealingLR(opt, T_max)
Plateau衰减 lr_scheduler.ReduceLROnPlateau(opt, mode='min')
训练模式 model.train()
评估模式 model.eval()
禁用梯度(推理) with torch.no_grad():

\[pytorch-总览\|← 返回总览\]

相关推荐
vivo互联网技术1 小时前
Octopus:基于无历史数据的梯度正交化的学习框架|CVPR 2026
深度学习·计算机视觉·llm
Zane19941 小时前
daemon 线程说没就没?一文讲透 threading 的适用场景与线程安全
后端·python
up up day1 小时前
Python enchant 模块使用教程
python
TMT星球1 小时前
快手Q2总收入355亿元:月活近8亿,核心商业收入同比增长7.4%
大数据·人工智能
zoujiahui_20181 小时前
Python 包与环境管理工具 uv
开发语言·python·uv
霸道流氓气质1 小时前
ima.copilot-AI知识库-完整使用手册与最新动态
人工智能·copilot
牧羊人.3331 小时前
计算机视觉基础 第 9 章|实战:银行卡号识别
图像处理·人工智能·opencv·计算机视觉·图搜索算法
吴声子夜歌1 小时前
Java面试题——JVM(二)
java·开发语言·jvm
CRMEB1 小时前
商城大促活动策划全流程:从定目标到复盘四阶段
java·开发语言·人工智能·ai·开源·php
denglinqings1 小时前
Java开发从入门到精通所有课程
java·开发语言