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

MaxVit-t迁移学习报错RuntimeError:矩阵形状不匹配求解

问题排查与解决

错误原因分析

RuntimeError: mat1 and mat2 shapes cannot be multiplied (114688x7 and 512x8) 本质是矩阵乘法维度不匹配:前一层输出特征的列维度(7)与自定义分类器Linear层的输入行维度(512)不兼容,导致无法完成矩阵运算。

具体排查与解决步骤

1. 确认MaxVit-T的特征输出维度

先查看预训练MaxVit-T的分类器原始结构,明确特征提取部分的输出维度:

import torchvision.models as models
model = models.maxvit_t(pretrained=True)
print(model.classifier)

通常输出类似:

Sequential(
  (0): AdaptiveAvgPool2d(output_size=1)
  (1): Flatten(start_dim=1, end_dim=-1)
  (2): Linear(in_features=768, out_features=1000, bias=True)
)

这说明特征经过池化、展平后,输入到Linear层的维度是768,而非你设置的512。

2. 正确替换分类器

不要直接用错误维度的Linear层替换整个分类器,推荐两种方式:

  • 方式一:仅替换最后一层Linear(保留原池化、展平逻辑):
    from torch import nn
    # 自动匹配原输入维度,替换输出为8类
    model.classifier[-1] = nn.Linear(model.classifier[-1].in_features, 8)
    
  • 方式二:重新定义完整分类器(确保包含池化和展平):
    from torch import nn
    model.classifier = nn.Sequential(
        nn.AdaptiveAvgPool2d((1, 1)),
        nn.Flatten(),
        nn.Linear(768, 8)  # 768对应原分类器Linear的输入维度
    )
    

3. 验证输入图像尺寸

MaxVit-T默认要求输入图像尺寸为224×224,如果数据集图像尺寸不符,会导致特征图的高/宽维度异常,进而影响展平后的特征维度。确保数据预处理时将图像resize到指定尺寸:

from torchvision.transforms import Compose, Resize, ToTensor
transform = Compose([
    Resize((224, 224)),
    ToTensor()
])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 18:05:28