使用PyTorch更新函数时,参数矩阵sub_param为何未更新?
问题分析:sub_param未更新的原因及解决方法
核心问题是代码中多余的loss.requires_grad_(True)调用:
- 当通过
loss_criterion(n1, sub_param)计算损失时,loss已经是计算图的一部分,自带完整的梯度传播路径(从输出的grad_fn=<MseLossBackward0>可验证)。 - 手动执行
loss.requires_grad_(True)会破坏原有计算图的关联,导致loss.backward()无法将梯度传递到sub_param,参数自然无法更新。
另外需要注意:如果n1的requires_grad为True,梯度会同时流向n1,若你的目标仅更新sub_param,需确保n1的梯度被禁用(比如用n1.detach()剥离计算图)。
修正后的代码如下:
import torch import torch.nn as nn import torch.optim as optim # 确保n1不参与梯度更新 n1 = n1.detach() sub_param = nn.Parameter(n2) optimizer = optim.SGD([sub_param], lr=0.1) loss_criterion = nn.MSELoss() for epoch in range(5): optimizer.zero_grad() loss = loss_criterion(n1, sub_param) # 移除多余的loss.requires_grad_(True) loss.backward() optimizer.step() print(loss.item()) # 此时loss会逐步下降,参数也会被更新
内容的提问来源于stack exchange,提问作者An Min Su
相关产品推荐
相关产品推荐

