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

PyTorch多输入CNN遇BatchNorm1d报错及维度问题排查

问题描述

我在复现多输入神经网络教程时,将原教程的PyTorch Lightning替换为原生PyTorch实现,已完成DataLoader与SimpleCNN模型的构建。测试单样本时遇到两个问题:

  • 执行self.batchnorm(img)时触发错误:ValueError: Expected more than 1 value per channel when training, got input size torch.Size([1, 16])
  • 移除BatchNorm1d后,出现张量拼接维度不匹配错误RuntimeError: Tensors must have same number of dimensions: got 2 and 1,已通过对tab_x执行unsqueeze解决维度问题,但BatchNorm1d的报错仍未解决。

相关代码

DataLoader定义

# Define loaders
from torch.utils.data import DataLoader
train_loader = DataLoader(train_set, batch_size=64, num_workers=2, drop_last=True, shuffle=True)
val_loader   = DataLoader(val_set,   batch_size=64, num_workers=2, drop_last=False, shuffle=False)
test_loader  = DataLoader(test_set,  batch_size=64, num_workers=2, drop_last=False, shuffle=False)

模型定义

def conv_block(input_size, output_size):
    block = nn.Sequential(
        nn.Conv2d(input_size, output_size, (3, 3)), nn.BatchNorm2d(output_size), nn.ReLU(), nn.MaxPool2d((2, 2)),
    )

    return block

class SimpleCNN(nn.Module):
    # Constructor
    def __init__(self):
        # Call parent constructor
        super().__init__()
        self.conv1 = conv_block(3, 16)
        self.conv2 = conv_block(16, 32)
        self.conv3 = conv_block(32, 64)

        self.ln1 = nn.Linear(64 * 26 * 26, 16)
        self.relu = nn.ReLU()
        self.batchnorm = nn.BatchNorm1d(16)
        self.dropout = nn.Dropout2d(0.5)
        self.ln2 = nn.Linear(16, 5)

        self.ln4 = nn.Linear(5, 10)
        self.ln5 = nn.Linear(10, 10)
        self.ln6 = nn.Linear(10, 5)
        self.ln7 = nn.Linear(10, 1)
    
    # Forward
    def forward(self, img, tab):
        img = self.conv1(img)
        img = self.conv2(img)
        img = self.conv3(img)
        img = img.reshape(img.shape[0], -1)
        img = self.ln1(img)
        img = self.relu(img)
        img = self.batchnorm(img)
        img = self.dropout(img)
        img = self.ln2(img)
        img = self.relu(img)

        tab = self.ln4(tab)
        tab = self.relu(tab)
        tab = self.ln5(tab)
        tab = self.relu(tab)
        tab = self.ln6(tab)
        tab = self.relu(tab)

        x = torch.cat((img, tab), dim=1)
        x = self.relu(x)

        return self.ln7(x)

测试代码

# Create the model
model = SimpleCNN()
img_x, tab_x, label_x = train_set[0]
print(img_x.shape, tab_x, label_x)
img_x = img_x.unsqueeze(dim=0)
output = model(img_x, tab_x)
output.shape

打印的张量维度信息

torch.Size([1, 16, 111, 111])
torch.Size([1, 32, 54, 54])
torch.Size([1, 64, 26, 26])
torch.Size([1, 43264])
torch.Size([1, 16])
报错原因分析
  1. BatchNorm1d报错原因:BatchNorm层在训练模式下,依赖当前批次的样本统计量(均值、方差)进行归一化。当测试单样本时,batch size为1,每个通道(此处为16个通道)仅对应1个样本值,无法计算有效方差(方差为0),因此PyTorch直接抛出错误。默认情况下模型处于train()模式,会强制要求每个通道的样本数大于1。
  2. 维度不匹配原因:tab_x初始为1维张量(如torch.Size([5])),而img经过处理后为2维张量(torch.Size([1,5])),拼接时维度不一致,因此需要给tab_x增加batch维度。
解决方案

解决BatchNorm1d单样本报错

有两种可行方案:

  • 方案1:测试时切换至评估模式:
    在测试单样本前调用model.eval(),此时BatchNorm会使用训练阶段学习到的移动均值和方差,不再依赖当前批次统计量,即使batch size为1也能正常运行。测试完成后若需继续训练,再切回model.train()。
    修改后的测试代码:
    model = SimpleCNN()
    img_x, tab_x, label_x = train_set[0]
    img_x = img_x.unsqueeze(dim=0)
    tab_x = tab_x.unsqueeze(dim=0)  # 确保tab_x为2维张量
    model.eval()  # 切换到评估模式
    with torch.no_grad():  # 测试阶段无需计算梯度,节省资源
        output = model(img_x, tab_x)
    print(output.shape)
    model.train()  # 后续训练切回训练模式
    
  • 方案2:修改BatchNorm参数适配单样本:
    初始化BatchNorm1d时设置track_running_stats=False和affine=False,此时BatchNorm会退化为对单个样本做归一化(类似LayerNorm),但这会改变BatchNorm的原有行为,仅适合特殊场景。
    修改模型中的BatchNorm定义:
    self.batchnorm = nn.BatchNorm1d(16, track_running_stats=False, affine=False)
    

彻底解决维度不匹配问题

确保tab_x传入模型前始终为2维张量(batch_size, feature_num),可在测试代码中固定执行unsqueeze(0),或在模型的forward方法开头统一处理:

def forward(self, img, tab):
    # 确保tab是2维张量
    if tab.dim() == 1:
        tab = tab.unsqueeze(0)
    # 后续原有代码...

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 19:41:57