【问题标题】:In PyTorch, how do I update a neural network via the average gradient from a list of losses?在 PyTorch 中,如何通过损失列表中的平均梯度更新神经网络?
【发布时间】:2022-10-02 16:10:13
【问题描述】:

我有一个基于 REINFORCE 算法的玩具强化学习项目(这里是PyTorch\'s implementation),我想添加批量更新。在 RL 中,“目标”只能在“预测”完成后创建,因此标准批处理技术不适用。因此,我为每一集累积损失并将它们附加到一个列表l_losses,其中每个项目都是一个零维张量。我推迟致电.backward()optimizer.step(),直到经过一定数量的剧集以创建一种伪批次。

鉴于此损失列表,我如何让 PyTorch 根据其平均梯度更新网络?或者基于平均梯度的更新与平均损失的更新相同(我似乎在其他地方读过)?

我目前的方法是从torch.stack(l_losses) 创建一个新张量t_loss,然后运行t_loss = t_loss.mean()t_loss.backward()optimizer.step(),并将梯度归零,但我不确定这是否等同于我的意图?我也不清楚我是否应该在每个单独的损失上运行.backward(),而不是将它们连接到一个列表中(但坚持.step() 部分直到最后?

    标签: python deep-learning pytorch gradient-descent


    【解决方案1】:

    梯度是一种线性运算,因此平均梯度与梯度的平均值相同。

    拿一些示例数据

    import torch
    a = torch.randn(1, 4, requires_grad=True);
    b = torch.randn(5, 4);
    

    您可以存储所有损失并计算平均值,

    a.grad = None
    x = (a * b).mean(axis=1)
    x.mean().backward() # gradient of the mean
    print(a.grad)
    

    或者每次迭代计算反向传播以获得该损失对梯度的贡献。

    a.grad = None
    for bi in b:
        (a * bi / len(b)).mean().backward()
    print(a.grad)
    

    表现

    我不知道 pytorch 向后实现的内部细节,但我可以说

    (1) 将ratain_graph=Truecreate_graph=True 向后传递到backward() 后默认销毁图。

    (2)除了叶子张量,不保留梯度,除非你指定retain_grad

    (3) 如果您使用不同的输入对模型进行两次评估,您可以对单个变量执行反向传递,这意味着它们具有单独的图。这可以使用以下代码进行验证。

    a.grad = None
    # compute all the variables in advance
    r = [ (a * b / len(b)).mean() for bi in b ]
    for ri in r:
        # This depends on the graph of r[i] but the graph or r[i-1]
        # was already destroyed, it means that r[i] graph is independent
        # of r[i-1] graph, hence they require separate memory.
        ri.backward()  # this will remove the graph of ri
    print(a.grad)
    

    因此,如果您在每一集之后更新梯度,它将累积叶节点的梯度,这就是下一个优化步骤所需的所有信息,因此您可以丢弃该损失,从而释放资源以进行进一步计算。如果内存分配可以有效地将刚刚释放的页面用于下一次分配,我希望内存使用量减少,甚至可能更快地执行。

    【讨论】:

    • 不应该for bi in b: 然后实际使用bi?如果是这样,我注意到我得到了不同的渐变。
    • 没错,谢谢你的关注。
    • 谢谢。为了使这完全全面,我注意到,如果我修改您的代码以将 (a * bi).mean() 附加到列表中,torch.stack() 该列表和 .mean().backward() 那些结果,我也会得到相同的渐变,这很好。为了结束这个问题,因为所有这些都是等价的,这里在计算速度或某种三重危险方面是否有任何偏好?
    • 回复为对答案的编辑。
    猜你喜欢
    • 2018-09-01
    • 2021-01-13
    • 1970-01-01
    • 2015-07-11
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2020-06-03
    • 2011-08-24
    相关资源
    最近更新 更多