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

如何访问预训练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即可,具体步骤如下:

  1. 获取temporal_attn模块的输出
    直接定位到Block下的temporal_attn子模块,路径为model.model.blocks[4].temporal_attn,注册Hook即可捕获其输出。

  2. 获取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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 22:06:23