提取CatAndDogConvNet首个Linear层特征时遇RuntimeError求助
问题原因与解决办法
错误根源
你遇到的形状不匹配错误,核心原因是截取的子模块缺少了将卷积输出的特征图展平为一维向量的关键步骤:
- 原模型的
fc1是全连接层,它的权重形状2304x500说明,它预期输入是长度为2304的一维向量(这个数值通常是卷积池化后特征图的通道数×高×宽的乘积)。 - 你用
list(model.children())[:7]截取的模块只包含了卷积和池化层,没有包含原模型中把多维特征图转成一维向量的操作(比如显式的Flatten层,或者forward函数里的x.view逻辑),导致fc1接收到的是多维张量(形状384x6),和它需要的输入形状不匹配,触发矩阵乘法错误。
解决步骤
先确认原模型结构
打印原模型的完整结构,找到负责展平特征图的模块:print(model)如果结构里有显式的
Flatten层,把它也包含到截取的子模块中。比如原模型的前8个children包含了卷积、池化和Flatten,就改成:model_new = torch.nn.Sequential(*list(model.children())[:8])如果原模型没有显式展平层
不少模型会在forward函数里用x = x.view(x.size(0), -1)来展平特征图,这种逻辑不会出现在model.children()里,你需要手动给model_new添加展平层:model_new = torch.nn.Sequential( *list(model.children())[:7], torch.nn.Flatten(start_dim=1) )更稳妥的自定义模型方式
直接创建自定义模型类,明确控制前向传播流程,避免截取子模块的不确定性:import torch.nn as nn class ExtractFc1Output(nn.Module): def __init__(self, original_model): super().__init__() # 取原模型的卷积特征提取部分 self.feature_extractor = nn.Sequential(*list(original_model.children())[:7]) # 复用原模型的fc1层 self.fc1 = original_model.fc1 # 添加展平层 self.flatten = nn.Flatten(start_dim=1) def forward(self, x): x = self.feature_extractor(x) x = self.flatten(x) x = self.fc1(x) return x model_new = ExtractFc1Output(model)
内容的提问来源于stack exchange,提问作者Sarvagya Gupta
相关产品推荐
相关产品推荐

