问题与目标
将全部数据一次性丢给模型,只能验证最小代码。真正的训练需要按批次读取数据、在训练与验证模式之间切换、汇总整个数据集的指标,并保存可以在新进程重建的模型产物。
本篇使用自建数值分类数据,完成 Dataset -> DataLoader -> train -> validate -> state_dict -> reload -> predict 闭环。数据量和网络都很小,CPU 即可运行。

上层是每个批次的参数更新,下层是验证、最佳权重保存和新进程推理,两条流程的模式和梯度状态不同。
核心概念
Dataset 管单条样本,DataLoader 组织批次
Dataset 通过 __len__() 报告样本数,通过 __getitem__() 返回一条样本。DataLoader 在其上处理批次组合、顺序打乱和多进程读取。
from torch.utils.data import Dataset
class PairDataset(Dataset):
def __init__(self, features, labels):
self.features = features
self.labels = labels
def __len__(self):
return len(self.labels)
def __getitem__(self, index):
return self.features[index], self.labels[index]
训练集通常打乱,验证与测试集通常不需要。时间序列、序列切片和分组数据的打乱还要遵守时间与主体边界。
训练循环和验证循环的职责
| 步骤 | 训练 | 验证/推理 |
|---|---|---|
| 模式 | model.train() | model.eval() |
| 梯度 | 启用 | torch.inference_mode() |
| 反向传播 | 有 | 无 |
| 参数更新 | 有 | 无 |
| Dropout | 随机失活 | 关闭 |
| BatchNorm | 更新运行统计 | 使用已保存统计 |
整个 epoch 的平均损失应按样本数加权,不要简单平均“各批损失”后忽略最后一个小批次。
保存权重而不是依赖当前 Python 对象
推理产物常保存 model.state_dict(),加载时先用相同构造参数重建网络,再加载权重。这比保存整个 Python 模型对象更容易跨文件重构和检查。
只加载可信来源的文件。对只包含张量权重的产物,使用 torch.load(path, weights_only=True) 减少反序列化暴露面。
“最佳模型”来自验证集
每个 epoch 后计算验证损失,只在它改善时写入检查点。测试集不用来选 epoch、学习率或网络结构。
可运行实现
from pathlib import Path
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset, random_split
torch.manual_seed(42)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
features = torch.randn(600, 4)
noise = torch.randn(600) * 0.35
labels = (
features[:, 0] - 0.8 * features[:, 1]
+ 0.5 * features[:, 2] ** 2 + noise > 0.4
).long()
dataset = TensorDataset(features, labels)
train_data, valid_data = random_split(
dataset, [480, 120], generator=torch.Generator().manual_seed(42)
)
train_loader = DataLoader(train_data, batch_size=32, shuffle=True)
valid_loader = DataLoader(valid_data, batch_size=64, shuffle=False)
class RiskNet(nn.Module):
def __init__(self) -> None:
super().__init__()
self.layers = nn.Sequential(
nn.Linear(4, 16), nn.ReLU(), nn.Linear(16, 2)
)
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
return self.layers(inputs)
def train_epoch(model, loader, loss_fn, optimizer) -> float:
model.train()
total_loss = 0.0
for batch_features, batch_labels in loader:
batch_features = batch_features.to(device)
batch_labels = batch_labels.to(device)
optimizer.zero_grad()
loss = loss_fn(model(batch_features), batch_labels)
loss.backward()
optimizer.step()
total_loss += loss.item() * len(batch_labels)
return total_loss / len(loader.dataset)
def evaluate(model, loader, loss_fn) -> tuple[float, float]:
model.eval()
total_loss, correct = 0.0, 0
with torch.inference_mode():
for batch_features, batch_labels in loader:
batch_features = batch_features.to(device)
batch_labels = batch_labels.to(device)
logits = model(batch_features)
total_loss += loss_fn(logits, batch_labels).item() * len(batch_labels)
correct += (logits.argmax(dim=1) == batch_labels).sum().item()
return total_loss / len(loader.dataset), correct / len(loader.dataset)
model = RiskNet().to(device)
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
checkpoint = Path("model_output/risk_net.pt")
checkpoint.parent.mkdir(parents=True, exist_ok=True)
best_loss = float("inf")
for epoch in range(1, 31):
train_loss = train_epoch(model, train_loader, loss_fn, optimizer)
valid_loss, valid_accuracy = evaluate(model, valid_loader, loss_fn)
if valid_loss < best_loss:
best_loss = valid_loss
torch.save(model.state_dict(), checkpoint)
if epoch in (1, 10, 20, 30):
print(epoch, round(train_loss, 4), round(valid_loss, 4), round(valid_accuracy, 3))
loaded = RiskNet().to(device)
loaded.load_state_dict(torch.load(checkpoint, map_location=device, weights_only=True))
loaded.eval()
with torch.inference_mode():
sample_probability = loaded(features[:1].to(device)).softmax(dim=1)
print("probability:", sample_probability.cpu().numpy().round(3))
输出应显示训练与验证损失总体下降,最后的概率形状为 (1, 2) 且和接近 1。由于设备与底层库实现可能存在非确定差异,不要将某个小数位当作固定答案。
常见问题与排查
- 验证结果每次不一样:确认调用
eval(),且验证集没有使用随机数据增强。 - DataLoader 多进程卡住:先用
num_workers=0验证逻辑,再根据操作系统调整。 - 保存的“最佳权重”最后变成最后一轮:直接
torch.save()到文件,或对内存中的state_dict做深拷贝。 - 加载时设备不可用:使用
map_location=device。 - 加载权重后维度不匹配:模型构造参数必须与训练时一致,并一起保存元数据。
- 只保存权重却无法处理原始输入:数据归一化、类别映射和输入契约也是推理产物。
小结
Dataset 提供单条样本,DataLoader 组织批次,训练循环负责梯度和参数更新,验证循环在不改参数的前提下选择方案。保存产物后重建模型并独立预测,才说明训练代码已经形成可复用流程。
License: CC BY-NC 4.0
Updated 2 hours ago
Was this article helpful? Give it a like.
0 comments


