PyTorch Quickstart教程训练准确率打印修改方案是否正确
结论
你的修改逻辑完全正确,可直接用于常规分类任务的神经网络训练。
核心逻辑正确性说明
- 训练流程顺序合规:进入函数后先调用
model.train()将模型切换至训练模式,保证Dropout、BatchNorm等训练专属层正常生效;前向推理计算损失、梯度清零、反向传播、优化器更新参数的顺序没有错误,符合PyTorch标准训练流程。 - 指标计算逻辑正确:
- 平均训练损失为所有batch损失累加后除以batch总数,计算逻辑和官方教程一致
- 训练准确率通过累加每个batch中预测类别(取
predictions.argmax(1)作为模型输出的预测类别)和真实标签匹配的样本数,最后除以总样本数得到,和Kerasfit()方法输出的训练准确率计算逻辑完全一致。
可选优化建议(非错误修正,可按需调整)
- 可以将
device设置为函数入参,避免依赖全局变量导致的报错,修改后函数签名可以改为def train(dataloader, model, optimizer, loss_fn, device):,后续传入对应设备参数即可。 - 可以在函数末尾返回
correct, training_loss两个指标值,方便后续记录训练曲线、做早停判断等,不必仅做打印输出。 - 如果你需要实时观测训练进度,也可以在batch循环中增加每N个batch打印一次当前batch指标的逻辑,目前每跑完一个完整epoch打印一次的逻辑也完全符合常规使用需求。
内容的提问来源于stack exchange,提问作者Abdullah Celik
相关产品推荐
相关产品推荐

