解决PyTorch模型前向传播中的“Dimension out of range”错误
PyTorch融合ResNet与颜色信息的维度错误解决
错误根源
触发IndexError: Dimension out of range的直接原因是:你试图访问ResNet输出特征的第三个维度(features1.size(2)),但ResNet去掉全连接层后的输出是全局平均池化后的二维张量(形状[batch_size, num_features]),根本没有空间维度(H/W)。同时get_color_info函数返回的张量多了一个冗余的batch维度,进一步导致维度不匹配。
代码问题拆解
1. ResNet特征维度误解
ResNet默认结构里,fc层前是avgpool全局平均池化,所以把fc替换为Identity()后,model1(x)输出的是[batch_size, 512](ResNet18)的二维特征,不存在size(2)和size(3)这两个维度。
2. color_info维度冗余
get_color_info最后多了一个.unsqueeze(0),导致返回张量形状是[1, batch_size, 3, H, W],而模型需要的是[batch_size, 3, H, W]的四维输入。
修正方案
推荐将颜色信息转换为全局特征,和ResNet的全局特征拼接,符合分类任务的特征融合逻辑,实现起来也更简单。
修正后的完整代码
1. 调整ClothesRecognizer模型
import torch import torch.nn as nn from torchvision import models, transforms from PIL import Image class ClothesRecognizer(nn.Module): def __init__(self, num_classes): super(ClothesRecognizer, self).__init__() # 加载ResNet-18,替换fc为Identity保留全局特征 self.model1 = models.resnet18(pretrained=True) self.model1.fc = nn.Identity() self.in_features1 = self.model1.fc.in_features # 512 # 加载ResNet-50,同理替换fc self.model2 = models.resnet50(pretrained=True) self.model2.fc = nn.Identity() self.in_features2 = self.model2.fc.in_features # 2048 # 颜色信息全局池化层,把[B,3,H,W]转为[B,3] self.color_global_pool = nn.AdaptiveAvgPool2d(1) # 融合三类特征的分类器 self.classifier = nn.Sequential( nn.Linear(self.in_features1 + self.in_features2 + 3, 512), nn.ReLU(), nn.Dropout(0.5), nn.Linear(512, num_classes) ) def forward(self, x, color_info): # 获取两个ResNet的全局特征(二维) features1 = self.model1(x) # shape: [B, 512] features2 = self.model2(x) # shape: [B, 2048] # 处理颜色信息:全局池化后展平为一维特征 color_global = self.color_global_pool(color_info).flatten(1) # shape: [B, 3] # 按特征维度拼接所有特征 combined_features = torch.cat([features1, features2, color_global], dim=1) return self.classifier(combined_features)
2. 修正get_color_info函数
去掉最后的.unsqueeze(0),确保返回正确的四维张量:
def get_color_info(batch): images = batch[0] color_info_list = [] for image in images: image_pil = transforms.ToPILImage()(image) if image_pil.mode != 'RGB': image_pil = image_pil.convert('RGB') r, g, b = image_pil.split() r = transforms.ToTensor()(r) g = transforms.ToTensor()(g) b = transforms.ToTensor()(b) color_info = torch.stack([r, g, b]).unsqueeze(0) color_info_list.append(color_info) # 拼接后得到[B,3,H,W],无需额外增加维度 return torch.cat(color_info_list, dim=0)
维度匹配验证
- 输入图像x:
[batch_size, 3, H, W] - ResNet特征:
[batch_size, 512]+[batch_size, 2048] - 颜色全局特征:
[batch_size, 3] - 拼接后特征:
[batch_size, 512+2048+3=2563],完全匹配分类器输入维度
内容的提问来源于stack exchange,提问作者Okuhle Ntloyiya
相关产品推荐
相关产品推荐

