使用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
相关产品推荐
相关产品推荐

