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

自定义PyTorch损失函数无法反向传播至模型参数的问题排查

问题

自定义损失函数代码如下:

class CustomIndicesEdgeAccuracyLoss(torch.nn.Module):
    def __init__(self, num_classes: int, selected_indices: list):
        super(CustomIndicesEdgeAccuracyLoss, self).__init__()
        self.num_classes = num_classes
        self.selected_indices = selected_indices

    def forward(self, input: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
        batch_size, num_classes, feature_size = input.shape
        selected_input = input[::, ::, self.selected_indices]
        selected_target = target[::, self.selected_indices]
        selected_preds = torch.argmax(selected_input, dim=1)
        edge_acc = torch.eq(selected_preds, selected_target).sum()/torch.numel(selected_preds)
        loss = 1 - edge_acc
        loss.requires_grad = True

        return loss

该损失无法反向传播,模型参数梯度始终为0,无法完成更新。已知各变量形状:

input.shape: torch.Size([64, 3, 5])
target.shape:torch.Size([64, 5])
selected_input.shape: torch.Size([64, 3, 2]) 
selected_target.shape:torch.Size([64, 2])

原因分析

  • torch.argmax是不可导操作:它返回离散的类别索引,会直接切断计算图,导致梯度无法从损失回传到模型参数。
  • 手动设置loss.requires_grad = True无效:前面的计算已经破坏了梯度流,强行设置无法恢复梯度传播能力。
  • 基于硬准确率的损失是离散的:准确率是0/1判断的平均值,损失值为1-准确率,属于离散阶梯状数值,几乎没有有效梯度信号(多数情况下梯度为0),无法驱动参数更新。

修改方案

改用可导的交叉熵损失针对选中索引计算,保留连续梯度信号。修改后的代码如下:

class CustomIndicesEdgeAccuracyLoss(torch.nn.Module):
    def __init__(self, num_classes: int, selected_indices: list):
        super(CustomIndicesEdgeAccuracyLoss, self).__init__()
        self.num_classes = num_classes
        self.selected_indices = selected_indices
        self.ce_loss = torch.nn.CrossEntropyLoss()

    def forward(self, input: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
        # 提取选中索引对应的输入和目标
        selected_input = input[:, :, self.selected_indices]
        selected_target = target[:, self.selected_indices]
        
        # 调整形状适配CrossEntropyLoss输入要求:(N, C, d) -> (N*d, C),目标转为(N*d,)
        batch_size, num_classes, num_selected = selected_input.shape
        selected_input_flat = selected_input.permute(0, 2, 1).reshape(-1, num_classes)
        selected_target_flat = selected_target.reshape(-1)
        
        # 计算交叉熵损失
        loss = self.ce_loss(selected_input_flat, selected_target_flat)
        return loss

补充说明

  • 交叉熵损失具备可导性,能为模型提供连续梯度信号,保证参数正常更新。
  • 调整输入形状是因为CrossEntropyLoss默认接受(样本数, 类别数)格式,需将选中的每个位置视为独立样本计算损失再平均。
  • 若需监控准确率,可在训练过程中单独计算该指标,但不要用它直接作为损失函数。

内容的提问来源于stack exchange,提问作者theabc50111

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 20:34:57