You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch Quickstart教程训练准确率打印修改方案是否正确

结论

你的修改逻辑完全正确,可直接用于常规分类任务的神经网络训练。

核心逻辑正确性说明

  • 训练流程顺序合规:进入函数后先调用model.train()将模型切换至训练模式,保证Dropout、BatchNorm等训练专属层正常生效;前向推理计算损失、梯度清零、反向传播、优化器更新参数的顺序没有错误,符合PyTorch标准训练流程。
  • 指标计算逻辑正确:
    • 平均训练损失为所有batch损失累加后除以batch总数,计算逻辑和官方教程一致
    • 训练准确率通过累加每个batch中预测类别(取predictions.argmax(1)作为模型输出的预测类别)和真实标签匹配的样本数,最后除以总样本数得到,和Keras fit()方法输出的训练准确率计算逻辑完全一致。

可选优化建议(非错误修正,可按需调整)

  • 可以将device设置为函数入参,避免依赖全局变量导致的报错,修改后函数签名可以改为def train(dataloader, model, optimizer, loss_fn, device):,后续传入对应设备参数即可。
  • 可以在函数末尾返回correct, training_loss两个指标值,方便后续记录训练曲线、做早停判断等,不必仅做打印输出。
  • 如果你需要实时观测训练进度,也可以在batch循环中增加每N个batch打印一次当前batch指标的逻辑,目前每跑完一个完整epoch打印一次的逻辑也完全符合常规使用需求。

内容的提问来源于stack exchange,提问作者Abdullah Celik

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.24 16:45:02