PyTorch R2Plus1D-18模型输出形状异常问题及优化方案咨询
R2Plus1D-18模型输出形状疑问及解决方法
问题描述
我使用PyTorch官方的R2Plus1D-18 3D ResNet模型,已注释resnet.py源码中的flatten行,理论上输出不应为1D形状。
我的模型代码:
class VideoModel(nn.Module): def __init__(self,num_channels=3): super(VideoModel, self).__init__() self.r2plus1d = models.video.r2plus1d_18(pretrained=True) self.r2plus1d.fc = Identity() for layer in self.r2plus1d.children(): layer.requires_grad_ = False def forward(self, x): print(x.shape) x = self.r2plus1d(x) print(x.shape) return x
用于替换fc层的Identity类:
class Identity(nn.Module): def __init__(self): super().__init__() def forward(self, x): return x
当输入torch.randn(1, 3, 8, 112, 112)时,输出为:
torch.Size([1, 3, 8, 112, 112]) torch.Size([1, 512, 1, 1, 1])
疑问:为何移除fc层和flatten操作后仍得到这种形状?是否有更优的移除flatten操作的方法?
解答
输出形状的原因
R2Plus1D-18的主干网络最后一层是全局自适应3D平均池化层(AdaptiveAvgPool3d),它会自动将特征图的时间、高度、宽度维度压缩为1x1x1。你看到的[1,512,1,1,1]是池化层的直接输出,和flatten无关——原模型中flatten只是把池化后的形状转成1D向量[B,512],再传入fc层。你注释了flatten但保留了池化层,自然会得到池化后的形状。
更优的特征获取方法
不需要修改源码,直接通过截取模型子模块或重写forward的方式,就能灵活获取不同阶段的特征:
方法1:重写forward跳过池化、flatten和fc
class VideoModel(nn.Module): def __init__(self, num_channels=3): super().__init__() self.r2plus1d = models.video.r2plus1d_18(pretrained=True) # 冻结所有参数 for param in self.r2plus1d.parameters(): param.requires_grad = False def forward(self, x): # 只执行到最后一个残差层,跳过池化、flatten和fc x = self.r2plus1d.stem(x) x = self.r2plus1d.layer1(x) x = self.r2plus1d.layer2(x) x = self.r2plus1d.layer3(x) x = self.r2plus1d.layer4(x) print(x.shape) # 输出: torch.Size([1, 512, 1, 7, 7]) return x
方法2:直接构建特征提取器
通过nn.Sequential截取模型的前半部分(去掉池化和fc层):
# 构建只包含特征提取主干的模型 r2plus1d_backbone = nn.Sequential(*list(models.video.r2plus1d_18(pretrained=True).children())[:-2]) # 冻结参数 for param in r2plus1d_backbone.parameters(): param.requires_grad = False # 测试输入 x = torch.randn(1, 3, 8, 112, 112) out = r2plus1d_backbone(x) print(out.shape) # 输出: torch.Size([1, 512, 1, 7, 7])
这两种方法都不需要修改源码,既保留了预训练模型的权重,又能获取到未经过池化压缩的完整特征图,比修改源码的方式更灵活、更安全。
内容的提问来源于stack exchange,提问作者Alex Mehta
相关产品推荐
相关产品推荐

