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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 11:45:29