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

PyTorch Lightning多任务模型训练报错:目标尺寸不匹配RuntimeError

解决PyTorch Lightning多任务分类模型的RuntimeError:Expected target size [20, 2], got [20]

这个错误核心是损失函数的输入维度不匹配:模型输出的张量形状为[20,2](20为batch size,2为分类类别数),但传入损失函数的标签(target)是形状[20]的一维张量,两者维度无法兼容。

常见问题定位

  • 损失函数选型错误:若误用BCEWithLogitsLoss(适用于多标签/二分类的独热标签场景)替代CrossEntropyLoss(适用于单标签多分类的类别索引场景),会要求标签为二维独热张量,和现有一维标签不匹配。
  • 多任务标签处理疏漏:多任务场景中,若未按任务拆分标签,或未对特定任务的标签做维度转换,会导致某任务的标签维度和对应模型头部输出不匹配。
  • 模型头部输出维度错误:任务头部的输出神经元数量设置错误(比如二分类误设为1而非2),导致输出张量维度和预期不符。

具体解决方法

1. 匹配损失函数与标签格式

  • 若为单标签多分类(每个样本仅属于一个类别):使用CrossEntropyLoss,标签保持一维类别索引即可,无需独热编码。
    # 正确示例
    loss_fn = nn.CrossEntropyLoss()
    # 模型输出logits: [20,2],标签target: [20]
    loss = loss_fn(model_output, target)
    
  • 若为多标签分类(每个样本可属于多个类别):使用BCEWithLogitsLoss,需将一维标签转为二维独热编码张量:
    # 正确示例
    loss_fn = nn.BCEWithLogitsLoss()
    # 转换标签为独热编码并转为float类型
    target = torch.nn.functional.one_hot(target, num_classes=2).float()
    # 模型输出logits: [20,2],标签target: [20,2]
    loss = loss_fn(model_output, target)
    

2. 多任务场景下的标签与输出对齐

多任务模型需确保每个任务的标签和对应头部输出维度匹配,拆分任务标签分别计算损失:

# 模型forward示例
def forward(self, x):
    features = self.backbone(x)
    task1_logits = self.task1_head(features)  # [20,2]
    task2_logits = self.task2_head(features)  # [20,2]
    return task1_logits, task2_logits

# 训练步骤示例
def training_step(self, batch, batch_idx):
    x, task1_target, task2_target = batch  # 两个任务的标签均为[20]的一维张量
    task1_logits, task2_logits = self(x)
    loss1 = nn.CrossEntropyLoss()(task1_logits, task1_target)
    loss2 = nn.CrossEntropyLoss()(task2_logits, task2_target)
    total_loss = loss1 + loss2
    self.log("train_loss", total_loss)
    return total_loss

3. 修正模型头部输出维度

检查任务头部的输出神经元数量,确保和任务类别数一致:

# 错误示例:二分类误设为1个输出神经元
self.task_head = nn.Linear(512, 1)
# 正确示例:二分类设置2个输出神经元
self.task_head = nn.Linear(512, 2)

快速验证技巧

从报错堆栈中找到触发错误的损失计算代码行,确认是哪个任务的损失维度不匹配,再对应上述方法逐一排查。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 17:16:10