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

PyTorch多分类(7类)AFL损失函数报错:维度1张量索引过多

问题与解决方案

问题描述

我正在使用PyTorch开发7分类任务的损失计算程序,实现了自定义AFL损失类,但运行时出现IndexError:too many indices for tensor of dimension 1。相关代码如下,其中y_pred形状为[32,7],y_true形状为[32],我希望分别计算每个类的back_ce,请问问题出在哪里?

class AFL(nn.Module):
    
    def __init__(self, delta=0.7, gamma=2., epsilon=1e-07):
        super(AFL, self).__init__()
        self.delta = delta
        self.gamma = gamma
        self.epsilon = epsilon

    def forward(self, y_pred, y_true):
        #y_pred=y_pred.size()[1]
        print(y_pred.shape) #[32,7]
        print(y_true.shape) #[32]
        y_pred = torch.clamp(y_pred, self.epsilon, 1. - self.epsilon)
        cross_entropy = np.empty(y_pred.shape)
        for i in range(len(y_pred)):
            for j in range(len(y_pred[i])):
                cross_entropy[i][j] = -y_true * torch.log(y_pred[i][j])
        #cross_entropy = -y_true * torch.log(y_pred[0][0]) #here i want to calculate cross_entropy for for each class
        
    # Calculate losses separately for each class, only suppressing background class
        back_ce = torch.pow(1 - y_pred[:,0], self.gamma) * cross_entropy[:,0]
        back_ce =  (1 - self.delta) * back_ce

        fore_ce = cross_entropy[:,1,:,:]
        fore_ce = self.delta * fore_ce

        loss = torch.mean(torch.sum(torch.stack([back_ce, fore_ce], axis=-1), axis=-1))
        return loss

报错信息:

back_ce = torch.pow(1 - y_pred[:,0], self.gamma) * cross_entropy[:,0]
IndexError: too many indices for tensor of dimension 1

问题分析

  1. numpy与PyTorch张量混合操作,交叉熵计算逻辑错误

    • 用np.empty创建numpy数组cross_entropy,但赋值时-y_true * torch.log(y_pred[i][j])得到的是[32]形状的PyTorch张量,强行赋值给numpy数组单个元素会引发隐式转换异常,导致cross_entropy维度或结构错乱,最终触发索引错误。
    • 嵌套循环计算交叉熵冗余且逻辑错误:y_true是类别索引,直接和单元素y_pred[i][j]相乘无法得到对应类别的交叉熵。
  2. 索引越界错误

    • fore_ce = cross_entropy[:,1,:,:]中,cross_entropy是二维张量[32,7],却使用了四维索引,即使交叉熵计算正确,这行也会报错。
  3. 未处理类别索引到one-hot编码的转换

    • y_true是[32]的类别索引张量,要计算每个类别的交叉熵,需先转为[32,7]的one-hot编码张量,才能与y_pred逐元素相乘。

修复后的代码

import torch
import torch.nn as nn

class AFL(nn.Module):
    
    def __init__(self, delta=0.7, gamma=2., epsilon=1e-07):
        super(AFL, self).__init__()
        self.delta = delta
        self.gamma = gamma
        self.epsilon = epsilon

    def forward(self, y_pred, y_true):
        # y_pred shape: [32,7], y_true shape: [32]
        y_pred = torch.clamp(y_pred, self.epsilon, 1. - self.epsilon)
        
        # 将y_true转换为one-hot编码,形状变为[32,7]
        y_onehot = nn.functional.one_hot(y_true, num_classes=y_pred.size(1)).float()
        
        # 计算每个类别的交叉熵,形状[32,7]
        cross_entropy = -y_onehot * torch.log(y_pred)
        
        # 计算背景类(第0类)损失
        back_ce = torch.pow(1 - y_pred[:, 0], self.gamma) * cross_entropy[:, 0]
        back_ce = (1 - self.delta) * back_ce
        
        # 计算前景类(第1到6类)损失,形状[32,6]
        fore_ce = cross_entropy[:, 1:]
        fore_ce = self.delta * fore_ce
        
        # 合并背景和前景损失,计算均值
        total_ce = torch.cat([back_ce.unsqueeze(1), fore_ce], dim=1)
        loss = torch.mean(torch.sum(total_ce, dim=1))
        
        return loss

关键修复点

  • 全程使用PyTorch张量操作,避免numpy与张量混合,保证计算效率和正确性。
  • 用one_hot将类别索引转换为one-hot编码,正确计算每个类别的交叉熵。
  • 修正fore_ce的索引为cross_entropy[:,1:],匹配二维张量维度。
  • 用cat替代stack高效合并损失,避免维度错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 05:50:11