PyTorch 训练基础:Tensor、自动求导、训练循环与模型保存

Tensor

Tensor 是 PyTorch 的核心数据结构,类似 NumPy 数组但支持 GPU 运算:

import torch

# 0 维标量
t0 = torch.tensor(3.14)

# 1 维向量
t1 = torch.tensor([1.0, 2.0, 3.0])

# 2 维矩阵
t2 = torch.zeros(3, 4)

# 移到 GPU
if torch.cuda.is_available():
    t2 = t2.cuda()

Dynamic Graph 与 Autograd

PyTorch 使用动态计算图(Define-by-Run),每次前向传播都重新构建计算图,便于调试和条件分支。

开启梯度追踪:

x = torch.tensor(2.0, requires_grad=True)
y = x ** 2 + 3 * x

y.backward()     # 反向传播
print(x.grad)    # dy/dx = 2x + 3 = 7.0

标准训练循环

model = MyModel()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
loss_fn = torch.nn.CrossEntropyLoss()

for epoch in range(100):
    for x_batch, y_batch in dataloader:
        pred = model(x_batch)           # 1. 前向传播
        loss = loss_fn(pred, y_batch)   # 2. 计算损失
        
        optimizer.zero_grad()           # 3. 清空上一步梯度
        loss.backward()                 # 4. 反向传播
        optimizer.step()                # 5. 更新参数
    
    print(f"Epoch {epoch}, Loss: {loss.item():.4f}")

zero_grad() 必须在 backward() 前调用,否则梯度会累加。

损失函数选择

任务损失函数激活函数
二分类BCELossSigmoid
多分类CrossEntropyLoss无(内置 Softmax)
回归MSELoss
# 二分类:Logistic Regression
model = torch.nn.Sequential(
    torch.nn.Linear(10, 1),
    torch.nn.Sigmoid()
)
loss_fn = torch.nn.BCELoss()

# 多分类
model = torch.nn.Sequential(
    torch.nn.Linear(10, 5)  # 5 个类别
)
loss_fn = torch.nn.CrossEntropyLoss()  # 内部含 Softmax

# 回归
model = torch.nn.Linear(10, 1)
loss_fn = torch.nn.MSELoss()

优化器

# SGD(基础,需手动调 lr)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

# SGD + Momentum(更稳定)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9)

# Adam(最常用,lr=0.001 通常无需调整)
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

# AdamW(Transformer 推荐,带权重衰减修正)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01)

大多数情况从 Adam + lr=0.001 开始,效果不好再调。

保存与加载模型

推荐:只保存参数(state_dict)

# 保存
torch.save(model.state_dict(), "model.pth")

# 加载
model = MyModel()
model.load_state_dict(torch.load("model.pth"))
model.eval()  # 切换到推理模式

保存完整模型(不推荐,依赖类定义路径):

torch.save(model, "model_full.pth")
model = torch.load("model_full.pth")

导出 ONNX

ONNX 格式可在 TensorFlow、ONNX Runtime、TensorRT 等框架中使用:

dummy_input = torch.randn(1, 3, 224, 224)  # 与实际输入形状一致

torch.onnx.export(
    model,
    dummy_input,
    "model.onnx",
    input_names=["input"],
    output_names=["output"],
    opset_version=17
)

导出后可用 onnxruntime 推理,速度通常比 PyTorch 原生快。

推理模式

model.eval()

with torch.no_grad():  # 禁用梯度计算,节省内存
    output = model(input_tensor)

推理时必须调 model.eval()torch.no_grad(),否则 BatchNorm 和 Dropout 行为与训练时不同。