PyTorch中损失函数无requires_grad=True致反向传播报错求助
PyTorch反向传播报错:RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn
可复现错误代码
import numpy as np from numpy import linalg as LA import torch import torch.optim as optim import torch.nn as nn def func(x,pars): a = pars[0] b = pars[1] c = pars[2] d = pars[3] x = x.int() H = torch.tensor([[a,b,1],[2,3,c],[4,d,7]]) eigenvalues, eigenvectors = np.linalg.eigh(H) trans_freq = eigenvalues[x] return torch.tensor(trans_freq) x_index = torch.tensor([1,2]) y_vals = torch.tensor([0.5,12]) params = torch.tensor([1.,2.,3.,4.]) params.requires_grad=True opt = optim.SGD([params], lr=100) mse_loss = nn.MSELoss() for i in range(10): opt.zero_grad() loss = mse_loss(func(x_index,params),y_vals) print(x_index.requires_grad) print(params.requires_grad) print(y_vals.requires_grad) print(loss.requires_grad) loss.backward() opt.step() print(loss)
运行输出
False True False False
触发错误
执行loss.backward()时出现:
RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn
原因分析
- 计算图被打断:代码混用了NumPy与PyTorch操作,
np.linalg.eigh(H)会将PyTorch张量H转换为NumPy数组,而NumPy操作不会参与PyTorch梯度计算图的构建,直接切断了梯度传递路径。 - 返回张量无梯度追踪:从NumPy数组
eigenvalues得到trans_freq后,用torch.tensor()新建的张量完全脱离了params所在的计算图,没有绑定梯度计算函数(grad_fn),导致后续计算的loss无法追踪梯度,因此loss.requires_grad为False。
解决方法
全程使用PyTorch原生操作,确保计算图完整:
- 将
np.linalg.eigh替换为PyTorch对应函数torch.linalg.eigh; - 调整索引张量类型为
long(PyTorch索引要求使用长整型张量); - 直接返回PyTorch计算得到的张量,无需用
torch.tensor()重新包裹。
修改后的代码:
import torch import torch.optim as optim import torch.nn as nn def func(x, pars): a = pars[0] b = pars[1] c = pars[2] d = pars[3] x = x.long() # 索引需使用长整型张量 # 构造PyTorch张量时显式开启梯度追踪,确保梯度传递 H = torch.tensor([[a, b, 1.], [2., 3., c], [4., d, 7.]], requires_grad=True) eigenvalues, eigenvectors = torch.linalg.eigh(H) trans_freq = eigenvalues[x] return trans_freq # 直接返回PyTorch张量,保留计算图信息 x_index = torch.tensor([1,2]) y_vals = torch.tensor([0.5,12.]) params = torch.tensor([1.,2.,3.,4.], requires_grad=True) opt = optim.SGD([params], lr=100) mse_loss = nn.MSELoss() for i in range(10): opt.zero_grad() pred = func(x_index, params) loss = mse_loss(pred, y_vals) print(f"Epoch {i+1}, Loss: {loss.item()}") loss.backward() opt.step()
内容的提问来源于stack exchange,提问作者Silviu
相关产品推荐
相关产品推荐

