torch计算梯度

#计算x的梯度
import torch
x = torch.arange(4.0)
print(x)

x.requires_grad_(True)   #声明x需要计算梯度
print(x.grad)    #默认值是None

#定义y的值
y = 2 * torch.dot(x,x)
print(y)

y.backward()
print(x.grad)

#清空梯度
x.grad.zero_()

#计算新的梯度
y = x.sum()
y.backward()
print(x.grad)

import torch
x1 = torch.arange(5.1)
print(x1)
x1.requires_grad_(True)
print(x1.grad)
y1 = torch.dot(x1 , x1)
y1.backward()
print(x1.grad)

  

posted @ 2021-11-12 11:19  小猪猪。。。  阅读(528)  评论(0)    收藏  举报