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

解决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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 19:40:37