PyTorch中如何基于多损失平均梯度更新REINFORCE网络?
关于REINFORCE批量更新的问题解答
1. 基于损失平均梯度更新网络的实现方式
在PyTorch里,有两种实用的方式实现平均梯度更新:
- 方式一:累加梯度后取平均
遍历l_losses里的每个损失张量,依次调用loss.backward(retain_graph=True)(如果损失间共享计算图,必须保留图);全部反向传播完成后,遍历优化器的所有参数,把每个参数的梯度除以损失的数量(批量大小),再执行optimizer.step(),最后清零梯度。 - 方式二:计算平均损失后反向传播
用torch.stack(l_losses)把损失堆叠成张量,计算均值后调用backward(),再执行optimizer.step()并清零梯度。这种方式本质和第一种等价,因为均值损失的梯度就是各损失梯度的平均值。
2. 平均梯度更新与平均损失更新是否等价?
完全等价。从数学逻辑看,假设存在N个损失$L_1, L_2, ..., L_N$,平均损失$\bar{L} = \frac{1}{N}\sum_{i=1}^N L_i$,对$\bar{L}$求导得到的梯度$\nabla\bar{L} = \frac{1}{N}\sum_{i=1}^N \nabla L_i$,正好是各损失梯度的平均值,所以两种方式最终的参数更新量完全一致。
3. 当前方法是否符合需求?两种实现方式的对比
你当前用torch.stack(l_losses).mean().backward()的方法完全符合需求,和逐个调用backward()再平均梯度的效果完全相同。
两种方式的细微差异:
- 内存占用:逐个
backward()若需保留计算图,会占用更多内存;计算平均损失后反向传播只需处理一次计算图,内存效率更高。 - 代码简洁度:平均损失的方式代码更简洁、逻辑更清晰,更推荐使用。
注意事项:
- 如果每个episode的损失张量相互独立(无共享计算图),逐个
backward()时无需加retain_graph=True;若有共享部分,必须添加该参数避免报错。 - 无论用哪种方式,更新完成后都要调用
optimizer.zero_grad()清零梯度,防止下一次更新时累积旧梯度。
内容的提问来源于stack exchange,提问作者Josh
相关产品推荐
相关产品推荐

