【问题标题】:PyTorch: How to check if some weights are not changed during training?PyTorch:如何检查训练期间是否有一些权重没有改变?
【发布时间】:2020-09-21 22:48:55
【问题描述】:

如何检查在 PyTorch 中训练期间某些权重是否未更改?

据我了解,一种选择可以是在某些时期转储模型权重并检查它们是否在权重上迭代,但也许有一些更简单的方法?

【问题讨论】:

  • 你可以设置一个钩子来检查梯度,然后你可以通过计算零的数量来跟踪一个计数器是否更新了梯度值。
  • @AnuragReddy 的例子会很棒。
  • 检查答案。 @mrgloom 如果有问题或令人困惑,请随时询问或纠正我。我也可以这样从你身上学到一些东西:)

标签: pytorch


【解决方案1】:

有两种方法可以解决这个问题:

第一

        for name, param in model.named_parameters():
            if 'weight' in name:
                temp = torch.zeros(param.grad.shape)
                temp[param.grad != 0] += 1
                count_dict[name] += temp

此步骤在您在培训模块中的loss.backward() 步骤之后进行。 count_dict[name] 字典跟踪梯度更新。你可以在训练开始之前这样初始化它:

    for name, param in model.named_parameters():
        if 'weight' in name:
            count_dict[name] = torch.zeros(param.grad.shape)

现在还有一种方法是注册一个钩子函数,然后创建钩子函数,您甚至可以在其中更新或修改渐变(如果需要)。这不是跟踪权重更新所必需的,但是如果你想对梯度做一些事情,它就派上用场了。 假设,我在这里随机稀疏渐变。

def hook_fn(grad):
    '''
    Randomly sparsify the gradients
    :param grad: Input gradient of the layer
    :return: grad_clone - the sparsified FC layer gradients
    '''
    grad_clone = grad.clone()
    temp = torch.cuda.FloatTensor(grad_clone.shape).uniform_()
    grad_clone[temp < 0.8] = 0
    return grad_clone

在这里我给模型一个钩子。

for name, param in model.named_parameters():
    if 'weight' in name:
            param.register_hook(hook_fn)

所以,这可能只是为您稀疏渐变,您可以通过这种方式在挂钩函数本身中跟踪渐变:

def hook_func(module, input, output):
    temp = torch.zeros(output.shape)
    temp[output != 0] += 1
    count_dict[module] += temp

虽然,我不建议这样做。这在可视化前向传递特征/激活的情况下通常很有用。而且,输入和输出可能会混淆,因为梯度和参数输入和输出是相反的。

【讨论】:

  • 似乎钩子并不总是可以称为stackoverflow.com/questions/63998318/…loss.backward() 之后使用register_hook 而不是.grad 的优缺点是什么?
  • 如果.grad 不是None 是否总是意味着权重会更新?
  • @mrgloom 回答您的第一个问题 - 因此,通常钩子在可视化中间激活(前向钩子)或向后钩子的情况下很有用,它们在训练自身时修改梯度很有用。假设你想对某些更新进行门控,或者有一些可以乘以或添加到梯度的参数或值,那么钩子是很好的。来到问题 2 - 如果.grad 不是None,我假设您设置了requires_grad=False,因为表示更新的是渐变本身的值。
  • 我的意思是,.grad 值可以不是 None 但同时它可以是0s 和其他值的混合。因此,只要值为 0,则表示突触将保持不变并且不会更新
  • 在你分享的post 中,我建议直接使用.grad 而不是使用钩子。因为,您不想修改渐变。是的,你可以在钩子中给出 no grad 选项(基本上在钩子函数中将所有渐变设置为 0):P.
猜你喜欢
  • 2021-03-28
  • 2018-12-08
  • 1970-01-01
  • 2018-01-03
  • 2020-08-30
  • 2021-08-06
  • 2019-08-10
  • 1970-01-01
  • 2022-10-07
相关资源
最近更新 更多