如何访问预训练TimeSformer模型Block内子模块的输出
问题
已知可以通过PyTorch Forward Hook获取预训练TimeSformer模型中指定Block的输出,示例代码如下:
import torch from timesformer.models.vit import TimeSformer model = TimeSformer(img_size=224, num_classes=400, num_frames=8, attention_type='divided_space_time', pretrained_model='/path/to/pretrained/model.pyth') activation = {} def get_activation(name): def hook(model, input, output): activation[name] = output.detach() return hook model.model.blocks[4].register_forward_hook(get_activation('block4')) x = torch.randn(3,3,224,224) output = model(x) block4_output = activation['block4']
现在需要进一步获取Block内子模块的输出,例如*标记的temporal_attn模块,以及@标记的temporal_attn内部的proj层的输出,该如何实现?
解决方案
核心逻辑和获取Block输出一致:精准定位到目标子模块,给它们注册Forward Hook即可,具体步骤如下:
获取
temporal_attn模块的输出
直接定位到Block下的temporal_attn子模块,路径为model.model.blocks[4].temporal_attn,注册Hook即可捕获其输出。获取
temporal_attn内部proj层的输出
继续深入到temporal_attn下的proj层,路径为model.model.blocks[4].temporal_attn.proj,同样注册Hook即可。
完整示例代码
import torch from timesformer.models.vit import TimeSformer # 初始化预训练模型 model = TimeSformer(img_size=224, num_classes=400, num_frames=8, attention_type='divided_space_time', pretrained_model='/path/to/pretrained/model.pyth') model.eval() # 推理模式下开启eval,避免BatchNorm/Dropout的随机影响 activation = {} def get_activation(name): def hook(model, input, output): activation[name] = output.detach() return hook # 为目标模块注册Hook model.model.blocks[4].register_forward_hook(get_activation('block4')) model.model.blocks[4].temporal_attn.register_forward_hook(get_activation('block4_temporal_attn')) model.model.blocks[4].temporal_attn.proj.register_forward_hook(get_activation('block4_temporal_attn_proj')) # 输入符合TimeSformer要求的张量(维度:[批量B, 帧数T, 通道C, 高H, 宽W]) x = torch.randn(3, 8, 3, 224, 224) output = model(x) # 提取各模块的输出 block4_output = activation['block4'] temporal_attn_output = activation['block4_temporal_attn'] temporal_attn_proj_output = activation['block4_temporal_attn_proj'] # 可选:打印输出维度验证 print(f"Block4输出维度: {block4_output.shape}") print(f"Temporal Attn输出维度: {temporal_attn_output.shape}") print(f"Temporal Attn Proj输出维度: {temporal_attn_proj_output.shape}")
注意事项
- 若不确定子模块路径,可通过
print(model.model.blocks[4].temporal_attn)打印模块结构,确认proj层的位置。 - 原示例的输入维度
torch.randn(3,3,224,224)不符合TimeSformer的输入规范,正确输入需包含帧数维度,示例中已修正。
内容的提问来源于stack exchange,提问作者dtr43
相关产品推荐
相关产品推荐

