PyTorch 速查

QUICK REFERENCE / PYTORCH

形状、设备、梯度、模式

遇到错误时先检查这四项,再怀疑模型结构。

张量的三项检查

1
2
3
print(x.shape)
print(x.dtype)
print(x.device)

输入、标签与模型应使用兼容的形状、类型和设备。x.to(device) 返回转换结果,通常需要重新赋值;不是只写一行就自动移动原变量。见 PyTorch 张量教程。

一次训练更新

1
2
3
4
5
6
model.train()
optimizer.zero_grad(set_to_none=True)
prediction = model(x)
loss = loss_fn(prediction, y)
loss.backward()
optimizer.step()

梯度会累积;每步清理还是多步累积,应由训练方案明确决定。分类任务的损失函数对输入和标签有具体要求,不能把回归示例直接套过去。见 官方优化教程。

推理模式

1
2
3
model.eval()
with torch.no_grad():
prediction = model(x)

eval() 影响 Dropout、BatchNorm 等层的行为;no_grad() 控制梯度记录。两者不是替代关系。见 Module API。

可复现记录最小清单

  • 数据划分及其来源,尤其是病人或来源分组;
  • 软件版本、随机种子、模型与超参数;
  • 验证集选择规则和最终测试时机;
  • 运行命令、日志和模型保存位置。

随机种子不承诺所有设备与版本之间逐位一致;不要把一次运行的漂亮指标当成可靠结论。

打开 可运行的最小训练循环或返回 全部速查。

引用到评论
随便逛逛博客分类文章标签
复制地址关闭热评深色模式轉為繁體