在PyTorch中用另一模型计算的损失训练模型时参数不更新如何解决?
问题原因
- 核心问题:梯度传递链路被类型强制转换截断
你创建state张量时指定了dtype=torch.long,但模型A输出的调度方案pred是浮点类型,将浮点张量赋值给长整型张量时会触发不可导的类型强制转换,直接断了从损失回传到模型A的梯度路径,导致A的参数梯度始终为空,自然不会更新。 - 代码笔误:你写的
optimizer.zero_gard()是拼写错误,正确写法为optimizer.zero_grad(),如果实际代码存在该错误会直接导致梯度没有清零逻辑,也会影响参数更新。 - 模型B没有正确冻结:仅给B设置学习率为0不是最优方案,B的参数如果没有显式关闭
requires_grad,依然会产生不必要的梯度计算,极端情况下还会干扰梯度回传。 - 设备不一致:如果
state在CPU创建、deploy在GPU上,跨设备赋值也会触发张量拷贝,截断梯度链路。 - 学习率设置不当:如果优化器学习率设置过小,会导致参数更新幅度极低,看起来像没有更新。
解决方法
1. 调整张量类型与设备,保留梯度
将state的类型改为和输入deploy一致的浮点类型,同时保证两者在同一设备上,避免类型转换、跨设备拷贝截断梯度。如果模型B需要输入离散的站点ID,建议将ID和连续的调度值拆分输入,不要拼到同一个张量里避免类型冲突。
2. 显式冻结模型B并开启推理模式
预训练好的B直接关闭所有参数的梯度更新,同时切换到eval模式,避免BN、Dropout等层的推理/训练行为不一致干扰评估结果。
3. 修正代码笔误,调整合理学习率
4. 可选:增加梯度裁剪避免梯度消失/爆炸
修改后参考代码
import torch A = modelA() B = modelB() # 冻结B的所有参数,切换到推理模式 B.eval() for param in B.parameters(): param.requires_grad = False # 建议给A设置合理的学习率,不要太小 optimizer = torch.optim.Adam(A.parameters(), lr=1e-4) def my_loss(deploy): shape = deploy.size() # 与deploy保持同类型、同设备,不强制转long state = torch.zeros( (shape[0], shape[1], 2 + shape[1]), dtype=deploy.dtype, device=deploy.device ) # 同类型赋值,梯度可正常传递 state[:, :, 2:] = torch.reshape(deploy, (shape[0], 1, shape[1])) # 站点ID转成和state一致的类型 state[:, :, 0] = torch.arange(0, shape[1], dtype=deploy.dtype, device=deploy.device) state = torch.reshape(state, (-1, 2 + shape[1])) eval = B(state) eval = torch.reshape(eval, (shape[0], shape[1])) return torch.mean(eval) # 训练流程 EPOCHS = 100 for epoch in range(EPOCHS): A.train() for batch_idx, (x, useless_y) in enumerate(dataloader): optimizer.zero_grad() # 修正拼写错误 pred = A(x) loss = my_loss(pred) loss.backward() # 调试时可打开下面的代码,确认A的参数存在梯度 # print(next(A.parameters()).grad) # 梯度裁剪避免梯度爆炸 torch.nn.utils.clip_grad_norm_(A.parameters(), max_norm=1.0) optimizer.step()
内容的提问来源于stack exchange,提问作者DDullahan
相关产品推荐
相关产品推荐

