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
问题分析
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]相乘无法得到对应类别的交叉熵。
- 用
索引越界错误
fore_ce = cross_entropy[:,1,:,:]中,cross_entropy是二维张量[32,7],却使用了四维索引,即使交叉熵计算正确,这行也会报错。
未处理类别索引到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
相关产品推荐
相关产品推荐

