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

CNN-ViT混合模型训练时张量维度不匹配RuntimeError求助

问题描述

基于SparK项目预训练文档,在自定义数据集上训练CNN-ViT混合模型时,反复遇到RuntimeError:The size of tensor a (222) must match the size of tensor b (168) at non-singleton dimension 3。已确认输入图片尺寸均为224x224,重写数据集加载器后问题仍存在,查阅类似错误方案均不匹配当前场景。


报错堆栈信息

Traceback (most recent call last):
  File "main.py", line 193, in <module>
    main_pt()
  File "main.py", line 103, in main_pt
    stats = pre_train_one_ep(ep, args, tb_lg, itrt_train, iters_train, model, optimizer)
  File "main.py", line 159, in pre_train_one_ep
    loss = model(inp, active_b1ff=None, vis=False)
  File "/global/common/software/nersc/shasta2105/pytorch/1.9.0/lib/python3.8/site-packages/torch/nn/modules/module.py", line 1051, in _call_impl
    return forward_call(*input, **kwargs)
  File "main.py", line 36, in forward
    return self.module(*args, **kwargs)
  File "/global/common/software/nersc/shasta2105/pytorch/1.9.0/lib/python3.8/site-packages/torch/nn/modules/module.py", line 1051, in _call_impl
    return forward_call(*input, **kwargs)
  File "/global/cfs/cdirs/dune/www/data/2x2/simulation/silentc_work/SparK/SparK/SparK/pretrain/spark.py", line 96, in forward
    fea_bcffs: List[torch.Tensor] = self.sparse_encoder(masked_bchw)
  File "/global/common/software/nersc/shasta2105/pytorch/1.9.0/lib/python3.8/site-packages/torch/nn/modules/module.py", line 1051, in _call_impl
    return forward_call(*input, **kwargs)
  File "/global/cfs/cdirs/dune/www/data/2x2/simulation/silentc_work/SparK/SparK/SparK/pretrain/encoder.py", line 209, in forward
    return self.sp_cnn(x, hierarchical=True)
  File "/global/common/software/nersc/shasta2105/pytorch/1.9.0/lib/python3.8/site-packages/torch/nn/modules/module.py", line 1051, in _call_impl
    return forward_call(*input, **kwargs)
  File "/global/cfs/cdirs/dune/www/data/2x2/simulation/silentc_work/SparK/SparK/SparK/pretrain/models/custom.py", line 42, in forward
    x = self.conv1(inp_bchw)
  File "/global/common/software/nersc/shasta2105/pytorch/1.9.0/lib/python3.8/site-packages/torch/nn/modules/module.py", line 1051, in _call_impl
    return forward_call(*input, **kwargs)
  File "/global/cfs/cdirs/dune/www/data/2x2/simulation/silentc_work/SparK/SparK/SparK/pretrain/encoder.py", line 23, in sp_conv_forward
    x *= _get_active_ex_or_ii(H=x.shape[2], W=x.shape[3], returning_active_ex=True)    # (BCHW) *= (B1HW), mask the output of conv
RuntimeError: The size of tensor a (222) must match the size of tensor b (168) at non-singleton dimension 3

数据集加载器代码

import os
from typing import Any, Callable, Optional, Tuple

import PIL.Image as PImage
from timm.data import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD
from torchvision.datasets.folder import DatasetFolder, IMG_EXTENSIONS
from torchvision.transforms import transforms
from torch.utils.data import Dataset

try:
    from torchvision.transforms import InterpolationMode
    interpolation = InterpolationMode.BICUBIC
except:
    import PIL
    interpolation = PIL.Image.BICUBIC


def pil_loader(path):
    # open path as file to avoid ResourceWarning (https://github.com/python-pillow/Pillow/issues/835)
    with open(path, 'rb') as f: img: PImage.Image = PImage.open(f).convert('L')
    return img


import os
from PIL import Image
import torch
from torch.utils.data import Dataset
from torchvision import transforms

class YourcnnDataset(Dataset):
    def __init__(self, data_path, input_size, transform=None):
        self.data_path = data_path
        self.input_size = input_size
        self.transform = transform
        self.images = self._load_images()

    def _load_images(self):
        # Load images from the data_path directory
        image_list = []
        for root, dirs, files in os.walk(self.data_path):
            for file in files:
                if file.endswith(".jpg") or file.endswith(".png"):
                    image_list.append(os.path.join(root, file))
        return image_list

    def __len__(self):
        return len(self.images)

    def __getitem__(self, idx):
        img_path = self.images[idx]
        image = Image.open(img_path).convert("L")

        if self.transform:
            image = self.transform(image)
        image = transforms.ToTensor()(image)

        return image

def build_your_dataset(data_path, input_size, batch_size):
    # Define your transformation for the dataset
    your_transform = transforms.Compose([transforms.RandomHorizontalFlip(),
        transforms.ToTensor(), transforms.Normalize(mean=(0.5), std=(0.5)),]) #transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(), transforms.Normalize(mean=(0.5), std=(0.5)),   https://github.com/pytorch/pytorch/issues/9446

    # Create an instance of your custom dataset
    your_dataset = YourCustomDataset(data_path, input_size, transform=your_transform)

    # Data loader to handle batching
    data_loader = torch.utils.data.DataLoader(your_dataset, batch_size=batch_size, shuffle=True)

    return data_loader

模型架构代码

import torch
import torch.nn as nn
from typing import List
from timm.models.registry import register_model


class YourConvNet(nn.Module):
    def __init__(self, num_classes=0, global_pool=''):
        super(YourConvNet, self).__init__()
        self.conv1 = nn.Conv2d(in_channels=1, out_channels=32, kernel_size=3, stride=1)
        self.relu = nn.ReLU()
        self.maxpool1 = nn.MaxPool2d(kernel_size=2, padding=0)
        self.conv2 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3, stride=1)
        self.maxpool2 = nn.MaxPool2d(kernel_size=2, padding=0)

    def get_downsample_ratio(self) -> int:
        # Define the downsampling ratio
        return 4  # Update this based on your actual architecture

    def get_feature_map_channels(self) -> List[int]:
        # Define the number of channels of each feature map
        return [32, 64]  # for conv1 and conv2

    def forward(self, inp_bchw: torch.Tensor, hierarchical=False):
        """
        The forward with `hierarchical=True` would ONLY be used in `SparseEncoder.forward` (see `pretrain/encoder.py`).

        :param inp_bchw: input image tensor, shape: (batch_size, channels, height, width).
        :param hierarchical: return the logits (not hierarchical), or the feature maps (hierarchical).
        :return:
            - hierarchical == False: return the logits of the classification task, shape: (batch_size, num_classes).
            - hierarchical == True: return a list of all feature maps, which should have the same length as the return value of `get_feature_map_channels`.
              E.g., for the provided ConvNet, it should return a list [1st_feat_map, 2nd_feat_map].
                    for an input size, the shapes could be [(B, 32, 56, 56), (B, 64, 28, 28)] based on provided architecture.
        """
        x = self.conv1(inp_bchw)
        x = self.relu(x)
        x = self.maxpool1(x)
        x = self.conv2(x)
        x = self.relu(x)
        x = self.maxpool2(x)

        # Depending on the hierarchical flag, return feature maps or logits
        if hierarchical:
            # Return feature maps
            #return [(1, 64, 12, 12)]  # Update this based on your actual architecture
            return [(1, 64, 54, 54)]
        else:
            # Perform further operations for classification logits, if needed
            # Example: Flatten x and add fully connected layers
            return x


@register_model
def your_convnet_small(pretrained=False, **kwargs):
    return YourConvNet(**kwargs)

@register_model
def your_cnn(pretrained=False, **kwargs):
    return YourConvNet(**kwargs)


@torch.no_grad()
def convnet_test():
    from timm.models import create_model
    cnn = create_model('your_convnet_small')
    print('get_downsample_ratio:', cnn.get_downsample_ratio())
    print('get_feature_map_channels:', cnn.get_feature_map_channels())


    downsample_ratio = cnn.get_downsample_ratio()
    feature_map_channels = cnn.get_feature_map_channels()
    
    # check the forward function
    B, C, H, W = 100, 1, 224, 224
    inp = torch.rand(B, C, H, W)
    feats = cnn(inp, hierarchical=True)
    assert isinstance(feats, list)
    assert len(feats) == len(feature_map_channels)
    print([tuple(t.shape) for t in feats])
    
    # check the downsample ratio
    feats = cnn(inp, hierarchical=True)
    assert feats[-1].shape[-2] == H // downsample_ratio
    assert feats[-1].shape[-1] == W // downsample_ratio
    
    # check the channel number
    for feat, ch in zip(feats, feature_map_channels):
        assert feat.ndim == 4
        assert feat.shape[1] == ch


if __name__ == '__main__':
    convnet_test()

模型测试代码

import torch
from custom import YourConvNet  # Replace 'YourConvNet_module_file' with your actual module file name

# Instantiate your YourConvNet model
model = YourConvNet()

# Define sample input data
sample_input = torch.randn(1, 1, 224, 224)  # Assuming input shape (batch_size, channels, height, width)

# Perform a forward pass to get the output feature maps
with torch.no_grad():
    feature_maps = model(sample_input, hierarchical=True)

# Print the shapes of the output feature maps
for i, fmap in enumerate(feature_maps):
    print(i, fmap)

解决方案

问题根源

  1. 数据集加载器缺失Resize操作:build_your_dataset的transform中注释掉了transforms.Resize((224, 224)),导致部分图片实际尺寸不达标,输入网络后特征图尺寸不符合预期;同时存在类名笔误(YourCustomDataset应为YourcnnDataset)。
  2. 模型forward函数返回错误:hierarchical=True时硬编码返回固定形状的特征图,既不匹配get_feature_map_channels的长度要求,也和实际卷积输出尺寸不符,导致SparK的mask生成函数输出的张量维度与卷积结果不匹配。
  3. 卷积层无padding导致尺寸不整除:两次3x1无padding卷积后,特征图尺寸无法被下采样率4整除,触发断言失效和维度不匹配问题。

修复步骤

1. 修复数据集加载器

恢复Resize操作并修正类名笔误:

def build_your_dataset(data_path, input_size, batch_size):
    your_transform = transforms.Compose([
        transforms.Resize((224, 224)),
        transforms.RandomHorizontalFlip(),
        transforms.ToTensor(), 
        transforms.Normalize(mean=(0.5), std=(0.5))
    ])
    your_dataset = YourcnnDataset(data_path, input_size, transform=your_transform)
    data_loader = torch.utils.data.DataLoader(your_dataset, batch_size=batch_size, shuffle=True)
    return data_loader

2. 修正模型forward函数

返回与get_feature_map_channels匹配的分层特征图:

def forward(self, inp_bchw: torch.Tensor, hierarchical=False):
    # 保存conv1池化后的特征图
    x1 = self.conv1(inp_bchw)
    x1 = self.relu(x1)
    x1_pooled = self.maxpool1(x1)
    
    # 保存conv2池化后的特征图
    x2 = self.conv2(x1_pooled)
    x2 = self.relu(x2)
    x2_pooled = self.maxpool2(x2)

    if hierarchical:
        # 返回对应通道数的两个特征图
        return [x1_pooled, x2_pooled]
    else:
        return x2_pooled

3. 调整卷积层padding保证尺寸整除

修改卷积层添加padding,让特征图尺寸严格匹配下采样率:

# 初始化时修改conv1和conv2
self.conv1 = nn.Conv2d(in_channels=1, out_channels=32, kernel_size=3, stride=1, padding=1)
self.conv2 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3, stride=1, padding=1)

此时输入224x224图片,经过两次卷积+池化后,最终特征图尺寸为56x56,正好是224/4的结果,符合get_downsample_ratio=4的定义。

4. 运行模型测试验证

执行convnet_test,确保所有断言通过,验证特征图尺寸和通道数完全符合要求。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 12:37:01