代码之家  ›  专栏  ›  技术社区  ›  oarfish

为什么这个函数参数在每次调用中都是相同的,尽管传递的值不同(在循环中创建闭包)

  •  0
  • oarfish  · 技术社区  · 7 年前

    我正在使用PyTorch并尝试注册模型参数上的钩子。下面的代码创建lambda函数来添加到每个模型参数中,因此我可以在hook中看到梯度属于哪个张量

    import torch
    import torchvision
    
    # define model and random train batch
    model = torchvision.models.alexnet()
    input = torch.rand(10, 3, 224, 224)   # batch of 10 images
    targets = torch.zeros(10).long()
    
    def grad_hook_template(param, name, grad):
        print(f'Receive grad for {name} w whape {grad.shape}')
    
    # add one lambda hook to each parameter
    for name, param in model.named_parameters():
        print(f'Register hook for {name}')
    
        # use a lambda so we can pass additional information to the hook, which should only take one parameter
        param.register_hook(lambda grad: grad_hook_template(param, name, grad))
    
    loss_fn = torch.nn.CrossEntropyLoss()
    optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
    optimizer.zero_grad()
    
    prediction = model(input)
    loss = loss_fn(prediction, targets)
    loss.backward()
    optimizer.step()
    

    结果是 name param 论据 grad_hook_template id ),但 grad

    我读书。 here 循环不创建新的作用域和闭包在Python中是词法性的,即 名称 param copy.copy() 变量?

    2 回复  |  直到 7 年前
        1
  •  0
  •   Cyphase    7 年前

    你遇到了 后期绑定闭包 . 变量 param name 在调用时查找,而不是在定义使用它们的函数时。在调用这些函数时, 名称 param 是循环中的最后一个值。要解决这个问题,您可以这样做:

    for name, param in model.named_parameters():
        print(f'Register hook for {name}')
        param.register_hook(lambda grad, name=name, param=param: grad_hook_template(param, name, grad))
    

    然而,我认为使用 functools.partial 这是正确的解决方案:

    from functools import partial
    
    for name, param in model.named_parameters():
        print(f'Register hook for {name}')
        param.register_hook(partial(grad_hook_template, name=name, param=param))
    

    你可以找到更多关于 late binding closures at the Common Gotchas page of the Hitchhiker's Guide to Python 以及 in the Python docs .

    def 关键词。

        2
  •  0
  •   oarfish    7 年前

    这是一种由 FAQ .

    • 使用 functools.partial 而不是 lambda
    • 使用lambdas的默认参数来捕获变量的值