PyTorch中批量损失与Epoch损失求和报错的解决问询
问题解决:PyTorch梯度计算中的Inplace操作错误
问题根源
你遇到的RuntimeError是因为后续epoch中使用的loss_epoch是包含完整计算图的张量,当它与当前batch的损失相加并执行反向传播后,后续重新计算loss_epoch的操作会修改该张量的版本,导致梯度计算时出现版本不匹配的问题。
解决方案
核心思路是断开loss_epoch的计算图,只使用它的数值参与batch损失求和,不让它参与当前的梯度计算。有两种简洁的修改方式:
方式一:在batch求和时断开计算图
loss_epoch = torch.tensor(0.0) # 初始化为浮点张量,匹配损失类型 for epoch in epochs: for batch in batches: optimizer.zero_grad() loss_batch = criterion_batch(output_batch, target_batch) # 仅使用loss_epoch的数值,断开其计算图 loss = loss_batch + loss_epoch.detach() loss.backward() optimizer.step() # 正常计算epoch损失,无需提前断开 loss_epoch = criterion_epoch(output_epoch, target_epoch)
方式二:保存epoch损失时直接断开计算图
loss_epoch = 0.0 for epoch in epochs: for batch in batches: optimizer.zero_grad() loss_batch = criterion_batch(output_batch, target_batch) loss = loss_batch + loss_epoch loss.backward() optimizer.step() # 计算后立即断开计算图,下一轮直接使用无梯度的数值 loss_epoch = criterion_epoch(output_epoch, target_epoch).detach()
原理说明
detach()方法会返回一个与原张量共享数据但不参与梯度计算的新张量,这样后续更新loss_epoch时,不会影响之前用于batch损失求和的张量版本,彻底避免了inplace操作导致的梯度计算冲突。
内容的提问来源于stack exchange,提问作者jaryl
相关产品推荐
相关产品推荐

