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

使用Hook函数提取预训练R2Plus1D模型特征时遇Net模块错误

问题场景

我需要从PyTorchVideo仓库的R2Plus1D预训练模型中提取部分层级特征,已移除模型的最终层并通过strict=False加载预训练权重,模型架构如下:

Net(
  (blocks): ModuleList(
    (0): ResNetBasicStem(...)
    (1): ResStage(...)  # 目标提取层1
    (2): ResStage(...)  # 目标提取层2
  )
)

尝试用Hook函数提取ResStage(1)和ResStage(2)的特征,代码如下:

class mymodel(nn.Module):
def __init__(self, pretrained=False):
    super(mymodel, self).__init__()
    
    self.activation = {}
    def get_activation(name):
        def hook(model, input, output):
            self.activation[name] = output.detach()
        return hook

    self.r2plus1d = create_r2plus1d()

    self.r2plus1d.Net.blocks[1].register_forward_hook(get_activation('ResBlock1'))
    self.r2plus1d.Net.blocks[2].register_forward_hook(get_activation('ResBlock2'))

def forward(self, x, out_consp = False):
    x = self.r2plus1d(x)
    
    block1_output = self.activation['ResBlock1']   # channel_num:256
    block2_output = self.activation['ResBlock2']   # channel_num:512
    
    return block1_output, block2_output

运行时报错,提示模型的state_dict中不存在Net模块。


解决方案

问题根源

create_r2plus1d()函数返回的对象本身就是Net类的实例,不需要额外通过.Net属性访问,直接调用.blocks即可获取模块列表。

修改后的代码

class mymodel(nn.Module):
def __init__(self, pretrained=False):
    super(mymodel, self).__init__()
    
    self.activation = {}
    def get_activation(name):
        def hook(model, input, output):
            self.activation[name] = output.detach()
        return hook

    self.r2plus1d = create_r2plus1d()

    # 移除多余的.Net,直接访问blocks
    self.r2plus1d.blocks[1].register_forward_hook(get_activation('ResBlock1'))
    self.r2plus1d.blocks[2].register_forward_hook(get_activation('ResBlock2'))

def forward(self, x, out_consp = False):
    # 前向传播触发hook
    _ = self.r2plus1d(x)
    
    block1_output = self.activation['ResBlock1']
    block2_output = self.activation['ResBlock2']
    
    return block1_output, block2_output

额外提示

  • 确保create_r2plus1d()返回的模型确实包含blocks属性(从你给出的模型架构来看是符合的)。
  • 如果后续遇到类似模块路径问题,可以通过print(self.r2plus1d)打印模型结构,确认层级关系后再编写访问路径。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 11:05:28