自定义AFL损失函数遇RuntimeError:张量维度不匹配问题求助
问题分析与解决:自定义AFL损失函数维度不匹配报错
问题背景
用户在修改自定义AFL损失函数时遇到维度不匹配报错,相关信息如下:
现有AFL损失函数代码
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] y_pred = torch.clamp(y_pred, self.epsilon, 1. - self.epsilon) cross_entropy = -y_true * torch.log(y_pred) # 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
调用代码
citerion = AFL() loss=criterion(outputs,labels.view(1, -1)) loss.backward() optimizer.step()
报错信息
RuntimeError Traceback (most recent call last) <ipython-input-178-fe998ec13e82> in <module> 20 print(outputs.shape) 21 print(labels.shape) ---> 22 loss=criterion(outputs,labels) 23 #model.zero_grad() 24 D:\software\anaconda3\lib\site-packages\torch\nn\modules\module.py in _call_impl(self, *input, **kwargs) 1128 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks 1129 or _global_forward_hooks or _global_forward_pre_hooks): -> 1130 return forward_call(*input, **kwargs) 1131 # Do not call functions when jit is used 1132 full_backward_hooks, non_full_backward_hooks = [], [] <ipython-input-173-0ace81c245e8> in forward(self, y_pred, y_true) 71 #y_pred=y_pred.size()[1] 72 y_pred = torch.clamp(y_pred, self.epsilon, 1. - self.epsilon) ---> 73 cross_entropy = -y_true * torch.log(y_pred) 74 75 # Calculate losses separately for each class, only suppressing background class RuntimeError: The size of tensor a (64) must match the size of tensor b (7) at non-singleton dimension 1
网络输出维度
网络最终输出为torch.Size([64,7]),其中64是batch size,7是类别数,属于分类任务。
错误原因
- 任务类型不匹配:原AFL损失函数是为图像分割任务设计的,假设输入
y_pred和y_true是4维张量([batch, class, H, W]),但你的任务是分类任务,输出是2维张量([batch, class]),代码中[:,0,:,:]这类针对空间维度的索引在分类任务中不存在,直接调用会导致维度错误。 - 标签格式错误:调用时
labels.view(1, -1)将标签转为[1, 64]格式,而y_pred是[64,7],两者维度不匹配;同时原代码期望y_true是one-hot编码格式,但你的标签可能是类别索引([64]),直接相乘会引发维度不兼容报错。 - 多分类适配缺失:原代码仅处理2类(背景0类、前景1类),但你的任务是7分类,逻辑上无法直接复用。
解决方案
修改AFL损失函数适配分类任务,同时修正标签处理逻辑:
修改后的AFL损失函数代码
class AFL(nn.Module): def __init__(self, delta=0.7, gamma=2., epsilon=1e-07, num_classes=7): super(AFL, self).__init__() self.delta = delta self.gamma = gamma self.epsilon = epsilon self.num_classes = num_classes # 适配多分类 def forward(self, y_pred, y_true): # y_pred shape: [batch_size, num_classes] # y_true shape: [batch_size] (类别索引) 或 [batch_size, num_classes] (one-hot编码) y_pred = torch.clamp(y_pred, self.epsilon, 1. - self.epsilon) # 自动将类别索引转换为one-hot编码 if y_true.dim() == 1: y_true = torch.nn.functional.one_hot(y_true, num_classes=self.num_classes).float() cross_entropy = -y_true * 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 # 前景类损失:对所有非背景类的交叉熵求和 fore_ce = cross_entropy[:, 1:].sum(dim=1) fore_ce = self.delta * fore_ce # 计算总损失 loss = torch.mean(back_ce + fore_ce) return loss
修改后的调用代码
criterion = AFL(num_classes=7) # 指定类别数 # labels为类别索引,shape [64],无需额外reshape loss = criterion(outputs, labels) loss.backward() optimizer.step()
关键修改点
- 新增
num_classes参数,适配多分类场景 - 自动处理标签格式:将类别索引转为one-hot编码,兼容两种输入格式
- 移除分割任务的空间维度索引,适配分类任务的2维输出
- 调整前景损失计算逻辑,对所有非背景类的损失求和,符合7分类需求
- 简化损失合并逻辑,直接对背景和前景损失求和后取均值,提升效率
内容的提问来源于stack exchange,提问作者shel coop
相关产品推荐
相关产品推荐

