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)

浙公网安备 33010602011771号