【问题标题】:Custom backward/optimization steps in pytorch-lightningpytorch-lightning 中的自定义后向/优化步骤
【发布时间】:2020-04-07 04:05:53
【问题描述】:

我想在 pytorch-lightning 中实现下面的训练循环(以伪代码形式阅读)。特殊之处在于,并不是每批都执行后向和优化步骤。

(背景:我正在尝试实现一个few-shots学习算法;虽然我需要在每一步都做出预测——forward方法 -- 我需要随机执行梯度更新 -- if- 块。

for batch in batches:
    x, y = batch
    loss = forward(x,y)

    optimizer.zero_grad()

    if np.random.rand() > 0.5:
        loss.backward()
        optimizer.step()

我提出的解决方案需要实现backwardoptimizer_step 方法,如下所示:

def backward(self, use_amp, loss, optimizer):
        self.compute_grads = False
        if np.random.rand() > 0.5:
            loss.backward()
            nn.utils.clip_grad_value_(self.enc.parameters(), 1)
            nn.utils.clip_grad_value_(self.dec.parameters(), 1)
            self.compute_grads = True
        return


    def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i, second_order_closure=None):
        if self.compute_grads:
            optimizer.step()
            optimizer.zero_grad()   
        return

注意:这样我需要在类级别存储一个compute_grads属性。

在 pytorch-lightning 中实现它的“最佳实践”方法是什么?有没有更好的使用钩子的方法?

【问题讨论】:

  • 这太模糊了。 forward 是什么?如果你不在loss 上运行.backward,那么你永远不会计算forward 参数的梯度(假设它是nn.Module 子类的一个实例)。在没有backward 的情况下运行forward 似乎是在浪费时间。 forward 可能会做的唯一持久的事情是更新任何批处理规范化层上的运行统计信息。

标签: pytorch


【解决方案1】:

这是一个很好的方法!这就是钩子的用途。

有一个新的回调模块也可能会有所帮助: https://pytorch-lightning.readthedocs.io/en/0.7.1/callbacks.html

【讨论】:

    猜你喜欢
    • 2022-06-19
    • 1970-01-01
    • 1970-01-01
    • 2021-12-27
    • 2021-08-25
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多