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

自定义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)
---&gt; 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):
-&gt; 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)
---&gt; 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是类别数,属于分类任务。

错误原因

  1. 任务类型不匹配:原AFL损失函数是为图像分割任务设计的,假设输入y_pred和y_true是4维张量([batch, class, H, W]),但你的任务是分类任务,输出是2维张量([batch, class]),代码中[:,0,:,:]这类针对空间维度的索引在分类任务中不存在,直接调用会导致维度错误。
  2. 标签格式错误:调用时labels.view(1, -1)将标签转为[1, 64]格式,而y_pred是[64,7],两者维度不匹配;同时原代码期望y_true是one-hot编码格式,但你的标签可能是类别索引([64]),直接相乘会引发维度不兼容报错。
  3. 多分类适配缺失:原代码仅处理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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 01:10:27