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

PyTorch多任务训练张量尺寸不匹配与交叉熵损失形状问题求解

问题概述

从TensorFlow迁移到PyTorch训练多任务模型时,首先触发尺寸不匹配警告:

UserWarning: Using a target size (torch.Size([400])) that is different to the input size (torch.Size([400, 1])). This will likely lead to incorrect results due to broadcasting. Please ensure they have the same size

尝试对目标张量执行unsqueeze(1)修复尺寸问题后,又因交叉熵损失对目标形状的要求,触发多任务维度适配错误。
附最小可复现代码:

import torch
import torch.nn as nn
import torch.optim as optim 
from torch.utils.data import Dataset, DataLoader, TensorDataset
import torch.nn.functional as F


X1 = torch.randn(400, 1, 9999)
X2 = torch.randn((400,1, 9999))
aux1 = torch.randn(400,1)
aux2 = torch.randn(400,1)
aux3 = torch.randn(400,1)
y1 = torch.rand(400,)
y2 = torch.rand(400,)
y3 = torch.rand(400,)


class MultiTaskDataset:
    def __init__(self, 
                 amplitude, 
                 phase, 
                 weight,
                 temperature,
                 humidity,
                 shelf_life_clf,
                 shelf_life_pred,
                 thickness_pred
                 ):
        self.amplitude = amplitude
        self.phase = phase
        self.weight = weight
        self.temperature = temperature
        self.humidity = humidity
        self.shelf_life_clf = shelf_life_clf
        self.shelf_life_pred = shelf_life_pred
        self.thickness_pred = thickness_pred
        

    def __len__(self):
        return self.amplitude.shape[0]

    def __getitem__(self, idx):
        #inputs
        amplitude = self.amplitude[idx]
        phase = self.phase[idx]
        weight = self.weight[idx]
        temperature = self.temperature[idx]
        humidity = self.humidity[idx]
        
        #outputs
        shelf_life_clf = self.shelf_life_clf[idx]
        shelf_life_reg = self.shelf_life_pred[idx]
        thickness_pred = self.thickness_pred[idx]
        
        return ([torch.tensor(amplitude, dtype=torch.float32),
                torch.tensor(phase, dtype=torch.float32),
                torch.tensor(weight, dtype=torch.float32),
                torch.tensor(temperature, dtype=torch.float32),
                torch.tensor(humidity, dtype=torch.float32)],
                [torch.tensor(shelf_life_clf, dtype=torch.long),
                torch.tensor(shelf_life_reg, dtype=torch.float32),
                torch.tensor(thickness_pred, dtype=torch.float32)])


# train loader
dataset = MultiTaskDataset(X1, X2, aux1, aux2, aux3, 
                           y1,y2,y3)
train_loader = DataLoader(dataset, batch_size=512, shuffle=True, num_workers=0)


class MyModel(nn.Module):
    def __init__(self):
        super(MyModel, self).__init__()
        self.features_amp = nn.Sequential(
            nn.LazyConv1d(1, 3, 1),
        )
        self.features_phase = nn.Sequential(
            nn.LazyConv1d(1, 3, 1),
        )
        
        
        self.backbone1 = nn.Sequential(
            nn.LazyConv1d(64,3,1),
            nn.LazyConv1d(64,3,1),
            nn.AvgPool1d(3),
            nn.Dropout(0.25),
        )
        
        self.backbone2 = nn.Sequential(
            nn.Conv1d(64, 32,3,1),
            nn.Conv1d(32, 32,3,1),
            nn.AvgPool1d(3),
            nn.Dropout(0.25),
        )
        
        self.backbone3 = nn.Sequential(
            nn.Conv1d(32, 16,3,1),
            nn.Conv1d(16, 16,3,1),
            nn.AvgPool1d(3),
            nn.Dropout(0.25),
        )
        
        
        self.classifier = nn.LazyLinear(2)
        self.shelf_life_reg = nn.LazyLinear(1)
        self.thickness_reg = nn.LazyLinear(1)

    def forward(self, x1, x2, aux1, aux2, aux3):
        x1 = self.features_amp(x1)
        x2 = self.features_phase(x2)
                                                                                                                                                                                                                                                
        x1 = x1.view(x1.size(0),-1)                                                                                                                                                                                                                 
                                                                                                                                                                                                                                                     
        x2 = x2.view(x2.size(0),-1)                                                                                                                                                                                                                                                    
        x = torch.cat((x1, x2), dim=-1)

        x = x.unsqueeze(1)
        x = self.backbone1(x)

        x = torch.flatten(x, start_dim=1, end_dim=-1)

 
        x = torch.cat([x, aux1, aux2, aux3], dim=-1)
        
        
        shelf_life_clf = self.classifier(x)     
        shelf_life_reg = self.shelf_life_reg(x)
        thickness_reg = self.thickness_reg(x)
        return (shelf_life_clf,
                shelf_life_reg,
                thickness_reg)


model = MyModel()

optimizer = optim.Adam(model.parameters(), lr=0.003)

criterion1 = nn.CrossEntropyLoss()
criterion2 = nn.MSELoss()
criterion3 = nn.MSELoss()


def train(epoch):
    model.train()
    arr_loss = []
    for batch_idx, (data, target) in enumerate(train_loader):
        clf, reg1, reg2 = target

        if torch.cuda.is_available():
            device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
            data = [data[i].cuda() for i in range(len(data))]
            target = [target[i].cuda() for i in range(len(target))]
            model.to(device)
  
        optimizer.zero_grad()
        output1, output2, output3 = model(*data)
        
        #losses
        loss = criterion1(output1, target[0].long())
        loss1 = criterion2(output2, target[1].float())
        loss2 = criterion3(output3, target[2].float())
        loss = loss + loss1 + loss2
        
        loss.backward()
        optimizer.step()

        print('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format(
                epoch, (batch_idx + 1) * len(data), len(train_loader.dataset),
                100. * (batch_idx + 1) / len(train_loader), loss.data))
        arr_loss.append(loss.data)
    return arr_loss

def averaged_accuracy(outputs, targets):
    assert len(outputs) != len(targets), "number of outputs should equal the number of targets"
    accuracy = []
    for i in range(len(outputs)):
        _, predicted = torch.max(output1.data, 1)
        total += target[0].size(0)
        correct += (predicted == target[0]).sum()
        acc = correct / total *100
        accuracy.append(acc)
    return torch.mean(accuracy)


optimizer = optim.Adam(model.parameters(), lr=0.00003)

criterion1 = nn.CrossEntropyLoss()
criterion2 = nn.MSELoss()
criterion3 = nn.MSELoss()


n_epochs = 10

for epoch in range(n_epochs):
    train(epoch)
问题根因
  • 尺寸不匹配警告来源:两个回归任务头输出形状为(batch_size, 1),但传入MSE损失的回归目标形状为(batch_size,),维度差1触发广播机制警告。
  • 交叉熵报错来源:nn.CrossEntropyLoss要求分类标签输入形状为(batch_size,),若对所有目标统一执行unsqueeze(1),分类标签会变成(batch_size, 1),不符合损失函数输入要求。
  • 附带代码bug:自定义的averaged_accuracy函数存在断言逻辑写反、变量未初始化、引用未传入的全局变量等问题,无法正常运行;示例中分类标签用torch.rand生成浮点数,不符合分类任务标签为类别整数索引的要求;CUDA设备判断逻辑写在训练循环内,每个batch重复判断浪费性能。
修复方案
  • 维度匹配二选一即可,优先推荐调整模型输出,无需改动数据加载逻辑:
    • 方案1(推荐):在模型forward方法返回结果时,压缩回归输出的最后一维,让回归输出形状和目标形状统一为(batch_size,):
      # 修改MyModel的forward返回部分
      return (
          shelf_life_clf,
          shelf_life_reg.squeeze(-1),
          thickness_reg.squeeze(-1)
      )
      
    • 方案2:若不修改模型输出,计算回归损失时单独给回归目标加维度,分类目标保持原形状不动:
      # 训练循环中计算损失部分
      loss = criterion1(output1, target[0].long()) # 分类目标保持(batch_size,)形状
      loss1 = criterion2(output2, target[1].float().unsqueeze(1)) # 仅回归目标加维度匹配输出
      loss2 = criterion3(output3, target[2].float().unsqueeze(1))
      
  • 修复准确率计算函数:
    def averaged_accuracy(outputs, targets):
        assert len(outputs) == len(targets), "输出数量需与目标数量一致"
        clf_output = outputs[0]
        clf_target = targets[0]
        _, predicted = torch.max(clf_output.data, 1)
        total = clf_target.size(0)
        correct = (predicted == clf_target).sum()
        acc = correct / total * 100
        return acc
    
  • 修复测试数据合法性:示例中分类标签为随机浮点数,测试时需替换为合法的整数类别标签:
    # 原y1 = torch.rand(400,) 替换为
    y1 = torch.randint(0, 2, (400,)) # 生成0、1两类的分类标签
    
  • 性能优化:将CUDA设备判断逻辑移到训练循环外,避免每个batch重复判断:
    # 训练前提前定义设备
    device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
    model.to(device)
    # 训练循环内直接复用device做数据迁移即可
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 20:39:18