Update net.parameters() without .data

Viewed 83

Is there any way to update the net parameters with some other tensors that carry gradients as well?

I want to do something like the following:

grads = torch.autograd.grad(loss, net.parameters(), 
                                    create_graph=True) 

for param gi in zip(net.parameters(), grads): 
       param -= eps * gi

And I want each param to carry the grad_fn of gi.

1 Answers

You can do this by wrapping the whole loop with torch.no_grad():

grads = torch.autograd.grad(loss, net.parameters(), create_graph=True) 
with torch.no_grad():
    for param, gi in zip(net.parameters(), grads):
        param -= eps*gi

Alternatively you can use the in-place copy_() on param's data property:

grads = torch.autograd.grad(loss, net.parameters(), create_graph=True) 
for param, gi in zip(net.parameters(), grads): 
    param.data.copy_(param.data - eps*gi)

As far as I have tested, both methods update the parameters the same way.


I haven't found any way to copy the grad_fn property though. As a workaround you could copy to gi instead of param, this will overwrite the values of grads with what would have been the new parameters of the model:

grads = torch.autograd.grad(loss, net.parameters(), create_graph=True) 
for param, gi in zip(net.parameters(), grads): 
    gi.data.copy_(param.data - eps*gi)

In case you need to keep grads unmodified, just clone it before the entering the loop.

Related