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)
解决方案
问题根源
- 数据集加载器缺失Resize操作:
build_your_dataset的transform中注释掉了transforms.Resize((224, 224)),导致部分图片实际尺寸不达标,输入网络后特征图尺寸不符合预期;同时存在类名笔误(YourCustomDataset应为YourcnnDataset)。 - 模型forward函数返回错误:
hierarchical=True时硬编码返回固定形状的特征图,既不匹配get_feature_map_channels的长度要求,也和实际卷积输出尺寸不符,导致SparK的mask生成函数输出的张量维度与卷积结果不匹配。 - 卷积层无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
相关产品推荐
相关产品推荐

