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

PyTorch batch_size不匹配问题:为何输出维度为56180?

问题分析与解决

核心错误原因

你的模型输出维度异常的根源在特征展平步骤和全连接层输入维度定义:

  1. 卷积后的输出形状是torch.Size([100, 20, 53, 53]),其中100是batch size,20是通道数,53×53是特征图尺寸。但你用了x.view(-1, x.size(0)),这会把维度顺序搞反——把batch size(100)当成了特征维度,把所有特征元素(20×53×53=56180)当成了新的batch维度,导致输出形状变成[56180,100],后续全连接层处理后最终输出维度自然变成[56180,1],和target的[100,1]不匹配。
  2. 全连接层self.fc1 = nn.Linear(100, 64)的输入维度定义错误,卷积后的特征展平后每个样本的特征数是20×53×53=56180,不是100。

修正后的代码

模型类修正

class Net(nn.Module):
    def __init__(self):
        super(Net, self).__init__()

        #input image 227x227x3
        self.conv1 = nn.Conv2d(3, 10, kernel_size=5)
        self.conv2 = nn.Conv2d(10, 20, kernel_size=5)
        self.conv2_drop = nn.Dropout2d()

        # 修正全连接层输入维度:20*53*53=56180
        self.fc1 = nn.Linear(56180, 64)
        self.fc3 = nn.Linear(64, 32)
        self.fc6 = nn.Linear(32, 1)

    def forward(self, x):
        x = F.relu(F.max_pool2d(self.conv1(x), 2))
        x = F.relu(F.max_pool2d(self.conv2_drop(self.conv2(x)), 2))
        # 修正展平方式:先保留batch维度,再展平特征
        x = x.view(x.size(0), -1)
        x = F.relu(self.fc1(x))
        x = F.dropout(x, training=self.training)
        x = self.fc3(x)
        x = F.dropout(x, training=self.training)
        x = self.fc6(x)
        return x

训练函数修正

二分类任务用F.cross_entropy不合适,该函数针对多分类场景,二分类应使用F.binary_cross_entropy_with_logits,同时target需要转为float类型:

def train(model, train_loader, optimizer):
    model.train()
    for batch_idx, (data, target) in enumerate(train_loader):
        data, target = data.to(DEVICE), target.to(DEVICE).float()  # 转为float类型
        optimizer.zero_grad()
        output = model(data)
        target = target.unsqueeze(-1)
        # 用二分类专用损失函数
        loss = F.binary_cross_entropy_with_logits(output, target)

        loss.backward()
        optimizer.step()

验证修正后的维度

修正后,各层输出形状应为:

  • 输入:torch.Size([100, 3, 227, 227])
  • conv1后:torch.Size([100, 10, 111, 111])
  • conv2后:torch.Size([100, 20, 53, 53])
  • 展平后:torch.Size([100, 56180])
  • fc1后:torch.Size([100, 64])
  • 最终输出:torch.Size([100, 1])
    此时输出batch size和target的100完全匹配,不会再报错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 23:31:07