You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.01 00:27:24