如何在PyTorch中变换神经网络输出且保证训练正常进行?
问题原因
你遇到的报错和梯度断裂问题根源出在自定义的transform_torch实现逻辑上:
- 你手动创建了独立的leaf张量
new_tensor并设置requires_grad=True,PyTorch默认禁止对需要求导的leaf节点执行inplace赋值操作(也就是new_tensor[i] = xxx这类写法),这是直接触发报错的原因 - 即使没有报错,这种独立创建张量再赋值的写法也会断开和原
predictions张量的计算图连接,梯度无法回传到神经网络的参数,加了torch.no_grad()之后更是完全切断梯度,自然训练失效
解决方案
推荐两种修复方式,都可以保留完整计算图,同时避免报错:
最优方案:矢量化操作(无循环,速度更快)
你的翻转-累加-翻转操作本身支持批量处理,完全不需要写循环,直接对整个输入张量操作即可,代码如下:
def transform_torch(predictions): # 如果你的单样本计算维度不是最后一维,把-1改成对应的维度编号即可 return torch.flip(torch.cumsum(torch.flip(predictions, dims = [-1]), dim = -1), dims = [-1])
这种写法没有任何inplace操作,自动保留输入的dtype、device和计算图,梯度可以正常反向传播。
兼容方案:列表存储结果再拼接(适合需要逐样本自定义逻辑的场景)
如果你的实际逻辑必须要逐样本循环处理,不要手动创建张量再赋值,而是把每个样本的计算结果存入列表,最后用torch.stack拼接,代码如下:
def transform_torch(predictions): res_list = [] for i in range(len(predictions)): a = predictions[i] transformed = torch.flip(torch.cumsum(torch.flip(a, dims = [0]), dim = 0), dims = [0]) res_list.append(transformed) return torch.stack(res_list)
额外注意事项
你的核心训练代码里scheduler.step()的位置有误,常规训练流程的顺序应该是:
- 计算损失
loss.backward()反向传播计算梯度optimizer.step()更新参数scheduler.step()更新学习率(如果是按batch更新的scheduler)
你当前把scheduler.step()放在反向传播和优化器更新之前,会导致学习率更新逻辑错误,也会影响训练效果。
内容的提问来源于stack exchange,提问作者Galen BlueTalon
相关产品推荐
相关产品推荐

