PyTorch如何实现不直接使用网络输出的自定义损失函数
PyTorch自定义序列损失函数修复方案
你的代码无法运行的核心原因是计算图断裂+梯度追踪失效,全程混用numpy数组、普通Python浮点数和torch张量,直接切断了反向传播的梯度路径,同时废弃的VariableAPI也会引发梯度异常。
现有代码核心问题
- 用
.detach().numpy()将带梯度的网络输出转为numpy数组计算YiPlus1,再转回张量时已经和原计算图完全脱钩,梯度无法回传到LSTM、全连接层的可训练参数。 self.maxRes用普通Python浮点数/numpy值存储,参与损失计算时会直接导致损失张量丢失requires_grad属性,无法调用.backward()完成反向传播。- PyTorch 0.4版本后
Variable已经和Tensor完全合并,手动封装Variable容易引发梯度追踪异常,属于废弃写法。 - 手动写for循环累加损失的逻辑冗余,用torch内置算子实现更高效,也不会破坏计算图结构。
修正后可运行实现
首先注意:你的自定义函数f(x)必须全部用PyTorch内置算子实现,不能混用numpy操作,否则会再次切断计算图。
import torch import torch.nn as nn # 替换为你自己的f(x)逻辑,全部使用torch API,禁止转numpy def f_torch(x): # 示例:如果原逻辑是对输出做维度调整+变换,直接用torch实现 return x.squeeze() class YourModel(nn.Module): # 原有模型初始化逻辑保留即可 def forward(self, x, y, hidden): c_0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size, device=x.device) y = y.reshape(y.shape[0], 1, 1) tmp = torch.cat((x, y), 2) output, (hn, cn) = self.lstm(tmp, (hidden, c_0)) out = self.fc(output) return out, hn def my_loss(self, target_seq, current_max_f): # 直接用torch算子计算平均差值,无需手动循环 return torch.mean(current_max_f - target_seq) def _train(self): num_epochs = 10 num_iteration = 10 save_loss_global = [] save_loss_epoch = [] for epoch in range(num_epochs): print("NOUVELLE EPOCH") X_train, Y_train = donneesAleatoires() # 用torch张量初始化最大值,不用普通Python数值 self.maxRes = torch.tensor(0.0, device=X_train.device) self.hidden = torch.zeros(self.num_layers, 1, self.hidden_size, device=X_train.device) # 用列表存储历史f(x)张量,比反复cat效率更高 tabY = [] first_y = Y_train[0].squeeze() tabY.append(first_y) self.maxRes = torch.max(self.maxRes, first_y) for iteration in range(num_iteration): x_i = X_train[0].reshape(x_i.shape[0], 1, x_i.shape[1]) y_i = Y_train[0] outputs, self.hidden = self(x_i, y_i, self.hidden) self.optimizer.zero_grad() # 禁止detach、禁止转numpy,直接用torch版f计算 YiPlus1 = f_torch(outputs).reshape(1) tabY.append(YiPlus1) # 最大值不需要参与梯度计算,detach后更新即可 self.maxRes = torch.max(self.maxRes, YiPlus1.detach()) target_seq = torch.stack(tabY) loss = self.my_loss(target_seq, self.maxRes) loss.backward() # 下一轮输入detach截断计算图,避免显存泄漏 X_train = outputs.detach() Y_train = YiPlus1.detach() self.optimizer.step() save_loss_global.append(loss.item()) if iteration == num_iteration - 1: save_loss_epoch.append(loss.item()) print(X_train)
关键注意事项
- 如果原f(x)是用numpy实现的,把所有numpy操作替换为对应torch操作即可,两者API几乎一致,不需要重写逻辑。
- 最大值项不需要参与梯度计算,更新
self.maxRes时必须对新的f(x)值做.detach(),否则梯度会向增大最大值的方向偏移,不符合你要计算「f(x)和max(f(x))差值」的设计目标。 - 序列拼接优先用列表存张量+最后
torch.stack()的方式,每一步迭代都调用torch.cat()会反复复制张量,训练速度会随迭代次数增加越来越慢。 - 如果训练目标是让f(x)尽可能大,也可以直接用
-torch.mean(target_seq)作为损失,去掉最大值项可以避免梯度计算的额外开销,收敛效果基本一致。
内容的提问来源于stack exchange,提问作者Alex
相关产品推荐
相关产品推荐

