PyTorch训练CIFAR100时train函数中10000数值的作用咨询
关于训练函数中数值10000的解释
先拆解这段代码的核心逻辑:
- 每处理一个mini-batch,就把当前batch的损失(
loss.item())累加到running_loss变量中 - 每完成100个mini-batch的处理(即
i % 100 == 99的判断触发时),就计算并打印平均损失,随后重置running_loss
10000的本质:一个错误的分母
这个10000完全是个无意义的错误数值,没有任何合理的技术或业务依据:
running_loss是连续100个mini-batch的损失总和,要得到这100个batch的平均损失,正确的分母应该是100,而非10000- 原代码里写10000大概率是手误,比如给100多打了两个0
修改数值后损失变化的原因
你修改这个数值后,打印出的损失值发生变化,本质就是平均计算的分母改变了:
- 比如改成100,打印的就是真实的100个batch的平均损失;改成1000,打印的就是真实平均损失的1/10
- 注意:这个数值只影响打印出来的损失值大小,不会改变模型的训练过程——模型参数更新只用到了每个batch的原始
loss,和这个打印用的平均计算无关
正确的写法
把打印行的running_loss / 10000改成running_loss / 100,这样打印的结果能准确反映连续100个mini-batch的平均损失,更利于观察训练状态:
print("[%d, %5d] loss: %.3f" % (epoch + 1, i + 1, running_loss / 100))
内容的提问来源于stack exchange,提问作者seni
相关产品推荐
相关产品推荐

