Skip to content

PyTorch 速查表

由 Meta 开发的开源深度学习框架,提供动态计算图与自动求导。

01

入门

安装与张量

通过 pip/conda 安装 PyTorch,建议按官网选择匹配 CUDA 版本的命令。torch.tensor() 总是复制数据,而 torch.as_tensor() 复用内存(更快)。张量类似 NumPy ndarray,但可在 GPU 上运算并支持自动求导。

pytorch
# install PyTorch (CUDA 11.8)
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118

import torch

# create a tensor
x = torch.tensor([[1, 2], [3, 4]])
y = torch.zeros(3, 3)
z = torch.randn(2, 3)  # standard normal

# tensor on GPU
device = "cuda" if torch.cuda.is_available() else "cpu"
x = x.to(device)

张量属性

dtype 决定数值精度(默认 float32),device 决定 CPU/GPU 位置。requires_grad 启用自动求导。ndim 是维度数,shape 返回各维度大小。tensor.item() 仅适用于单元素张量,会触发同步。

pytorch
x = torch.randn(3, 4, 5)
print(x.shape)      # torch.Size([3, 4, 5])
print(x.dtype)      # torch.float32
print(x.device)     # cpu or cuda:0
print(x.requires_grad)  # False

# specify dtype at creation
x = torch.zeros(3, dtype=torch.int64)
x = torch.tensor([1.0, 2.0], dtype=torch.float16)

# check number of elements
print(x.numel())    # 3

张量创建

zeros/ones/empty 创建基础张量;arange/linspace 生成等差序列。*_like 函数沿用输入张量的形状与 dtype。randn 生成标准正态分布,rand 生成 [0,1) 均匀分布。empty 不初始化内存,可能含任意值。

pytorch
# from Python lists
a = torch.tensor([1, 2, 3])

# special tensors
zeros = torch.zeros(2, 3)
ones = torch.ones(2, 3)
eye = torch.eye(3)            # identity matrix
full = torch.full((2, 3), 7.0)

# ranges
arange = torch.arange(0, 10, 2)   # [0, 2, 4, 6, 8]
linspace = torch.linspace(0, 1, 5)  # [0, 0.25, 0.5, 0.75, 1.0]

# random
rand = torch.rand(2, 3)       # uniform [0, 1)
randn = torch.randn(2, 3)     # normal N(0, 1)
randint = torch.randint(0, 10, (3, 3))

类型转换

.to(dtype) 是推荐的类型转换方式,也可用 .float()/.long()/.bool() 等便捷方法。整型张量做除法用 //(地板除)或先转 float。device 也可用 .to('cuda') 同时转换。注意 dtype 不匹配会触发自动类型提升。

pytorch
x = torch.tensor([1.5, 2.5, 3.5])

# change dtype
y = x.to(torch.int32)    # tensor([1, 2, 3])
y = x.int()              # shortcut
y = x.float()
y = x.bool()
y = x.long()             # int64

# to numpy and back
arr = x.numpy()          # shares memory (CPU)
t = torch.from_numpy(arr)

# 0-d tensor to Python scalar
val = x[0].item()        # 1.5

索引与切片

索引语义与 NumPy 一致。负索引从末尾开始。Advanced indexing(用张量/列表索引)会创建副本而非视图。使用 [...] 保留其余维度,用 : 选择整维。布尔掩码索引常用于过滤样本。

pytorch
x = torch.arange(12).reshape(3, 4)
# tensor([[ 0,  1,  2,  3],
#         [ 4,  5,  6,  7],
#         [ 8,  9, 10, 11]])

x[0]          # first row
x[1, 2]       # element at (1, 2) -> 6
x[:, 1]       # second column
x[0:2]        # first two rows
x[:, ::2]     # every other column

# boolean mask
mask = x > 5
x[mask]       # tensor([6, 7, 8, 9, 10, 11])

# index with LongTensor
idx = torch.tensor([0, 2])
x[idx]        # rows 0 and 2

变形与视图

view 要求内存连续,否则需先 contiguous()。reshape 自动处理不连续情况。permute 调换维度顺序。squeeze/remove 移除大小为 1 的维度,unsqueeze 在指定位置插入维度。view 不复制数据,与原张量共享内存。

pytorch
x = torch.arange(12)

# reshape (returns view when possible)
y = x.reshape(3, 4)
y = x.view(3, 4)      # only for contiguous tensors

# add/remove dimensions
y = x.unsqueeze(0)    # shape (1, 12)
y = x.squeeze()       # remove size-1 dims

# transpose / permute
a = torch.randn(3, 4, 5)
b = a.permute(2, 0, 1)  # shape (5, 3, 4)
b = a.transpose(0, 1)   # swap dims 0 and 1

# flatten
flat = a.flatten()      # shape (60,)
flat = a.reshape(-1)    # equivalent
02

张量运算

逐元素运算

运算符 +、-、*、/ 是逐元素的。** 是幂运算,@ 是矩阵乘法。torch.clamp 限制数值范围。原地运算(带 _ 后缀,如 add_)节省内存但会破坏自动求导,慎用。逐元素乘法 * 与矩阵乘法 @ 完全不同。

pytorch
a = torch.tensor([1.0, 2.0, 3.0])
b = torch.tensor([4.0, 5.0, 6.0])

# element-wise ops
a + b           # tensor([5., 7., 9.])
a * b           # tensor([ 4., 10., 18.])
a / b           # element-wise division
a ** 2          # tensor([1., 4., 9.])

# in-place (suffix _)
a.add_(1)       # a = a + 1, modifies in place
a.mul_(2)       # a = a * 2

# scalar ops
a + 10
b * 0.5

归约运算

sum/mean/max/min 沿指定维度归约。keepdim=True 保留该维度为 1,便于广播。argmax 返回最大值索引(分类任务取预测类别常用)。mean 默认对整个张量求平均,注意整型需先转 float。

pytorch
x = torch.tensor([-1.0, 0.0, 1.0, 2.0])

torch.abs(x)        # absolute value
torch.sqrt(torch.abs(x))
torch.exp(x)        # e^x
torch.log(x.abs() + 1e-8)
torch.sin(x)
torch.clamp(x, -0.5, 1.5)   # clip values

# activation functions
torch.sigmoid(x)    # 1 / (1 + e^-x)
torch.tanh(x)
torch.relu(x)       # max(0, x)

# rounding
y = torch.tensor([1.4, 1.5, 2.6])
torch.round(y)      # tensor([1., 2., 3.])
torch.floor(y)
torch.ceil(y)

矩阵运算

@ 或 torch.matmul 支持批量矩阵乘法(最后两维)。mm/bmm 仅限 2D/3D。einsum 用爱因斯坦求和约定表达复杂运算(如转置、双线性、注意力分数),可读性强。torch.linalg.inv/svd/qr 提供线性代数运算。

pytorch
x = torch.arange(12).reshape(3, 4).float()

x.sum()             # 66.0 (scalar tensor)
x.sum(dim=0)        # sum over rows -> shape (4,)
x.sum(dim=1)        # sum over columns -> shape (3,)

x.mean(dim=1)
x.max()             # returns value tensor
x.max(dim=1)        # returns (values, indices)
x.argmin(dim=1)     # indices of min along dim

# keep dimensions
x.sum(dim=1, keepdim=True)  # shape (3, 1)

# norm
x.norm()            # Frobenius norm
x.norm(dim=1)       # per-row norm

广播机制

广播规则与 NumPy 一致:从末尾维度对齐,维度为 1 或缺失时扩展。形状 (3,1) 与 (1,4) 广播为 (3,4)。避免意外的广播导致内存膨胀或错误结果。einsum 是显式维度的安全替代方案。

pytorch
a = torch.tensor([1, 2, 3, 4])
b = torch.tensor([3, 2, 1, 4])

a == b              # tensor([False, True, False, True])
a > b               # tensor([False, False, True, False])
a != b

torch.equal(a, b)   # False (single bool)

# where: pick from two tensors
torch.where(a > b, a, b)  # element-wise max

# boolean helpers
(a > 2).any()       # True
(a > 2).all()        # False
(a == b).sum()       # 2 (count of True)

# top-k
torch.topk(a, 2)     # returns top 2 values + indices

拼接与堆叠

cat 沿现有维度拼接,维度不变;stack 沿新维度堆叠,维度数 +1。cat 要求非拼接维度形状一致。stack 要求所有张量形状完全相同。split/chunk 是其逆操作。

pytorch
a = torch.randn(2, 3)
b = torch.randn(3, 4)

# matrix multiplication
c = a @ b              # shape (2, 4)
c = torch.matmul(a, b) # equivalent

# batched matmul
A = torch.randn(10, 2, 3)
B = torch.randn(10, 3, 4)
C = A @ B              # shape (10, 2, 4)

# element-wise product
e = a * a              # shape (2, 3)

# other ops
torch.inverse(a[:, :3])   # inverse of square slice
torch.det(a[:, :3])
torch.svd(a)               # singular value decomposition
torch.einsum('bi,i->b', torch.randn(5,3), torch.randn(3))  # batched dot

比较运算

比较运算返回布尔张量。torch.where(cond, x, y) 按条件选择元素。topk 返回前 k 大的值与索引。isnan/isinf 用于检测数值异常。eq/equal 做精确比较,allclose 做带容差的近似比较。

pytorch
a = torch.randn(2, 3)
b = torch.randn(2, 3)

# concat along existing dim
c = torch.cat([a, b], dim=0)   # shape (4, 3)
c = torch.cat([a, b], dim=1)   # shape (2, 6)

# stack along new dim
s = torch.stack([a, b], dim=0) # shape (2, 2, 3)

# split / chunk
x = torch.arange(6).reshape(2, 3)
parts = torch.split(x, 2, dim=0)   # split into chunks of size 2
chunks = torch.chunk(x, 3, dim=1)  # split into 3 chunks

# repeat / expand
r = a.repeat(2, 3)        # repeat whole tensor
e = a.unsqueeze(0).expand(4, 2, 3)  # broadcast without copy
03

自动求导

自动求导基础

requires_grad=True 启用梯度追踪。backward() 从标量损失反向传播,自动填充 .grad 属性。叶子张量(用户创建)才有 .grad,中间结果默认释放。grad_fn 指向生成该张量的反向函数。

pytorch
import torch

# tensors that require gradient tracking
x = torch.tensor([2.0], requires_grad=True)
y = torch.tensor([3.0], requires_grad=True)

# build a computation graph
z = x * y + x ** 2     # z = xy + x^2
print(z.grad_fn)        # <AddBackward0 object>

# backprop
z.backward()

# dz/dx = y + 2x = 3 + 4 = 7
print(x.grad)   # tensor([7.])
# dz/dy = x = 2
print(y.grad)   # tensor([2.])

梯度计算

对非标量张量调用 backward 需传入 gradient 参数(与该张量同形状的权重)。更常见的做法是对标量 loss 直接 backward()。.grad 累加而非覆盖,每次反向前应 optimizer.zero_grad() 或 model.zero_grad()。

pytorch
import torch

x = torch.linspace(-3, 3, steps=10, requires_grad=True)
y = torch.sin(x)

# sum needed for backward on non-scalar
y.sum().backward()
print(x.grad)   # cos(x)

# gradient w.r.t. intermediate tensor
a = torch.tensor(1.0, requires_grad=True)
b = a * 2
c = b ** 2
grads = torch.autograd.grad(c, a)
print(grads)    # (tensor(8.),)  dc/da = 4b * 2 = 8

计算图

PyTorch 默认动态图,每次前向都构建新图。retain_graph=True 保留图供多次反向(如双重反向求 Hessian)。 backward 后图默认释放以节省内存。retain_grad() 让非叶子张量也保存 .grad。

pytorch
x = torch.tensor([1.0], requires_grad=True)

# option 1: context manager (preferred)
with torch.no_grad():
    y = x * 2          # y.requires_grad == False

# option 2: decorator
@torch.no_grad()
def inference(x):
    return model(x)

# option 3: detach from graph
y = (x * 2).detach()   # y is a new tensor, no grad

# enable grad inside no_grad
with torch.no_grad():
    with torch.enable_grad():
        z = x * 3      # z requires grad

禁用梯度

推理/验证时用 with torch.no_grad() 或 model.eval() 关闭梯度计算,显著降低内存与时间开销。@torch.no_grad() 装饰器形式用于函数。torch.enable_grad() 在 no_grad 上下文中重新启用。inference_mode 比 no_grad 更快但限制更多。

pytorch
x = torch.tensor(0.5, requires_grad=True)

# first derivative
y = torch.sin(x)
y.backward(create_graph=True)
print(x.grad)   # cos(0.5)

# second derivative (gradient of gradient)
x.grad.zero_()
g = torch.autograd.grad(torch.sin(x), x, create_graph=True)[0]
g2 = torch.autograd.grad(g, x)[0]
print(g2)       # -sin(0.5)

# practical use: penalize gradient (e.g. gradient penalty)
loss = g.norm()

自定义 Autograd 函数

继承 Function 并实现 forward 和 backward 静态方法。backward 接收上游梯度,返回各输入的梯度。ctx.save_for_backward 保存前向张量供反向使用。用于实现不可导为普通算子的自定义层(如分段线性、稀疏操作)。

pytorch
x = torch.tensor([1.0, 2.0], requires_grad=True)
y = (x ** 2).sum()

# full backward hook on a tensor
def print_grad(grad):
    print("grad:", grad)
    return grad * 2  # can modify gradient

y.register_hook(print_grad)
y.backward()
print(x.grad)   # 2x but doubled by hook

# module hook
def hook_fn(module, grad_input, grad_output):
    print(module.__class__.__name__, grad_output)

model.layer.register_full_backward_hook(hook_fn)

梯度钩子

register_hook 在反向传播时触发回调,可查看或修改梯度(返回新梯度即替换)。常用于梯度分析、调试 NaN、实现梯度惩罚(WGAN-GP)。张量钩子在第一次反向后自动移除,模块钩子持久。register_full_backward_hook 更精确。

pytorch
x = torch.tensor(2.0, requires_grad=True)
w = torch.tensor(3.0, requires_grad=True)
b = torch.tensor(1.0, requires_grad=True)

y = w * x + b
print(y.is_leaf)        # False
print(x.is_leaf)        # True

# retain graph for multiple backward passes
loss = (y - 5) ** 2
loss.backward(retain_graph=True)
# can call backward again
loss.backward()

# double backward
loss.backward(create_graph=True)
print(x.grad)
04

神经网络模块

nn.Module 基础

所有自定义模型继承 nn.Module。在 __init__ 中定义子模块和参数,在 forward 中定义计算流。model.parameters() 迭代所有可学习参数,model.to(device) 递归迁移。模块可嵌套形成计算树。

pytorch
from torch.utils.data import Dataset, DataLoader

# built-in datasets
from torchvision import datasets
mnist = datasets.MNIST("./data", train=True, download=True)

# wrap in DataLoader
loader = DataLoader(
    mnist,
    batch_size=64,
    shuffle=True,
    num_workers=4,
    pin_memory=True,
)

for x, y in loader:
    print(x.shape, y.shape)   # (64, 1, 28, 28), (64,)
    break

常用层

Linear 做仿射变换(in_features → out_features)。Conv2d 用 in/out_channels 控制通道数。batch_first=True 让 RNN 输入为 (batch, seq, feature)。Embedding 将整数索引映射为稠密向量,权重即查找表。

pytorch
from torch.utils.data import DataLoader, BatchSampler, RandomSampler

dataset = list(range(100))

sampler = RandomSampler(dataset)
batch_sampler = BatchSampler(sampler, batch_size=10, drop_last=True)

loader = DataLoader(dataset, batch_sampler=batch_sampler)

# custom sampler
class EvenSampler:
    def __init__(self, data): self.data = data
    def __iter__(self): return iter(range(0, len(self.data), 2))
    def __len__(self): return len(self.data) // 2

# weighted sampling for imbalanced classes
from torch.utils.data import WeightedRandomSampler
weights = [0.1] * 50 + [1.0] * 50
sampler = WeightedRandomSampler(weights, num_samples=100)

激活函数

ReLU 是 CNN 的默认选择,LeakyReLU 防止神经元死亡。GELU 在 Transformer 中流行。Sigmoid 输出 [0,1],Tanh 输出 [-1,1]。现代网络多用 inplace=True 节省内存,但慎用于需要原始输入的 autograd 场景。

pytorch
from torch.utils.data import DataLoader, random_split

# train/val split
train, val = random_split(dataset, [80, 20])

train_loader = DataLoader(train, batch_size=16, shuffle=True)
val_loader = DataLoader(val, batch_size=32, shuffle=False)

# reproducible split
from torch.utils.data import Subset
import numpy as np
gen = torch.Generator().manual_seed(42)
idx = torch.randperm(len(dataset), generator=gen).tolist()
train_idx, val_idx = idx[:80], idx[80:]
train, val = Subset(dataset, train_idx), Subset(dataset, val_idx)

Sequential API

Sequential 按顺序串联模块,前向自动依次调用。OrderedDict 可为每层命名便于访问(model.features,model.classifier)。适合线性堆叠的简单模型,复杂分支结构仍需自定义 forward。

pytorch
from torch.utils.data import default_collate

# variable-length sequences: pad in collate
def pad_collate(batch):
    # batch is list of (tensor, label)
    seqs, labels = zip(*batch)
    lens = torch.tensor([len(s) for s in seqs])
    max_len = lens.max()
    padded = torch.zeros(len(seqs), max_len)
    for i, s in enumerate(seqs):
        padded[i, :len(s)] = s
    return padded, lens, torch.tensor(labels)

loader = DataLoader(dataset, batch_size=4, collate_fn=pad_collate)

# stack batch as dict
def dict_collate(batch):
    return default_collate(batch)

参数与缓冲区

nn.Parameter 是自动注册的 Tensor,requires_grad 默认 True。register_buffer 注册非学习但需随模型迁移的张量(如 BN 的 running_mean)。state_dict() 同时返回参数与缓冲区,便于保存/加载。

pytorch
loader = DataLoader(
    dataset,
    batch_size=128,
    num_workers=8,       # one process per worker
    pin_memory=True,     # page-locked memory for fast H2D copy
    persistent_workers=True,  # keep workers alive across epochs
    prefetch_factor=4,   # batches prefetched per worker
)

# avoid CPU bottleneck
import os
os.cpu_count()  # check available cores

# in Jupyter / Windows: set start method
import torch.multiprocessing as mp
# mp.set_start_method('spawn')  # uncomment if needed
05

损失函数

MSE 与 L1 损失

MSELoss 用于回归,对大误差更敏感(平方放大)。L1Loss 对异常值更鲁棒。SmoothL1Loss(Huber)在误差小时用平方、大时用线性,兼顾两者。reduction 默认为 'mean',求和用 'sum',不归约用 'none'。

pytorch
import torch.nn as nn

class MLP(nn.Module):
    def __init__(self, in_dim, hidden, out_dim):
        super().__init__()
        self.fc1 = nn.Linear(in_dim, hidden)
        self.fc2 = nn.Linear(hidden, hidden)
        self.fc3 = nn.Linear(hidden, out_dim)
        self.act = nn.ReLU()

    def forward(self, x):
        x = self.act(self.fc1(x))
        x = self.act(self.fc2(x))
        return self.fc3(x)   # logits

model = MLP(784, 128, 10)
print(model)   # prints structure

交叉熵损失

CrossEntropyLoss 内部已包含 LogSoftmax,不要在模型最后再加 Softmax!输入应为未归一化的 logits。label_smoothing>0 防止过拟合。忽略某些类别用 ignore_index(如 padding)。多标签分类应用 BCEWithLogitsLoss。

pytorch
import torch.nn as nn

linear = nn.Linear(in_features=20, out_features=10, bias=True)
print(linear.weight.shape)   # (10, 20)
print(linear.bias.shape)     # (10,)

# common activations
nn.ReLU()         # max(0, x)
nn.LeakyReLU(0.01)
nn.GELU()         # smooth, used in transformers
nn.Sigmoid()      # (0, 1)
nn.Tanh()         # (-1, 1)
nn.Softmax(dim=1)
nn.ELU()

# functional API (no parameters)
import torch.nn.functional as F
out = F.relu(linear(x))   # same as nn.ReLU()(linear(x))

二元交叉熵损失

BCEWithLogitsLoss 内置 Sigmoid,数值上比 Sigmoid+BCELoss 更稳定(避免 log(0))。pos_weight 处理类别不平衡(如正样本稀少)。多标签分类(每样本多类别独立为 0/1)用此损失,每维独立判断。

pytorch
model = MLP(784, 128, 10)

# iterate parameters
for name, p in model.named_parameters():
    print(name, p.shape, p.requires_grad)

# only trainable params
trainable = filter(lambda p: p.requires_grad, model.parameters())

# count parameters
n_params = sum(p.numel() for p in model.parameters())
print(f"{n_params:,} parameters")

# access submodules by name
model.fc1         # the Linear layer
model['fc1']      # equivalent if using nn.ModuleDict

# children vs modules
list(model.children())    # direct children
list(model.modules())     # recursively all sub-modules

负对数似然损失

NLLLoss 要求模型最后输出 LogSoftmax。等价于 CrossEntropyLoss 但分开两步,便于自定义 log-softmax。ignore_index 可跳过不参与损失的样本(如 padding token)。旧代码中常见,新代码建议直接用 CrossEntropyLoss。

pytorch
import torch.nn as nn

# simple stack
model = nn.Sequential(
    nn.Linear(784, 256),
    nn.ReLU(),
    nn.Dropout(0.5),
    nn.Linear(256, 10),
)

# named layers via OrderedDict
from collections import OrderedDict
model = nn.Sequential(OrderedDict([
    ('fc1', nn.Linear(784, 256)),
    ('relu', nn.ReLU()),
    ('dropout', nn.Dropout(0.5)),
    ('fc2', nn.Linear(256, 10)),
]))
print(model.fc1)   # access by name

# append at runtime
model.append(nn.Softmax(dim=1))

自定义损失

损失函数也继承 nn.Module,在 forward 中实现。所有运算用 torch 函数以保证自动求导可反向。inplace 操作可能破坏梯度。可学习损失(带参数)也用 nn.Parameter 注册。

pytorch
model = MLP(784, 128, 10)

# move to device / dtype
model = model.to('cuda')
model = model.to(torch.float16)

# train vs eval mode
model.train()    # enable dropout, update BN stats
model.eval()     # disable dropout, freeze BN stats

# state dict
sd = model.state_dict()
model.load_state_dict(sd)

# apply a function to all submodules
model.apply(lambda m: print(type(m).__name__))

# set requires_grad on all parameters
for p in model.parameters():
    p.requires_grad_(False)

三元组与对比损失

TripletMarginLoss 用于度量学习,拉近锚点与正样本、推远与负样本。margin 是关键超参。CosineEmbeddingLoss 基于余弦相似度。对比学习(SimCLR/MoCo)常用 InfoNCE 损失,温度参数 τ 控制难度。

pytorch
import torch.nn as nn
import torch.nn.init as init

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(100, 50)
        self.fc2 = nn.Linear(50, 10)
        self._init_weights()

    def _init_weights(self):
        init.kaiming_normal_(self.fc1.weight, nonlinearity='relu')
        init.zeros_(self.fc1.bias)
        init.xavier_uniform_(self.fc2.weight)
        init.zeros_(self.fc2.bias)

# apply to existing model
def init_all(m):
    if isinstance(m, nn.Linear):
        init.kaiming_normal_(m.weight)
        if m.bias is not None:
            init.zeros_(m.bias)
model.apply(init_all)
06

优化器

SGD

momentum 加速收敛并减少震荡,典型值 0.9。weight_decay 实现 L2 正则化(等价于梯度加上 λ·w)。Nesterov 在动量更新前先向前看一步,常略优。SGD+Momentum 泛化常优于 Adam,是大模型预训练首选。

pytorch
import torch
import torch.nn as nn

# logits shape (N, C), targets shape (N,) with class indices
logits = torch.randn(4, 3)        # 4 samples, 3 classes
targets = torch.tensor([0, 2, 1, 0])

criterion = nn.CrossEntropyLoss()
loss = criterion(logits, targets)

# equivalently, combine LogSoftmax + NLLLoss
log_softmax = nn.LogSoftmax(dim=1)
nll = nn.NLLLoss()
loss2 = nll(log_softmax(logits), targets)

# weight classes (handle imbalance)
weights = torch.tensor([1.0, 2.0, 1.0])
criterion = nn.CrossEntropyLoss(weight=weights)

# ignore padding index (NLP)
criterion = nn.CrossEntropyLoss(ignore_index=-100)

Adam 与 AdamW

Adam 自适应学习率,收敛快,默认超参对大多任务可用。AdamW 修正了权重衰减的实现(与 L2 正则不同),是 Transformer 微调的标配。amsgrad 防止学习率单调不降导致的发散。betas 控制一二阶矩的平滑系数。

pytorch
import torch.nn as nn

pred = torch.randn(4, 1)
target = torch.randn(4, 1)

# mean squared error (L2)
mse = nn.MSELoss()
loss = mse(pred, target)        # mean over all elements

# L1 (mean absolute error)
l1 = nn.L1Loss()
loss = l1(pred, target)

# smooth L1 (Huber) — robust to outliers
smooth = nn.SmoothL1Loss(beta=0.1)

# reduction options
mse_sum = nn.MSELoss(reduction='sum')
mse_none = nn.MSELoss(reduction='none')   # per-element

学习率调度器

StepLR 按固定步长衰减。CosineAnnealingLR 平滑衰减到 eta_min。OneCycleLR 配合 warmup 实现超级收敛(先升后降)。ReduceLROnPlateau 根据指标停滞自动降学习率。调用 scheduler.step() 应在 optimizer.step() 之后。

pytorch
import torch.nn as nn

# binary classification, multi-label, or sigmoid outputs
logits = torch.randn(4, 1)
target = torch.tensor([[1.0], [0.0], [1.0], [0.0]])

# preferred: combines sigmoid + BCE, numerically stable
criterion = nn.BCEWithLogitsLoss()
loss = criterion(logits, target)

# BCELoss expects probabilities (apply sigmoid first)
criterion = nn.BCELoss()
probs = torch.sigmoid(logits)
loss = criterion(probs, target)

# positive weight for imbalanced binary tasks
pos_weight = torch.tensor([5.0])
criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)

逐层学习率

不同层用不同学习率常用于微调预训练模型(底层特征通用,学习率小;新分类头大)。param_groups 是列表,每个组可单独设置 lr、weight_decay 等。也可为不同参数组设置不同 weight_decay(如 LayerNorm 和 Bias 不衰减)。

pytorch
import torch
import torch.nn as nn
import torch.nn.functional as F

# as a function
def huber_loss(pred, target, delta=1.0):
    error = pred - target
    abs_err = error.abs()
    quad = torch.where(abs_err <= delta,
                       0.5 * error ** 2,
                       delta * (abs_err - 0.5 * delta))
    return quad.mean()

# as a module (so it can have parameters / state)
class FocalLoss(nn.Module):
    def __init__(self, alpha=0.25, gamma=2.0):
        super().__init__()
        self.alpha = alpha
        self.gamma = gamma

    def forward(self, logits, targets):
        bce = F.binary_cross_entropy_with_logits(logits, targets, reduction='none')
        p = torch.sigmoid(logits)
        pt = p * targets + (1 - p) * (1 - targets)
        loss = self.alpha * (1 - pt) ** self.gamma * bce
        return loss.mean()

criterion = FocalLoss()

梯度裁剪

clip_grad_norm_ 按全局 L2 范数裁剪(更常用),clip_grad_value_ 按元素裁剪到 [-v, v]。防止梯度爆炸导致训练崩溃,RNN/Transformer 必备。在 backward() 之后、optimizer.step() 之前调用。

pytorch
import torch.nn as nn
import torch

criterion = nn.CrossEntropyLoss(reduction='none')
loss_per = criterion(logits, targets)   # shape (N,)
loss = loss_per.mean()

# class-wise mask
mask = targets != -100
loss = criterion(logits[mask], targets[mask])

# multi-task loss
loss_cls = nn.CrossEntropyLoss()(logits, labels)
loss_box = nn.L1Loss()(boxes_pred, boxes_gt)
total = 1.0 * loss_cls + 0.5 * loss_box

# gradient-weighted
total = loss_cls + 0.5 * loss_box
total.backward()

优化器状态

Adam 的 state 保存一二阶矩的滑动平均。load_state_dict 恢复时检查参数形状匹配。state_dict 可单独保存用于断点续训。空 state_dict 表示从零开始。不同优化器的 state 结构不同,不可互换。

pytorch
import torch.nn as nn

# NLLLoss expects log-probabilities (use LogSoftmax first)
log_probs = torch.log_softmax(logits, dim=1)
loss = nn.NLLLoss()(log_probs, targets)

# triplet loss for embeddings
anchor = torch.randn(4, 128)
positive = torch.randn(4, 128)
negative = torch.randn(4, 128)
loss = nn.TripletMarginLoss(margin=1.0)(anchor, positive, negative)

# KL divergence (distributions)
loss = nn.KLDivLoss(reduction='batchmean')(log_probs, target_dist)

# cosine embedding loss
loss = nn.CosineEmbeddingLoss()(x1, x2, torch.tensor([1, -1, 1, -1]))

# CTCLoss for sequence alignment (ASR/OCR)
loss = nn.CTCLoss(blank=0, zero_infinity=True)(log_probs, targets, in_lens, out_lens)
07

数据加载

TensorDataset 与 DataLoader

TensorDataset 将多个张量打包为数据集,要求第一维(样本维)大小一致。DataLoader 处理 batching、shuffling、并行加载。默认 collate_fn 将样本堆叠为 batched 张量。num_workers>0 启用多进程预取。

pytorch
import torch.optim as optim

# basic SGD
optimizer = optim.SGD(model.parameters(), lr=0.01)

# SGD with momentum (classic: 0.9)
optimizer = optim.SGD(model.parameters(), lr=0.01, momentum=0.9)

# Nesterov momentum
optimizer = optim.SGD(
    model.parameters(),
    lr=0.01,
    momentum=0.9,
    nesterov=True,
)

# weight decay (L2 regularization)
optimizer = optim.SGD(model.parameters(), lr=0.01, momentum=0.9, weight_decay=1e-4)

DataLoader 参数

batch_size 控制每批样本数。shuffle=True 每轮重排(训练集常用)。num_workers 多进程加载(注意 Windows 上需 if __name__=='__main__' 保护)。pin_memory=True 加速 H2D 传输。drop_last=True 丢弃不完整批次(影响 BN 统计)。

pytorch
import torch.optim as optim

# Adam — adaptive, default-friendly
optimizer = optim.Adam(model.parameters(), lr=1e-3)

# AdamW — decoupled weight decay (preferred for transformers)
optimizer = optim.AdamW(
    model.parameters(),
    lr=3e-4,
    betas=(0.9, 0.999),
    weight_decay=0.01,
)

# per-parameter group settings
optimizer = optim.Adam([
    {'params': model.encoder.parameters(), 'lr': 1e-4},
    {'params': model.head.parameters(), 'lr': 1e-3},
], lr=1e-3)

自定义 Collate

默认 collate 堆叠等形状样本。变长序列(NLP)需自定义 collate_fn,常用 pad_sequence 填充并对齐。可返回 dict 结构的 batch。collate 是 DataLoader 单线程中的逻辑,避免在其中做重计算以免成为瓶颈。

pytorch
optimizer = optim.Adam(model.parameters(), lr=1e-3)

# standard training step
optimizer.zero_grad()              # clear old grads
loss = criterion(model(x), y)
loss.backward()                    # compute gradients
optimizer.step()                   # update parameters

# set gradients to None (faster, recommended)
optimizer.zero_grad(set_to_none=True)

# manual gradient modification
for p in model.parameters():
    if p.grad is not None:
        p.grad.data.clamp_(-1, 1)  # gradient clipping
optimizer.step()

采样器

RandomSampler 随机采样,WeightedRandomSampler 按权重采样(处理类别不平衡)。DistributedSampler 为 DDP 切分数据。BatchSampler 生成 batch 索引。sampler 与 shuffle 互斥(指定 sampler 时需 shuffle=False)。

pytorch
# gradient norm clipping (most common)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

# gradient value clipping
torch.nn.utils.clip_grad_value_(model.parameters(), clip_value=0.5)

# in a training loop
optimizer.zero_grad()
loss = criterion(model(x), y)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()

# inspect gradient norm
total_norm = torch.norm(torch.stack([
    p.grad.norm(2) for p in model.parameters() if p.grad is not None
]), 2)

可迭代数据集

IterableDataset 适用于流式数据(无法随机访问的场景,如日志、网络流)。实现 __iter__ 返回迭代器。多 worker 时需自行分片,否则各 worker 会重复读取相同数据。无法 shuffle,需在缓冲区内做近似 shuffle。

pytorch
import torch.optim as optim

# get / set learning rate
for g in optimizer.param_groups:
    print(g['lr'])
    g['lr'] = 1e-4

# manual warmup
for step in range(warmup_steps):
    lr = base_lr * (step + 1) / warmup_steps
    for g in optimizer.param_groups:
        g['lr'] = lr

# optimizer also accepts schedulers (see lr-scheduler section)
from torch.optim.lr_scheduler import StepLR
scheduler = StepLR(optimizer, step_size=10, gamma=0.1)

Other Optimizers

RMSprop is the classic RNN optimizer and still works well. Adagrad accumulates squared gradients and can stall early — fine for sparse features but not deep nets. LBFGS needs a closure and full-batch evaluation — great for small convex problems or fine-tuning, too slow for large deep nets. For most new projects, start with AdamW.

pytorch
import torch.optim as optim

# RMSprop — good for RNNs
optim.RMSprop(model.parameters(), lr=1e-3, alpha=0.99)

# Adagrad — adapts per-parameter, accumulates squared grads
optim.Adagrad(model.parameters(), lr=1e-2)

# Adadelta — Adagrad variant with running average
optim.Adadelta(model.parameters(), rho=0.9)

# LBFGS — full-batch quasi-Newton, second-order
optim.LBFGS(model.parameters(), lr=1.0, max_iter=20)

# NAdam — Nesterov-accelerated Adam
optim.NAdam(model.parameters(), lr=1e-3)
08

卷积神经网络

Conv2d

kernel_size 是滤波器尺寸,stride 控制步长,padding 控制零填充。输出尺寸 = (W - k + 2p) / s + 1。groups>1 实现分组卷积(深度可分离卷积的基石)。dilation>1 实现空洞卷积,扩大感受野而不增参数。

pytorch
model = MLP(784, 128, 10).to(device)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()

for epoch in range(epochs):
    model.train()
    for x, y in train_loader:
        x, y = x.to(device), y.to(device)

        optimizer.zero_grad()
        logits = model(x)
        loss = criterion(logits, y)
        loss.backward()
        optimizer.step()

    print(f"epoch {epoch}, loss={loss.item():.4f}")

池化

MaxPool2d 取局部最大值,最常用。AvgPool2d 取平均。AdaptiveAvgPool2d 直接指定输出尺寸(如 1x1 实现全局平均池化 GAP),自动计算 stride/kernel。池化无参数,降低空间分辨率,提供平移不变性。

pytorch
for epoch in range(epochs):
    model.train()
    for x, y in train_loader:
        # ... training step ...

    # validation
    model.eval()
    val_loss, correct, total = 0.0, 0, 0
    with torch.no_grad():
        for x, y in val_loader:
            x, y = x.to(device), y.to(device)
            logits = model(x)
            val_loss += criterion(logits, y).item() * x.size(0)
            correct += (logits.argmax(1) == y).sum().item()
            total += x.size(0)

    print(f"val_loss={val_loss/total:.4f}, acc={correct/total:.4f}")

基础 CNN 架构

经典结构:卷积+激活+池化为一个 block,逐步降空间分辨率、升通道数。最后用 GAP 或 Flatten 接全连接分类。通道数常按 2 的幂递增(32→64→128)。输出尺寸需手动计算以确保 Flatten 后维度匹配。

pytorch
accum_steps = 4
optimizer.zero_grad()

for i, (x, y) in enumerate(train_loader):
    x, y = x.to(device), y.to(device)
    loss = criterion(model(x), y) / accum_steps
    loss.backward()

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

# equivalent effective batch = accum_steps * batch_size
# useful when GPU memory limits batch size

批归一化

BatchNorm2d 按 batch 维度归一化,含可学习的 γ、β。eval 时使用 running_mean/var 而非当前 batch 统计。小 batch 时 BN 不稳定,可考虑 GroupNorm。BN 在 train/eval 行为不同,记得调用 model.train()/eval()。

pytorch
from torch.amp import autocast, GradScaler

scaler = GradScaler('cuda')

for x, y in train_loader:
    x, y = x.to(device), y.to(device)
    optimizer.zero_grad()

    with autocast('cuda'):
        logits = model(x)
        loss = criterion(logits, y)

    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

# autocast casts select ops to float16/bf16 for speed
# GradScaler prevents underflow by scaling the loss

Dropout

训练时随机置零神经元(概率 p),测试时不丢弃但输出乘以 (1-p)。Dropout2d 丢弃整个通道(适合 CNN)。eval 时自动关闭。常放在全连接层之间。p 过大会欠拟合,过小无效果,典型 0.1-0.5。

pytorch
from tqdm import tqdm

for epoch in range(epochs):
    model.train()
    pbar = tqdm(train_loader, desc=f"epoch {epoch}")
    running = 0.0
    for i, (x, y) in enumerate(pbar):
        x, y = x.to(device), y.to(device)
        optimizer.zero_grad()
        loss = criterion(model(x), y)
        loss.backward()
        optimizer.step()

        running = 0.9 * running + 0.1 * loss.item()
        pbar.set_postfix(loss=f"{running:.4f}")

# TensorBoard logging (see visualization section)
writer.add_scalar('train/loss', loss.item(), global_step)

Reproducibility

Full reproducibility needs seeds for Python, NumPy, PyTorch CPU and CUDA. cudnn.deterministic=True disables non-deterministic algorithms (some conv backward ops). DataLoader workers each need seeding via worker_init_fn because they fork with their own RNG state. Even then, multi-GPU and some cuDNN ops may not be bit-reproducible — aim for statistical reproducibility.

pytorch
import torch, numpy as np, random

def set_seed(seed=42):
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    # cudnn determinism (may slow down)
    torch.backends.cudnn.deterministic = True
    torch.backends.cudnn.benchmark = False

set_seed(42)

# DataLoader worker seeding
def seed_worker(worker_id):
    worker_seed = torch.initial_seed() % 2**32
    np.random.seed(worker_seed)
    random.seed(worker_seed)

g = torch.Generator().manual_seed(42)
loader = DataLoader(ds, batch_size=32, worker_init_fn=seed_worker, generator=g)
09

RNN/LSTM

RNN 基础

nn.RNN 输入 (seq, batch, input_size) 或 (batch, seq, input_size)(batch_first=True)。返回 output(所有时间步的隐藏状态)和 h_n(最后时间步的隐藏状态)。RNN 难以训练长序列(梯度消失),实践中多用 LSTM/GRU。

pytorch
model.eval()
test_loss, correct, total = 0.0, 0, 0

with torch.no_grad():
    for x, y in test_loader:
        x, y = x.to(device), y.to(device)
        logits = model(x)
        test_loss += criterion(logits, y).item() * x.size(0)
        preds = logits.argmax(dim=1)
        correct += (preds == y).sum().item()
        total += x.size(0)

print(f"test_loss={test_loss/total:.4f}, acc={correct/total:.4f}")

LSTM

LSTM 通过门控机制(输入门、遗忘门、输出门)缓解梯度消失,能建模长距离依赖。隐藏状态包括 h(短时)和 c(长时记忆)。bidirectional=True 双向 LSTM 同时利用前后文,输出维度翻倍。是 NLP 任务的传统主力。

pytorch
logits = torch.randn(8, 10)
targets = torch.randint(0, 10, (8,))

# top-1 accuracy
preds = logits.argmax(dim=1)
acc = (preds == targets).float().mean().item()

# top-k accuracy
topk = 3
_, idx = logits.topk(topk, dim=1)         # shape (N, k)
correct = idx.eq(targets.view(-1, 1)).any(dim=1)
topk_acc = correct.float().mean().item()

# per-class accuracy
for c in range(10):
    mask = targets == c
    print(c, (preds[mask] == targets[mask]).float().mean().item())

GRU

GRU 是 LSTM 的简化版,合并了遗忘门与输入门为更新门,无单独的细胞状态。参数更少、训练更快,效果通常与 LSTM 相当。重置门控制前一时刻信息的影响。小数据集上常优于 LSTM。

pytorch
from sklearn.metrics import confusion_matrix, classification_report
import numpy as np

all_preds, all_labels = [], []
model.eval()
with torch.no_grad():
    for x, y in test_loader:
        logits = model(x.to(device))
        all_preds.append(logits.argmax(1).cpu())
        all_labels.append(y)

preds = torch.cat(all_preds).numpy()
labels = torch.cat(all_labels).numpy()

cm = confusion_matrix(labels, preds)
print(cm)
print(classification_report(labels, preds))

序列分类

取最后一个时间步的隐藏状态 h_n(或双向拼接)作为序列表示,接全连接层分类。注意 packed 序列时应取 pack 的最后输出而非 h_n。也可对隐藏状态做注意力加权求和(Self-Attention)。

pytorch
from torchmetrics import Accuracy, Precision, Recall, F1Score

acc = Accuracy(task='multiclass', num_classes=10).to(device)
prec = Precision(task='multiclass', num_classes=10, average='macro').to(device)
rec = Recall(task='multiclass', num_classes=10, average='macro').to(device)
f1 = F1Score(task='multiclass', num_classes=10, average='macro').to(device)

model.eval()
with torch.no_grad():
    for x, y in test_loader:
        logits = model(x.to(device))
        preds = logits.argmax(1)
        acc.update(preds, y.to(device))
        prec.update(preds, y.to(device))

print(acc.compute(), prec.compute(), f1.compute())

打包序列

pack_padded_sequence 将变长序列打包(去掉 padding 的计算),RNN 不会在 padding 上浪费算力。pad_packed_sequence 还原为 padded 张量。需按长度降序输入并传 enforce_sorted=True(或 False 自动排序)。强烈推荐用于变长序列任务。

pytorch
from sklearn.model_selection import KFold

kfold = KFold(n_splits=5, shuffle=True, random_state=42)
fold_accs = []

for fold, (train_idx, val_idx) in enumerate(kfold.split(dataset)):
    train = torch.utils.data.Subset(dataset, train_idx)
    val = torch.utils.data.Subset(dataset, val_idx)

    model = build_model()                  # fresh model per fold
    train_model(model, train)              # your training function

    acc = evaluate(model, val)
    fold_accs.append(acc)
    print(f"fold {fold}: acc={acc:.4f}")

print(f"mean={np.mean(fold_accs):.4f} +/- {np.std(fold_accs):.4f}")

Model Comparison

Always compare models on the same held-out test set with the same seed and metric. Beyond accuracy, look at calibration, inference latency, and memory. Use bootstrap confidence intervals when test-set differences are small. Keep a results table (CSV or W&B) so you can compare against future experiments instead of just the last run.

pytorch
results = {}
for name, model in models.items():
    model.eval()
    accs, losses = [], []
    with torch.no_grad():
        for x, y in test_loader:
            x, y = x.to(device), y.to(device)
            logits = model(x)
            losses.append(criterion(logits, y).item())
            accs.append((logits.argmax(1) == y).float().mean().item())
    results[name] = {
        'acc': np.mean(accs),
        'loss': np.mean(losses),
    }

for name, r in sorted(results.items(), key=lambda kv: -kv[1]['acc']):
    print(f"{name:20s} acc={r['acc']:.4f} loss={r['loss']:.4f}")
10

Transformer

多头注意力

MultiheadAttention 将 Q、K、V 投影到多个子空间分别做注意力再拼接。num_heads 通常为 8 的倍数,embedding_dim 必须能被 num_heads 整除。batch_first=True 让输入为 (batch, seq, dim)。attn_mask 可屏蔽特定位置(如 padding 或因果掩码)。

pytorch
import torch

# save (recommended: only state_dict)
torch.save(model.state_dict(), 'model.pth')

# load
model = MLP(784, 128, 10)        # must match the saved architecture
model.load_state_dict(torch.load('model.pth', map_location='cpu'))
model.eval()

# load on GPU that was trained on GPU
model.load_state_dict(torch.load('model.pth'))
model.to(device)

# strict=False allows partial loading (transfer learning)
model.load_state_dict(torch.load('model.pth'), strict=False)

位置编码

Transformer 无位置感知,需显式注入位置信息。正弦余弦编码是固定的(不可学习),适合短序列。可学习位置编码对固定长度序列有效。RoPE(旋转位置编码)支持外推,是现代 LLM 主流。ALiBi 通过注意力偏置线性外推。

pytorch
# save entire model (uses pickle) — NOT recommended
torch.save(model, 'model_full.pth')

# load
model = torch.load('model_full.pth', map_location='cpu')
model.eval()

# pickle limitations:
# - class definition must be importable at load time
# - breaks if you refactor or rename the class
# - not portable across PyTorch versions
# prefer state_dict for production

Transformer 编码器

TransformerEncoderLayer 包含自注意力 + FFN,各带残差连接与 LayerNorm。TransformerEncoder 堆叠多层。前置 LayerNorm(pre-norm)比后置更稳定,便于训练深层模型。FFN 中间维度通常为 4 倍 embedding_dim。

pytorch
checkpoint = {
    'epoch': epoch,
    'model_state': model.state_dict(),
    'optimizer_state': optimizer.state_dict(),
    'scheduler_state': scheduler.state_dict() if scheduler else None,
    'loss': loss.item(),
    'rng_state': torch.get_rng_state(),
    'cuda_rng_state': torch.cuda.get_rng_state_all(),
}
torch.save(checkpoint, f'ckpt_epoch{epoch}.pth')

# resume
ckpt = torch.load('ckpt_epoch10.pth', map_location='cpu')
model.load_state_dict(ckpt['model_state'])
optimizer.load_state_dict(ckpt['optimizer_state'])
start_epoch = ckpt['epoch'] + 1
torch.set_rng_state(ckpt['rng_state'])

nn.Transformer

nn.Transformer 集成了编码器和解码器,适合 seq2seq 任务。src_mask 屏蔽源序列 padding,tgt_mask 屏蔽目标 padding 并用因果掩码防止看到未来。generate_square_subsequent_mask 生成上三角因果掩码。新代码建议用 TransformerEncoder 配合自定义解码。

pytorch
best_acc = 0.0
for epoch in range(epochs):
    train_one_epoch(model, ...)
    acc = evaluate(model, val_loader)

    if acc > best_acc:
        best_acc = acc
        torch.save(model.state_dict(), 'best.pth')
        print(f"new best {acc:.4f} saved")

# load best at the end
model.load_state_dict(torch.load('best.pth'))
final_acc = evaluate(model, test_loader)

完整 Transformer 模型

典型结构:Token Embedding + 位置编码 + N 层 Encoder + 池化([CLS] 或平均)+ 分类头。warmup 学习率调度对 Transformer 训练稳定性至关重要。LayerNorm 优于 BatchNorm(不受序列长度影响)。Dropout 通常放在注意力与 FFN 内部。

pytorch
# transfer learning: load a backbone, drop the classifier
pretrained = torch.load('resnet_backbone.pth')
model = ResNet(num_classes=100)

# drop mismatched keys (e.g. final fc)
missing, unexpected = model.load_state_dict(pretrained, strict=False)
print('missing:', missing)        # e.g. ['fc.weight', 'fc.bias']
print('unexpected:', unexpected)

# rename keys before loading
pretrained = {k.replace('encoder.', 'backbone.'): v
              for k, v in pretrained.items()}

# freeze loaded weights
for name, p in model.named_parameters():
    if 'fc' not in name:
        p.requires_grad_(False)

Export to TorchScript

TorchScript produces a self-contained graph you can load in LibTorch (C++) or Python without the original class — ideal for deployment. trace captures one control-flow path; script handles data-dependent control flow but requires type annotations. ONNX exports for inference in ONNX Runtime, TensorRT, or OpenVINO. Use dynamic_axes for variable batch size.

pytorch
# scripted (production-ready, runs without Python)
scripted = torch.jit.script(model)
scripted.save('model_scripted.pt')

# traced (records a static graph from sample input)
example = torch.randn(1, 3, 224, 224)
traced = torch.jit.trace(model, example)
traced.save('model_traced.pt')

# load TorchScript in C++ / Python
loaded = torch.jit.load('model_scripted.pt')
loaded.eval()
out = loaded(example)

# export to ONNX (cross-framework)
torch.onnx.export(model, example, 'model.onnx',
                  input_names=['input'],
                  output_names=['output'],
                  dynamic_axes={'input': {0: 'batch'}})
11

训练循环

基础训练循环

标准流程:zero_grad → forward → loss → backward → step。loss.backward() 计算梯度,optimizer.step() 更新参数。loss.item() 取标量值(会触发 GPU 同步,避免在循环内频繁调用)。每个 epoch 后打印平均损失以监控训练。

pytorch
import torch.nn as nn

conv = nn.Conv2d(
    in_channels=3,        # RGB input
    out_channels=16,      # number of filters
    kernel_size=3,        # 3x3 filter
    stride=1,
    padding=1,            # 'same' for kernel 3, stride 1
    bias=True,
)
print(conv.weight.shape)   # (16, 3, 3, 3)

x = torch.randn(4, 3, 32, 32)   # NCHW
out = conv(x)
print(out.shape)                # (4, 16, 32, 32) with padding=1

# depthwise separable conv
dw = nn.Conv2d(16, 16, 3, padding=1, groups=16)

带验证的训练

每个 epoch 后在验证集评估,注意 model.eval() 关闭 BN/Dropout,并 torch.no_grad() 关闭梯度。验证指标用于早停、学习率调度、模型选择。最佳模型应保存 state_dict 并在训练结束后加载。验证集不应参与任何训练决策的反向传播。

pytorch
import torch.nn as nn

x = torch.randn(4, 16, 32, 32)

# max pooling
pool = nn.MaxPool2d(kernel_size=2, stride=2)
out = pool(x)              # shape (4, 16, 16, 16)

# average pooling
avg = nn.AvgPool2d(2, stride=2)

# adaptive pool — output size fixed regardless of input
gap = nn.AdaptiveAvgPool2d(1)   # global average pool
out = gap(x)                    # shape (4, 16, 1, 1)

# fractional max pool (slightly improves accuracy)
fmp = nn.FractionalMaxPool2d(2, output_ratio=0.5)

检查点

保存的不只是模型,还有 optimizer、scheduler、epoch、best_metric,用于断点续训。建议定期保存(如每 N 步)而非每 epoch,以防长训练中断。state_dict 是字典,可扩展保存任意训练状态(如 scaler、early_stopping 计数器)。

pytorch
import torch.nn as nn

class TinyCNN(nn.Module):
    def __init__(self, num_classes=10):
        super().__init__()
        self.features = nn.Sequential(
            nn.Conv2d(3, 32, 3, padding=1),
            nn.BatchNorm2d(32),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2),                       # 32 -> 16
            nn.Conv2d(32, 64, 3, padding=1),
            nn.BatchNorm2d(64),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2),                       # 16 -> 8
        )
        self.classifier = nn.Sequential(
            nn.AdaptiveAvgPool2d(1),
            nn.Flatten(),
            nn.Linear(64, num_classes),
        )

    def forward(self, x):
        return self.classifier(self.features(x))

早停

监控验证指标,连续 patience 个 epoch 无改善则停止,避免过拟合与算力浪费。保存 best_state 供恢复。min_delta 避免微小波动被误判为改善。mode 区分指标方向(min 用于 loss,max 用于 accuracy)。

pytorch
# visualize feature maps
model.eval()
x = images[0:1].to(device)
with torch.no_grad():
    feat = model.features[0](x)   # first conv output
print(feat.shape)                 # (1, 32, H, W)

import matplotlib.pyplot as plt
fig, axes = plt.subplots(4, 8, figsize=(12, 6))
for i, ax in enumerate(axes.flat):
    if i < feat.shape[1]:
        ax.imshow(feat[0, i].cpu(), cmap='viridis')
    ax.axis('off')

# hook to capture intermediate activations
acts = {}
def hook(module, inp, out): acts['layer'] = out
model.features[3].register_forward_hook(hook)

指标追踪

torchmetrics 提供标准化的指标计算,自动处理 batch 累积与设备迁移。accuracy、f1_score、precision、recall 是分类常用指标。compute() 在 epoch 末返回最终值,reset() 清零状态开始新 epoch。也可集成 TensorBoard/WandB 进行可视化。

pytorch
import torch.nn as nn

# 2D BN for conv outputs
bn = nn.BatchNorm2d(64)
bn.weight.shape    # (64,) — learnable scale gamma
bn.bias.shape      # (64,) — learnable shift beta
bn.running_mean.shape  # (64,) — EMA of batch means
bn.running_var.shape   # (64,) — EMA of batch vars

# 1D BN for linear / time-series
bn1d = nn.BatchNorm1d(128)

# in training mode: use batch stats, update running stats
model.train()
# in eval mode: use running stats (deterministic)
model.eval()

# alternatives
nn.LayerNorm(128)        # normalize over features (transformers)
nn.GroupNorm(8, 64)      # split channels into 8 groups
nn.InstanceNorm2d(64)    # per-sample, style transfer

训练函数模板

封装为可复用函数,参数化模型、数据加载器、损失、优化器、训练配置。tqdm 提供进度条。返回训练历史便于绘图。这是大多数训练脚本的起点,建议作为模板保存并按需扩展(如分布式、AMP、梯度累积)。

pytorch
import torchvision.models as M

# pretrained backbones
resnet = M.resnet50(weights=M.ResNet50_Weights.DEFAULT)
effnet = M.efficientnet_b0(weights=M.EfficientNet_B0_Weights.DEFAULT)
vgg = M.vgg16(weights=M.VGG16_Weights.DEFAULT)
mobnet = M.mobilenet_v3_small(weights=M.MobileNet_V3_Small_Weights.DEFAULT)

# replace classifier head
num_classes = 100
resnet.fc = nn.Linear(resnet.fc.in_features, num_classes)

# feature extractor (drop head)
backbone = nn.Sequential(*list(resnet.children())[:-1])  # outputs (N, 2048, 1, 1)
features = backbone(images).flatten(1)                    # (N, 2048)
12

评估与测试

评估函数

评估循环:model.eval() + torch.no_grad(),遍历测试集累计 loss 与指标。注意 BN 用 running 统计而非 batch 统计。yield 生成器形式便于逐样本分析。最终指标用整体加权平均,而非 batch 平均(batch 大小不同会导致偏差)。

pytorch
import torch.nn as nn

rnn = nn.RNN(
    input_size=64,
    hidden_size=128,
    num_layers=2,
    batch_first=True,        # input shape (batch, seq, feature)
    bidirectional=False,
    nonlinearity='tanh',
)

x = torch.randn(32, 10, 64)  # (batch, seq, input)
out, h_n = rnn(x)
print(out.shape)             # (32, 10, 128) — per-step hidden states
print(h_n.shape)             # (2, 32, 128) — final hidden state per layer

# last-step output for classification
last = out[:, -1, :]         # shape (32, 128)

混淆矩阵

混淆矩阵揭示每类被分到各类的次数,对角线为正确预测。可分析哪些类别易混淆(如 9 和 4、5 和 3)。多分类任务中比单一准确率更有信息量。可用 seaborn 热力图可视化。每类 precision/recall 可直接从矩阵计算。

pytorch
import torch.nn as nn

lstm = nn.LSTM(
    input_size=64,
    hidden_size=128,
    num_layers=2,
    batch_first=True,
    bidirectional=True,
    dropout=0.5,             # between stacked layers
)

x = torch.randn(32, 10, 64)
out, (h_n, c_n) = lstm(x)
print(out.shape)    # (32, 10, 256) — 128 * 2 directions
print(h_n.shape)    # (4, 32, 128) — 2 layers * 2 directions
print(c_n.shape)    # (4, 32, 128) — cell state

# concatenate final forward + backward for classification
last = torch.cat([h_n[-2], h_n[-1]], dim=1)   # shape (32, 256)

每类准确率

类别不平衡时整体准确率有误导性(如 99% 负样本时全预测负样本也有 99%)。每类准确率(召回率)暴露少数类的表现。macro-average 对各类等权平均,micro-average 受大类主导。还建议看 F1、AUC 等阈值无关指标。

pytorch
import torch.nn as nn

gru = nn.GRU(
    input_size=64,
    hidden_size=128,
    num_layers=2,
    batch_first=True,
)

x = torch.randn(32, 10, 64)
out, h_n = gru(x)            # GRU has no cell state
print(out.shape)             # (32, 10, 128)
print(h_n.shape)             # (2, 32, 128)

# GRU vs LSTM: fewer parameters, slightly faster,
# often similar performance on many tasks

推理与预测

单样本推理前需 model.eval() + no_grad(),并将输入张量化并添加 batch 维(unsqueeze(0))。生产环境建议批处理以利用并行。注意预处理(resize、normalize)必须与训练一致,否则性能骤降。可用 TorchScript 导出以脱离 Python。

pytorch
from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence

# variable-length sequences in a batch
seqs = [torch.randn(5, 64), torch.randn(3, 64), torch.randn(8, 64)]
lens = torch.tensor([5, 3, 8])

# pad to (batch, max_len, feature)
padded = torch.nn.utils.rnn.pad_sequence(seqs, batch_first=True)
# sort by length descending (required by older versions)
padded = padded[torch.argsort(lens, descending=True)]
lens_sorted = lens[torch.argsort(lens, descending=True)]

packed = pack_padded_sequence(padded, lens_sorted,
                              batch_first=True, enforce_sorted=True)
out_packed, h_n = lstm(packed)
out, _ = pad_packed_sequence(out_packed, batch_first=True)

测试时增强

对同一输入做多种增强(翻转、缩放)分别推理,取预测平均或投票。小幅提升准确率但增加推理时间。常用于竞赛刷分。注意增强应保持语义不变(如 MNIST 的垂直翻转会改变 6 和 9 的含义)。

pytorch
import torch.nn as nn

class TextClassifier(nn.Module):
    def __init__(self, vocab, emb_dim=128, hidden=256, classes=10):
        super().__init__()
        self.emb = nn.Embedding(vocab, emb_dim, padding_idx=0)
        self.lstm = nn.LSTM(emb_dim, hidden, batch_first=True,
                            bidirectional=True, num_layers=2)
        self.dropout = nn.Dropout(0.5)
        self.fc = nn.Linear(hidden * 2, classes)

    def forward(self, x, lengths):
        x = self.emb(x)
        packed = pack_padded_sequence(x, lengths.cpu(),
                                      batch_first=True, enforce_sorted=False)
        _, (h_n, _) = self.lstm(packed)
        last = torch.cat([h_n[-2], h_n[-1]], dim=1)
        return self.fc(self.dropout(last))

model = TextClassifier(vocab=10000)

Sequence Generation (Decoder)

A decoder runs one step at a time, conditioning on the previous output token. LSTMCell is a single-step LSTM — use it for fine-grained control over the loop. Teacher forcing feeds the ground-truth token at the next step during training, which converges much faster than feeding the model's own predictions. At inference, switch to feeding argmax (or sample) from the previous step's logits.

pytorch
import torch.nn as nn

class Decoder(nn.Module):
    def __init__(self, vocab, emb_dim=128, hidden=256):
        super().__init__()
        self.emb = nn.Embedding(vocab, emb_dim)
        self.lstm = nn.LSTMCell(emb_dim, hidden)
        self.fc = nn.Linear(hidden, vocab)

    def forward(self, tgt, hidden):
        # tgt shape (batch, T)
        outputs = []
        h, c = hidden
        for t in range(tgt.size(1)):
            emb = self.emb(tgt[:, t])
            h, c = self.lstm(emb, (h, c))
            outputs.append(self.fc(h))
        return torch.stack(outputs, dim=1)   # (batch, T, vocab)

# teacher forcing: feed ground truth next token
# at inference: feed argmax of previous output
13

保存与加载

保存与加载 state_dict

推荐方式:只保存 state_dict(参数字典),加载时需先实例化模型再 load_state_dict。模型代码必须可导入,否则无法恢复。strict=True 检查键完全匹配,部分加载(如迁移学习)用 strict=False。CPU 加载用 map_location='cpu'。

pytorch
import torch.nn as nn

attn = nn.MultiheadAttention(
    embed_dim=512,
    num_heads=8,
    dropout=0.1,
    batch_first=True,        # input (batch, seq, dim)
)

q = k = v = torch.randn(32, 10, 512)
out, weights = attn(q, k, v, need_weights=True)
print(out.shape)       # (32, 10, 512)
print(weights.shape)   # (32, 10, 10) — attention scores

# causal / self-attention mask
mask = nn.Transformer.generate_square_subsequent_mask(10)
out, _ = attn(q, k, v, attn_mask=mask)   # masked attention

# key padding mask (ignore pad positions)
key_padding_mask = (tokens == 0)   # True = ignore
out, _ = attn(q, k, v, key_padding_mask=key_padding_mask)

保存整个模型

用 pickle 保存整个模型对象,无需实例化即可加载。缺点:强依赖 Python 类定义与导入路径,跨版本/跨环境易出错。仅建议用于快速实验,生产部署仍用 state_dict 或 TorchScript/ONNX。

pytorch
import torch.nn as nn

enc_layer = nn.TransformerEncoderLayer(
    d_model=512,
    nhead=8,
    dim_feedforward=2048,
    dropout=0.1,
    activation='relu',
    batch_first=True,
    norm_first=True,        # pre-LN, more stable
)
encoder = nn.TransformerEncoder(enc_layer, num_layers=6)

x = torch.randn(32, 10, 512)
out = encoder(x)
print(out.shape)            # (32, 10, 512)

# with padding mask
pad_mask = (tokens == 0)
out = encoder(x, src_key_padding_mask=pad_mask)

训练检查点

断点续训需保存:模型 state_dict、优化器 state_dict(动量/Adam 矩阵)、当前 epoch、最佳指标、调度器状态。加载时按相同顺序恢复。这是长训练任务的必备保险。保存路径建议带 epoch 后缀以防覆盖。

pytorch
import torch
import torch.nn as nn
import math

class PositionalEncoding(nn.Module):
    def __init__(self, d_model=512, max_len=5000, dropout=0.1):
        super().__init__()
        self.dropout = nn.Dropout(dropout)
        pe = torch.zeros(max_len, d_model)
        pos = torch.arange(0, max_len).unsqueeze(1).float()
        div = torch.exp(torch.arange(0, d_model, 2).float()
                        * (-math.log(10000.0) / d_model))
        pe[:, 0::2] = torch.sin(pos * div)
        pe[:, 1::2] = torch.cos(pos * div)
        self.register_buffer('pe', pe.unsqueeze(0))   # (1, max_len, d_model)

    def forward(self, x):
        x = x + self.pe[:, :x.size(1)]
        return self.dropout(x)

# usage
emb = nn.Embedding(vocab, 512)
pos = PositionalEncoding(512)
x = pos(emb(tokens))   # (batch, seq, 512)

保存多个模型

GAN、师生网络等场景有多个模型。统一保存到一个字典,键为模型名。加载时分别取出。注意优化器也应一并保存。同一文件便于版本管理,但文件较大时可分文件存储。

pytorch
import torch.nn as nn

class Seq2SeqTransformer(nn.Module):
    def __init__(self, vocab_src, vocab_tgt, d=512, h=8, layers=6):
        super().__init__()
        self.src_emb = nn.Embedding(vocab_src, d, padding_idx=0)
        self.tgt_emb = nn.Embedding(vocab_tgt, d, padding_idx=0)
        self.pos = PositionalEncoding(d)
        enc_layer = nn.TransformerEncoderLayer(d, h, d*4, batch_first=True, norm_first=True)
        dec_layer = nn.TransformerDecoderLayer(d, h, d*4, batch_first=True, norm_first=True)
        self.encoder = nn.TransformerEncoder(enc_layer, layers)
        self.decoder = nn.TransformerDecoder(dec_layer, layers)
        self.head = nn.Linear(d, vocab_tgt)

    def forward(self, src, tgt, src_pad, tgt_pad):
        src = self.pos(self.src_emb(src))
        tgt = self.pos(self.tgt_emb(tgt))
        memory = self.encoder(src, src_key_padding_mask=src_pad)
        causal = nn.Transformer.generate_square_subsequent_mask(tgt.size(1))
        out = self.decoder(tgt, memory,
                           tgt_mask=causal,
                           tgt_key_padding_mask=tgt_pad,
                           memory_key_padding_mask=src_pad)
        return self.head(out)

设备感知加载

map_location='cpu' 将 GPU 模型加载到 CPU。map_location={'cuda:0':'cuda:1'} 重映射 GPU。无 GPU 环境加载 GPU 训练的模型必须用此参数,否则报错。加载后再 .to(device) 迁移到目标设备。

pytorch
import torch.nn as nn

class ViTBlock(nn.Module):
    def __init__(self, dim=768, heads=12, mlp=3072, dropout=0.1):
        super().__init__()
        self.norm1 = nn.LayerNorm(dim)
        self.attn = nn.MultiheadAttention(dim, heads, dropout=dropout, batch_first=True)
        self.norm2 = nn.LayerNorm(dim)
        self.mlp = nn.Sequential(
            nn.Linear(dim, mlp),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(mlp, dim),
            nn.Dropout(dropout),
        )

    def forward(self, x):
        h = self.norm1(x)
        a, _ = self.attn(h, h, h, need_weights=False)
        x = x + a
        x = x + self.mlp(self.norm2(x))
        return x

# patchify: (B, 3, 224, 224) -> (B, 196, 768) via Conv2d(3, 768, 16, 16)

Attention from Scratch

Scaled dot-product attention: scores = QK^T / sqrt(d_k), softmax, multiply by V. The 1/sqrt(d_k) keeps variance stable so softmax does not saturate. masked_fill with -inf before softmax zeroes out those positions. Multi-head attention splits the embedding into H heads, attends in parallel, and concatenates. This is exactly what nn.MultiheadAttention does — write it from scratch only to learn or to add variants like linear/flash attention.

pytorch
import torch
import torch.nn.functional as F
import math

def attention(q, k, v, mask=None, dropout=None):
    # q, k, v shape (batch, heads, seq, dim)
    d_k = q.size(-1)
    scores = (q @ k.transpose(-2, -1)) / math.sqrt(d_k)
    if mask is not None:
        scores = scores.masked_fill(mask == 0, float('-inf'))
    attn = F.softmax(scores, dim=-1)
    if dropout is not None:
        attn = dropout(attn)
    out = attn @ v
    return out, attn

# multi-head: reshape (batch, seq, dim) -> (batch, heads, seq, dim//heads)
def split_heads(x, heads):
    b, s, d = x.shape
    return x.view(b, s, heads, d // heads).transpose(1, 2)
14

GPU/设备管理

设备选择

torch.cuda.is_available() 检测 CUDA。model.to(device) 和 tensor.to(device) 迁移到目标设备。模型与输入必须在同一设备,否则报错。建议在脚本开头定义 device 变量统一使用。MPS 是 Apple Silicon 的后端。

pytorch
import torch

# pick device
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(device)                          # cuda or cpu
print(torch.cuda.device_count())       # number of GPUs
print(torch.cuda.get_device_name(0))   # GPU name

# use a specific GPU
torch.cuda.set_device(0)
x = torch.randn(3, device='cuda:1')    # directly on GPU 1

# default device context (PyTorch 2.x)
with torch.device('cuda'):
    y = torch.randn(4)   # on GPU

# Apple Silicon
device = 'mps' if torch.backends.mps.is_available() else 'cpu'

内存管理

torch.cuda.empty_cache() 释放缓存的显存(但不改变已分配峰值)。del tensor 后再调用才有效。OOM 常因计算图保留中间结果,用 backward 后立即释放或用 gradient checkpointing。监控显存用 nvidia-smi 或 torch.cuda.memory_allocated()。

pytorch
model = MLP(784, 128, 10).to(device)

for x, y in train_loader:
    x, y = x.to(device), y.to(device)
    # non-blocking copy if pin_memory=True
    x = x.to(device, non_blocking=True)
    ...

# check placement
print(next(model.parameters()).device)   # cuda:0
print(x.device)

# move a tensor to CPU for numpy / saving
arr = x.detach().cpu().numpy()

# half precision
model = model.half()        # float16
x = x.half()

多 GPU

nn.DataParallel 简单但效率低(GIL 限制、负载不均)。nn.parallel.DistributedDataParallel 是推荐方案,每 GPU 一个进程。device_ids 列表指定使用的 GPU。DP 在前向时复制模型,DDP 在初始化时一次性复制。新项目直接用 DDP。

pytorch
import torch.nn as nn

# simple, single-process data parallelism
model = MLP(784, 512, 10).to(device)
model = nn.DataParallel(model)   # wrap AFTER .to(device)

# now model behaves normally
out = model(x)
loss = criterion(out, y)
loss.backward()   # gradients are summed across GPUs

# access underlying module
inner = model.module

# specify which GPUs
model = nn.DataParallel(model, device_ids=[0, 1, 2])

CUDA 流

Stream 是异步执行的命令队列。不同 stream 间可并行(如计算与数据传输重叠)。synchronize() 等待所有 stream 完成(用于精确计时)。默认 stream 是阻塞的。非默认 stream 间需要 event 同步。高级优化时使用。

pytorch
import os
import torch.distributed as dist
import torch.nn as nn

# launch with: torchrun --nproc_per_node=4 train.py
dist.init_process_group(backend='nccl')
local_rank = int(os.environ['LOCAL_RANK'])
torch.cuda.set_device(local_rank)

model = MLP(784, 512, 10).to(local_rank)
model = nn.parallel.DistributedDataParallel(
    model, device_ids=[local_rank], output_device=local_rank,
)

sampler = torch.utils.data.distributed.DistributedSampler(dataset)
loader = DataLoader(dataset, batch_size=64, sampler=sampler)

for epoch in range(epochs):
    sampler.set_epoch(epoch)   # shuffle differently each epoch
    for x, y in loader:
        x, y = x.to(local_rank), y.to(local_rank)
        loss = criterion(model(x), y)
        loss.backward()
        optimizer.step()
        optimizer.zero_grad()

可复现性

设置随机种子保证可复现:torch.manual_seed、np.random.seed、random.seed。CUDA 还需 torch.cuda.manual_seed_all 和 cudnn.deterministic=True(可能降低性能)。DataLoader 的 worker_init_fn 保证多进程随机性一致。完全复现仍受硬件影响。

pytorch
import torch

# release unused memory
torch.cuda.empty_cache()

# current memory usage
print(torch.cuda.memory_allocated() / 1e9, 'GB allocated')
print(torch.cuda.memory_reserved() / 1e9, 'GB reserved')

# per-tensor memory
x = torch.randn(1000, 1000, device='cuda')
print(x.element_size() * x.nelement() / 1e6, 'MB')

# peak memory tracking
torch.cuda.reset_peak_memory_stats()
# ... run training ...
print(torch.cuda.max_memory_allocated() / 1e9, 'GB peak')

# out-of-memory debugging
import os
os.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'expandable_segments:True'

CuDNN & Performance

cudnn.benchmark=True lets cuDNN profile algorithms on the first forward and cache the fastest — great when input shapes are fixed (typical CNN), harmful if shapes change every step (RNNs with variable length). TF32 uses 19-bit mantissa on Ampere+ for ~3x matmul speedup with minimal accuracy loss; enable it unless you need strict fp32 reproducibility. Deterministic mode disables non-deterministic algorithms.

pytorch
import torch.backends.cudnn as cudnn

# benchmark finds the fastest conv algorithm (good for fixed input sizes)
cudnn.benchmark = True

# deterministic mode (slower, reproducible)
cudnn.deterministic = True
cudnn.benchmark = False

# allow TF32 on Ampere+ (big speedup, slight precision loss)
torch.backends.cuda.matmul.allow_tf32 = True
cudnn.allow_tf32 = True

# disable cudnn (debugging)
cudnn.enabled = False
15

分布式训练

DDP 设置

DistributedDataParallel 是多 GPU 训练的推荐方案:每 GPU 一个进程,模型副本独立,梯度通过 all-reduce 同步。init_process_group 初始化通信后端(nccl for GPU)。set_device 限定本进程使用的 GPU。DDP 性能优于 DataParallel。

pytorch
from torchvision import transforms

transform = transforms.Compose([
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),                       # PIL -> Tensor, scales to [0,1]
    transforms.Normalize(                        # per-channel standardize
        mean=[0.485, 0.456, 0.406],
        std=[0.229, 0.224, 0.225],
    ),
])

# apply to a PIL image
img = transform(pil_image)
print(img.shape)       # (3, 224, 224)
print(img.min(), img.max())   # standardized

dataset = datasets.ImageFolder('data/', transform=transform)

torchrun

torchrun 是启动分布式训练的标准工具,自动设置环境变量并管理进程。--nproc_per_node 指定每节点 GPU 数。比 mp.spawn 更稳健,支持容错(某进程崩溃时整体重启)。LOCAL_RANK 环境变量由 torchrun 注入。生产环境首选。

pytorch
from torchvision import transforms

train_tf = transforms.Compose([
    transforms.RandomResizedCrop(224),
    transforms.RandomHorizontalFlip(),
    transforms.ColorJitter(brightness=0.2, contrast=0.2,
                           saturation=0.2, hue=0.1),
    transforms.RandomRotation(15),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406],
                         std=[0.229, 0.224, 0.225]),
])

# stronger augmentation: RandAugment / TrivialAugment
train_tf = transforms.Compose([
    transforms.TrivialAugmentWide(),
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406],
                         std=[0.229, 0.224, 0.225]),
])

DistributedSampler

DistributedSampler 将数据集分片到各进程,避免重复计算。每个 epoch 调用 set_epoch 改变随机种子,否则 shuffle 顺序固定。配合 DDP 使用时 DataLoader 的 shuffle 必须为 False(由 sampler 负责)。drop_last 默认 True 保证各进程 batch 数一致。

pytorch
from torchvision.transforms import v2

transform = v2.Compose([
    v2.RandomResizedCrop(224, antialias=True),
    v2.RandomHorizontalFlip(),
    v2.ToDtype(torch.float32, scale=True),   # scale to [0,1]
    v2.Normalize(mean=[0.485, 0.456, 0.406],
                 std=[0.229, 0.224, 0.225]),
])

# v2 works on images, videos, masks, and bounding boxes together
imgs, masks = transform(imgs, masks)
imgs, bboxes, labels = transform(imgs, bboxes, labels)

# supports more ops and JIT scripting
transform = v2.RandomChoice([
    v2.RandomHorizontalFlip(),
    v2.RandomVerticalFlip(),
])

DataParallel

nn.DataParallel 是早期的多 GPU 方案,单进程多线程。使用简单(一行代码)但效率低:受 GIL 限制、每次前向复制模型、scatter/gather 开销大。仅适合快速实验,正式训练应迁移到 DDP。

pytorch
import torch
from torchvision import transforms

class RandomGaussianNoise:
    def __init__(self, std=0.1):
        self.std = std
    def __call__(self, x):
        if self.std > 0:
            return x + torch.randn_like(x) * self.std
        return x

# compose with built-ins
transform = transforms.Compose([
    transforms.ToTensor(),
    RandomGaussianNoise(0.05),
    transforms.Normalize([0.5], [0.5]),
])

# callable with probability
class RandomApply:
    def __init__(self, fn, p=0.5):
        self.fn, self.p = fn, p
    def __call__(self, x):
        return self.fn(x) if torch.rand(1) < self.p else x

梯度同步

DDP 默认在 backward 时自动同步梯度。find_unused_parameters=True 处理部分参数未参与前向的情况(有性能开销)。no_sync 上下文跳过同步,用于梯度累积(多步积累再同步一次)。this.rank() 返回本进程 rank,用于日志与检查点保存(仅 rank 0 保存)。

pytorch
import torch
from torchtext.vocab import build_vocab_from_iterator

# tokenize and numericalize text
tokens = ['hello', 'world', 'this', 'is', 'nlp']
vocab = build_vocab_from_iterator([tokens], specials=['<unk>', '<pad>'])
vocab.set_default_index(vocab['<unk>'])

# text -> token ids
ids = vocab(['hello', 'world', 'unknownword'])
print(ids)   # e.g. [3, 4, 0]

# tabular: standardize features
from sklearn.preprocessing import StandardScaler
scaler = StandardScaler().fit(X_train)
X_train_t = torch.tensor(scaler.transform(X_train), dtype=torch.float32)
X_test_t = torch.tensor(scaler.transform(X_test), dtype=torch.float32)

# fit scaler ONLY on training data, apply to val/test
16

自定义数据集

自定义 Dataset 类

继承 Dataset 并实现 __len__ 和 __getitem__。__getitem__ 应返回单个样本(通常是 (input, target) 元组)。支持索引访问,DataLoader 据此批量加载。避免在 __getitem__ 中做重计算(如读大文件),最好预加载到内存或用内存映射。

pytorch
from torch.utils.data import Dataset

class ImageDataset(Dataset):
    def __init__(self, paths, labels, transform=None):
        self.paths = paths
        self.labels = labels
        self.transform = transform

    def __len__(self):
        return len(self.paths)

    def __getitem__(self, idx):
        from PIL import Image
        img = Image.open(self.paths[idx]).convert('RGB')
        label = self.labels[idx]
        if self.transform:
            img = self.transform(img)
        return img, label

dataset = ImageDataset(paths, labels, transform=train_tf)
img, label = dataset[0]

图像数据集

PIL 读取图像,ToTensor 自动归一化到 [0,1] 并调整通道顺序 (H,W,C) → (C,H,W)。Normalize 用 ImageNet 均值方差标准化(迁移学习必备)。路径列表 + 标签的轻量实现比继承 ImageFolder 更灵活。图像可在 __getitem__ 中实时增强。

pytorch
from torchvision import datasets

# convention: data/train/cat/*.jpg, data/train/dog/*.jpg
dataset = datasets.ImageFolder('data/train', transform=train_tf)
print(dataset.class_to_idx)   # {'cat': 0, 'dog': 1}

# train / val split
from torch.utils.data import random_split
train, val = random_split(dataset, [0.8, 0.2])

# custom split with same class mapping
val_dataset = datasets.ImageFolder('data/val', transform=val_tf)
val_dataset.class_to_idx = dataset.class_to_idx   # keep consistent

# use with DataLoader
loader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=4)

变换

torchvision.transforms 提供常用图像变换。Compose 串联多个变换。RandomResizedCrop、RandomHorizontalFlip 做训练增强。Resize+CenterCrop 用于测试。新版推荐 v2 transforms(基于 Kornia/torch tensors,更快且支持视频)。注意测试变换不应含随机性。

pytorch
from torch.utils.data import Dataset
import torch

class TextDataset(Dataset):
    def __init__(self, texts, labels, vocab, max_len=128):
        self.texts = texts
        self.labels = labels
        self.vocab = vocab
        self.max_len = max_len

    def __len__(self):
        return len(self.texts)

    def __getitem__(self, idx):
        tokens = self.texts[idx].split()[:self.max_len]
        ids = [self.vocab[t] for t in tokens]
        ids += [0] * (self.max_len - len(ids))   # pad
        return torch.tensor(ids), self.labels[idx]

# variable length: use custom collate
def text_collate(batch):
    ids, labels = zip(*batch)
    lens = torch.tensor([len(x) for x in ids])
    padded = torch.nn.utils.rnn.pad_sequence(ids, batch_first=True)
    return padded, lens, torch.tensor(labels)

文本数据集

构建词表(vocab)将 token 映射为整数索引。特殊 token:<pad>(填充)、<unk>(未知词)、<bos>/<eos>。变长序列需 collate_fn 填充对齐。可加载预训练 embedding(如 GloVe)并按词表索引对齐。现代做法多用 subword tokenizer(BPE/WordPiece)避免 OOV。

pytorch
from torch.utils.data import IterableDataset, DataLoader
import itertools

class StreamDataset(IterableDataset):
    def __init__(self, file_path):
        self.file_path = file_path

    def __iter__(self):
        # split work across workers
        worker_info = torch.utils.data.get_worker_info()
        with open(self.file_path) as f:
            if worker_info is None:
                lines = f
            else:
                lines = itertools.islice(f, worker_info.id,
                                         None, worker_info.num_workers)
            for line in lines:
                yield self._parse(line)

    def _parse(self, line):
        # parse line into tensor + label
        return torch.tensor([...]), label

loader = DataLoader(StreamDataset('big.csv'), batch_size=32)

自定义 Dataset 的 DataLoader

与标准 DataLoader 用法一致。num_workers>0 多进程预取(Windows 下需 if __name__=='__main__' 保护)。pin_memory=True 加速 GPU 传输。自定义 collate_fn 处理变长样本或返回 dict。prefetch_factor 控制 worker 预取 batch 数。

pytorch
from torch.utils.data import TensorDataset, ConcatDataset, DataLoader

# wrap in-memory tensors
X = torch.randn(1000, 784)
y = torch.randint(0, 10, (1000,))
dataset = TensorDataset(X, y)
x, label = dataset[0]

# combine multiple datasets
full = ConcatDataset([dataset_a, dataset_b])
print(len(full))   # len(a) + len(b)

# Subset for indexing
from torch.utils.data import Subset
half = Subset(dataset, range(500))

# random splits
from torch.utils.data import random_split
train, val = random_split(dataset, [0.8, 0.2])
17

迁移学习

预训练模型

torchvision.models 提供经典模型(ResNet、VGG、EfficientNet)的预训练权重。pretrained=True 自动下载。特征已学习通用视觉表示(边缘、纹理、对象部件),迁移到小数据集效果远好于从零训练。新代码用 weights=ResNet18_Weights.DEFAULT 替代 deprecated 的 pretrained 参数。

pytorch
from torch.optim.lr_scheduler import StepLR, MultiStepLR

# decay by gamma every step_size epochs
scheduler = StepLR(optimizer, step_size=10, gamma=0.1)
# LR: 1e-3 -> 1e-4 (epoch 10) -> 1e-5 (epoch 20)

# decay at specific epochs
scheduler = MultiStepLR(optimizer, milestones=[30, 60, 80], gamma=0.1)

# in the training loop
for epoch in range(epochs):
    train(...)
    val(...)
    scheduler.step()    # call once per epoch

# inspect current LR
print(scheduler.get_last_lr())

特征提取

冻结骨干网络(requires_grad=False),仅训练新分类头。适合小数据集,避免过拟合。前向时 no_grad 进一步省内存。注意 BN 层即使 requires_grad=False,eval 模式才用 running 统计。数据量小且与原任务相似时首选此方案。

pytorch
from torch.optim.lr_scheduler import CosineAnnealingLR

# cosine decay to eta_min over T_max epochs
scheduler = CosineAnnealingLR(optimizer, T_max=50, eta_min=1e-6)
# LR follows cosine from base_lr down to eta_min

# with warm restarts
from torch.optim.lr_scheduler import CosineAnnealingWarmRestarts
scheduler = CosineAnnealingWarmRestarts(
    optimizer,
    T_0=10,             # length of first cycle
    T_mult=2,           # each cycle is T_mult longer
    eta_min=1e-6,
)

微调

解冻部分或全部层,用较小学习率(如骨干 1e-4、新头 1e-3)联合训练。适合数据量较大或与原任务差异较大的场景。可从最后几层开始逐步解冻(layer-wise fine-tuning)。学习率过大易破坏预训练特征。

pytorch
from torch.optim.lr_scheduler import ReduceLROnPlateau

scheduler = ReduceLROnPlateau(
    optimizer,
    mode='min',         # 'min' for loss, 'max' for accuracy
    factor=0.5,         # multiply LR by factor on plateau
    patience=3,         # epochs to wait before reducing
    min_lr=1e-7,
    threshold=1e-4,
)

for epoch in range(epochs):
    train(...)
    val_loss = validate(...)
    scheduler.step(val_loss)   # pass the metric!

冻结与解冻

requires_grad=False 冻结参数不参与梯度计算与更新。逐层解冻是渐进式微调的常见策略:先训新头收敛,再解冻高层,最后解冻全部。每阶段可降低学习率。child_name 精确控制子模块。state_dict 中仍包含冻结参数。

pytorch
from torch.optim.lr_scheduler import OneCycleLR

# one cycle: warm up, peak, anneal down
scheduler = OneCycleLR(
    optimizer,
    max_lr=1e-3,
    total_steps=epochs * len(loader),   # or steps_per_epoch + epochs
    pct_start=0.3,       # fraction spent warming up
    anneal_strategy='cos',
)

# step per BATCH (not per epoch)
for epoch in range(epochs):
    for x, y in loader:
        train_step(...)
        scheduler.step()   # call once per batch

# simpler cyclic LR
from torch.optim.lr_scheduler import CyclicLR
scheduler = CyclicLR(optimizer, base_lr=1e-5, max_lr=1e-3, step_size_up=500)

torchvision 模型

替换分类头:将最后 fc 层替换为输出新类别数的 Linear。新层随机初始化,需训练。也可添加自定义头(多层 MLP、Dropout)。in_features 从原 fc 读取保证维度匹配。Detection/Segmentation 模型(如 Mask R-CNN)同样支持 num_classes 替换。

pytorch
from torch.optim.lr_scheduler import LinearLR, SequentialLR, CosineAnnealingLR

# linear warmup for first N steps
warmup = LinearLR(optimizer, start_factor=0.1, total_iters=500)

# cosine decay after warmup
cosine = CosineAnnealingLR(optimizer, T_max=total_steps - 500, eta_min=1e-6)

# combine
scheduler = SequentialLR(
    optimizer,
    schedulers=[warmup, cosine],
    milestones=[500],
)

# manual warmup (transformers-style)
def lr_lambda(step):
    if step < 500:
        return step / 500
    return 0.5 * (1 + math.cos(math.pi * (step - 500) / (total - 500)))

from torch.optim.lr_scheduler import LambdaLR
scheduler = LambdaLR(optimizer, lr_lambda)

Custom Scheduler

LambdaLR multiplies the base LR by lambda(step) — the most flexible way to implement custom schedules. Pass a list of lambdas (one per param_group) for per-group schedules, e.g. decay weights but not biases, or use different schedules for backbone vs head in transfer learning. Lambdas must be cheap because they run every step. Avoid Python closures that capture mutable state — keep them pure.

pytorch
from torch.optim.lr_scheduler import LambdaLR
import math

# any function of step / epoch
def warmup_cosine(step, warmup=500, total=10000):
    if step < warmup:
        return step / warmup
    progress = (step - warmup) / (total - warmup)
    return 0.5 * (1 + math.cos(math.pi * progress))

scheduler = LambdaLR(optimizer, lr_lambda=warmup_cosine)

# per-parameter-group schedule (e.g. decay only weights, not bias)
scheduler = LambdaLR(
    optimizer,
    lr_lambda=[lambda s: 0.9 ** s, lambda s: 1.0],  # one per group
)

# inspect
for step in range(0, total, 1000):
    print(step, scheduler.get_last_lr())
    scheduler.step()
18

可视化

TensorBoard 设置

SummaryWriter 写入事件文件,tensorboard --logdir 启动 Web 界面查看。可记录标量(loss/acc)、图像、直方图、网络结构图、嵌入投影。add_scalar 在训练循环中调用。定期记录而非每步以减少 IO。进程结束前 writer.close() 刷盘。

pytorch
import torchvision.models as M

# load a model with pretrained ImageNet weights
weights = M.ResNet50_Weights.DEFAULT
model = M.resnet50(weights=weights)

# get the matching preprocessing transform
preprocess = weights.transforms()
print(preprocess)   # resize, crop, normalize with ImageNet stats

# model is ready for 1000-class ImageNet inference
model.eval()
with torch.no_grad():
    out = model(preprocess(img).unsqueeze(0))
    pred = out.argmax(1)
    print(weights.meta['categories'][pred])

记录图像

add_image 需要 (C,H,W) 格式张量。make_grid 将多张图拼成网格便于批量查看。可记录预测 vs 真实标签,或特征图、注意力热力图。dataformats 参数支持不同格式(如 'HWC')。图像增强后的样本可视化有助于验证变换正确性。

pytorch
import torch.nn as nn
import torchvision.models as M

# load backbone, freeze it, replace head
model = M.resnet50(weights=M.ResNet50_Weights.DEFAULT)

# freeze all parameters
for p in model.parameters():
    p.requires_grad_(False)

# replace classifier for new task (10 classes)
model.fc = nn.Linear(model.fc.in_features, 10)
# only the new fc has requires_grad=True

# train only the head
optimizer = torch.optim.Adam(model.fc.parameters(), lr=1e-3)

# inference-friendly: extract features once
features = torch.nn.Sequential(*list(model.children())[:-1])
feat = features(images).flatten(1)   # (N, 2048)

Visdom

Visdom 是 Meta 的可视化工具,支持实时更新图表,适合远程服务器训练监控。需先启动 visdom.server。提供 scatter、line、image 等接口。比 TensorBoard 更轻量但功能较少。也可考虑现代替代如Weights & Biases (WandB)、MLflow。

pytorch
import torch.nn as nn
import torch.optim as optim
import torchvision.models as M

model = M.resnet50(weights=M.ResNet50_Weights.DEFAULT)
model.fc = nn.Linear(model.fc.in_features, 10)

# differential learning rates: lower for backbone, higher for head
optimizer = optim.Adam([
    {'params': [p for n, p in model.named_parameters() if 'fc' not in n], 'lr': 1e-4},
    {'params': model.fc.parameters(), 'lr': 1e-3},
])

# gradual unfreezing: train head first, then unfreeze later layers
def unfreeze(layer_idx):
    for n, p in model.named_parameters():
        if f'layer{layer_idx}' in n:
            p.requires_grad_(True)

# warmup head only for a few epochs, then unfreeze layer4, etc.

模型检查

named_parameters 暴露参数名与形状,便于发现维度错误。weights & bias 的统计(均值、方差、范围)能反映训练健康度——权重爆炸/消失、NaN 都能及早发现。make_dot 生成计算图可视化(需 torchviz),有助于理解复杂模型结构与梯度流。

pytorch
model = M.resnet50(weights=M.ResNet50_Weights.DEFAULT)
model.fc = nn.Linear(model.fc.in_features, 10)

# freeze everything except fc
for name, p in model.named_parameters():
    p.requires_grad = ('fc' in name)

# freeze specific submodule
for p in model.layer1.parameters():
    p.requires_grad_(False)

# count trainable params
n = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f"trainable: {n:,}")

# unfreeze after warmup
for p in model.layer4.parameters():
    p.requires_grad_(True)

# trainable params still need grad to flow: ensure no torch.no_grad
# and that frozen layers don't break autograd

特征可视化

前向钩子捕获中间层输出而无需修改模型。可视化特征图可观察每层学到什么(浅层边缘、深层语义部件)。detach 输出避免追踪梯度。Grad-CAM 是更高级技术,高亮驱动预测的图像区域,常用于模型可解释性分析。

pytorch
import torch.nn as nn
import torchvision.models as M

model = M.resnet50(weights=M.ResNet50_Weights.DEFAULT)
# ResNet: model.fc
model.fc = nn.Linear(model.fc.in_features, num_classes)

# EfficientNet: model.classifier[1]
model = M.efficientnet_b0(weights=M.EfficientNet_B0_Weights.DEFAULT)
model.classifier[1] = nn.Linear(model.classifier[1].in_features, num_classes)

# ViT: model.heads.head
model = M.vit_b_16(weights=M.ViT_B_16_Weights.DEFAULT)
model.heads.head = nn.Linear(model.heads.head.in_features, num_classes)

# multi-task head
class MultiTask(nn.Module):
    def __init__(self, backbone, n_cls, n_reg):
        super().__init__()
        self.backbone = backbone
        in_f = backbone.fc.in_features
        backbone.fc = nn.Identity()
        self.cls = nn.Linear(in_f, n_cls)
        self.reg = nn.Linear(in_f, n_reg)
    def forward(self, x):
        f = self.backbone(x)
        return self.cls(f), self.reg(f)
19

混合精度 (AMP)

Autocast

自动混合精度(AMP)用 autocast 包装前向,自动为矩阵乘法等受益操作选择 fp16,同时保持归约在 fp32 以保证数值稳定。GradScaler 放大损失防止 fp16 梯度下溢,在 optimizer.step 前反缩放。在现代 GPU 上可获 2-3 倍加速且精度损失极小。

pytorch
from torch.utils.tensorboard import SummaryWriter

writer = SummaryWriter('runs/experiment_1')

# log scalars (loss, accuracy, LR)
writer.add_scalar('train/loss', loss.item(), step)
writer.add_scalar('val/accuracy', acc, step)
writer.add_scalar('lr', optimizer.param_groups[0]['lr'], step)

# log multiple metrics together
writer.add_scalars('losses', {
    'train': train_loss,
    'val': val_loss,
}, step)

# close at the end
writer.close()

# launch: tensorboard --logdir=runs

GradScaler

GradScaler 动态调整损失缩放因子。梯度溢出(inf/nan)时跳过 optimizer.step 并降低缩放。调用 unscale_ 后再做梯度裁剪,使裁剪看到真实梯度幅值。get_scale 返回当前缩放因子,比较 update 前后可知本步是否被跳过。

pytorch
writer = SummaryWriter('runs/exp')

# log model graph (needs a sample input)
model = MLP(784, 128, 10)
writer.add_graph(model, torch.randn(1, 784))

# weight histograms (per layer, over training)
for name, p in model.named_parameters():
    writer.add_histogram(f'weights/{name}', p, step)
    if p.grad is not None:
        writer.add_histogram(f'grads/{name}', p.grad, step)

# text
writer.add_text('config', 'lr=1e-3, batch=32', 0)

# close at the end
writer.close()

完整 AMP 训练循环

生产级 AMP 循环。zero_grad(set_to_none=True) 比置零更快。non_blocking=True 让 H2D 传输与计算重叠。正确顺序:unscale_ → clip_grad_norm_ → step → update。step 被跳过时本次不更新参数,update 会调整缩放供下次使用。

pytorch
writer = SummaryWriter('runs/exp')

# single image (CHW tensor in [0,1] or [0,255])
writer.add_image('val/sample', img_tensor, step)

# grid of images
from torchvision.utils import make_grid
grid = make_grid(images, nrow=8, normalize=True)
writer.add_image('val/batch', grid, step)

# figure from matplotlib
import matplotlib.pyplot as plt
fig, ax = plt.subplots()
ax.plot([1,2,3], [1,4,9])
writer.add_figure('curves/loss', fig, step)

# overlay segmentation mask on image
from torchvision.utils import draw_segmentation_masks
overlay = draw_segmentation_masks(img, masks, alpha=0.5)
writer.add_image('val/masks', overlay, step)

BFloat16

BF16 与 FP32 共享 8 位指数,动态范围相同——无溢出/下溢,不需要 GradScaler。尾数只有 7 位(FP16 有 10 位),精度略低。推荐在 Ampere 及以上 GPU(A100、H100)使用。用 torch.autocast(device_type='cuda', dtype=torch.bfloat16)。

pytorch
from torchviz import make_dot

x = torch.randn(1, 784, requires_grad=True)
y = model(x)
loss = criterion(y, torch.tensor([1]))
loss.backward()

# render the autograd graph
dot = make_dot(loss, params=dict(model.named_parameters()))
dot.render('model_graph', format='png')   # saves model_graph.png

# also show saved tensors (memory)
dot = make_dot(loss, params=dict(model.named_parameters()),
               show_attrs=True, show_saved=True)

AMP 最佳实践

autocast 仅包装前向与损失,backward 应在 autocast 外执行。autocast 自动将归约(sum/mean/softmax)保持 fp32 并在不支持的算子上回退。手动混合精度需保证 matmul 操作数 dtype 一致。TensorCore 数学需 fp16 输入且维度为 8 的倍数(更高 TFLOPS 需 16 倍数)。

pytorch
import matplotlib.pyplot as plt

# visualize conv filters
filters = model.features[0].weight.data.clone().cpu()
fig, axes = plt.subplots(4, 8, figsize=(10, 5))
for i, ax in enumerate(axes.flat):
    if i < filters.size(0):
        ax.imshow(filters[i, 0], cmap='gray')
    ax.axis('off')

# hook to capture activations
acts = {}
def hook(name):
    def fn(mod, inp, out): acts[name] = out.detach().cpu()
    return fn
model.features[3].register_forward_hook(hook('l1'))

# activation statistics over training
for name, p in model.named_parameters():
    print(name, p.data.mean().item(), p.data.std().item(),
          p.grad.norm().item() if p.grad is not None else None)
20

TorchScript 与 ONNX

torch.jit.trace

Tracing 记录在示例输入上实际执行的操作。局限:控制流(if/for)是静态的——只记录示例输入走到的分支。模型有数据相关控制流时用 script 模式。适合静态计算图模型(大多数 CNN、Transformer)。可用多组示例输入提升鲁棒性。

pytorch
import os, torch
import torch.distributed as dist
import torch.nn as nn

def setup():
    dist.init_process_group(backend='nccl')
    local_rank = int(os.environ['LOCAL_RANK'])
    torch.cuda.set_device(local_rank)
    return local_rank

def cleanup():
    dist.destroy_process_group()

# launch with torchrun
# torchrun --nproc_per_node=4 train.py
# rank, world_size, LOCAL_RANK are set by torchrun

local_rank = setup()
model = model.to(local_rank)
model = nn.parallel.DistributedDataParallel(
    model, device_ids=[local_rank], output_device=local_rank,
)

torch.jit.script

Script 模式直接将 Python 源码解析为 TorchScript IR,支持 if/for/while 与 TorchScript 类型注解(-> torch.Tensor)。限制:不支持第三方库、无动态类属性、Python 标准库支持有限。用 model.code 检查生成的代码。适合 trace 无法处理的含控制流模型。

pytorch
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler

sampler = DistributedSampler(
    dataset,
    num_replicas=dist.get_world_size(),
    rank=dist.get_rank(),
    shuffle=True,
    drop_last=True,
)
loader = DataLoader(dataset, batch_size=64, sampler=sampler,
                    num_workers=4, pin_memory=True)

for epoch in range(epochs):
    sampler.set_epoch(epoch)   # reshuffle per epoch
    for x, y in loader:
        # each GPU sees a non-overlapping subset
        ...

保存与加载 TorchScript

TorchScript 产物自包含,可在 libtorch (C++) 中加载而无需 Python 类定义。freeze() 折叠常量并移除训练专用算子以加速推理。部署时建议 freeze + optimize_for_inference + save 产出最精简产物。map_location 用于跨设备加载。

pytorch
# single-node, 4 GPUs
torchrun --nproc_per_node=4 train.py

# multi-node (2 nodes, 4 GPUs each = 8 total)
# node 0 (master):
torchrun --nproc_per_node=4 --nnodes=2 --node_rank=0 \
    --master_addr=10.0.0.1 --master_port=29500 train.py

# node 1:
torchrun --nproc_per_node=4 --nnodes=2 --node_rank=1 \
    --master_addr=10.0.0.1 --master_port=29500 train.py

# in code, world_size = nnodes * nproc_per_node
print(dist.get_world_size())   # 8

# also supports elastic restart on node failure
torchrun --nproc_per_node=4 --rdzv_backend=c10d \
    --rdzv_endpoint=10.0.0.1:29500 train.py

ONNX 导出

ONNX 是跨框架部署的交换格式(支持 TensorRT、OpenVINO、CoreML、ONNX Runtime)。opset_version 控制可用算子(现代算子建议 14+)。dynamic_axes 使维度可变以支持运行时变 batch。导出后用 onnxruntime 验证数值与 PyTorch 输出一致(torch.allclose)。自定义算子可能不可导出——新代码可试 torch.onnx.dynamo_export。

pytorch
# DDP all-reduces gradients after backward by default
loss.backward()   # gradients synchronized across GPUs

# skip sync for unused parameters (small speedup)
model = nn.parallel.DistributedDataParallel(
    model, device_ids=[local_rank],
    find_unused_parameters=False,   # set True if some params unused
)

# manual gradient sync (advanced)
for p in model.parameters():
    if p.grad is not None:
        dist.all_reduce(p.grad, op=dist.ReduceOp.SUM)
        p.grad /= dist.get_world_size()

# only sync on the last micro-batch (gradient accumulation)
# use model.no_sync() context for the non-final steps
with model.no_sync():
    loss.backward()

TorchScript 优化

optimize_for_inference 融合逐元素算子并转为推理友好布局。freeze() 折叠常量。Conv+BN 融合在 trace 时(model.eval())自动发生。最大吞吐需组合 TorchScript + TensorCore autocast (fp16) + pinned-memory 输入流水线。务必前后 benchmark 确认收益。

pytorch
# only log from rank 0 to avoid duplicates
if dist.get_rank() == 0:
    writer.add_scalar('train/loss', loss.item(), step)
    print(f'step {step}: loss={loss.item():.4f}')

# only save checkpoint from rank 0
if dist.get_rank() == 0:
    torch.save({
        'model': model.module.state_dict(),   # unwrap DDP
        'optimizer': optimizer.state_dict(),
        'epoch': epoch,
    }, 'ckpt.pt')

# load on all ranks, then wrap with DDP
state = torch.load('ckpt.pt', map_location='cpu')
model.load_state_dict(state['model'])
model = model.to(local_rank)
model = nn.parallel.DistributedDataParallel(model, device_ids=[local_rank])

FullyShardedDataParallel (FSDP)

FSDP shards parameters, gradients, and optimizer states across GPUs — so a model too big for one GPU can train across several. Unlike DDP (which keeps full replicas), FSDP gathers shards on demand during forward/backward. Use it when the model doesn't fit in a single GPU's memory. Saving is trickier: wrap state_dict collection in FULL_STATE_DICT mode so rank 0 materializes the full model. Mixed precision + FSDP + activation checkpointing lets you train very large models.

pytorch
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import ShardingStrategy

model = LargeModel().to(local_rank)
model = FSDP(
    model,
    sharding_strategy=ShardingStrategy.FULL_SHARD,
    device_id=local_rank,
    use_orig_params=True,   # needed for save/load compatibility
)

# forward / backward as usual
loss = criterion(model(x), y)
loss.backward()
optimizer.step()

# save: only the full state on rank 0
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
import torch.distributed.fsdp.api as fsdp_api
with FSDP.state_dict_type(model, fsdp_api.StateDictType.FULL_STATE_DICT):
    state = model.state_dict()
    if dist.get_rank() == 0:
        torch.save(state, 'fsdp_ckpt.pt')

这篇内容对您有帮助吗?

学习路径

从零开始学习

通过结构化课程从头学习这个语言。