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

PyTorch中提取TimeSformer指定模块特征并移除末层的方法

自定义TimeSformer特征提取实现方案

要实现从第4、8、11个Block提取特征并移除最后的norm和head层,以下是完善后的自定义模型代码,解决直接访问blocks的问题:

完整代码

import torch
import torch.nn as nn
from pytorchvideo.models.video_transformers import TimeSformer

class TimeSformerFeatureExtractor(nn.Module):
    def __init__(self, pretrained_model, target_blocks=[4, 8, 11]):
        super().__init__()
        # 映射1-based的Block编号到0-based索引
        self.target_indices = [idx - 1 for idx in target_blocks]
        
        # 复制原模型的前端组件
        self.patch_embed = pretrained_model.patch_embed
        self.pos_embed = pretrained_model.pos_embed
        self.cls_token = pretrained_model.cls_token
        
        # 兼容不同版本的Block访问路径
        self.all_blocks = pretrained_model.blocks if hasattr(pretrained_model, "blocks") else pretrained_model.backbone.blocks
        # 保留所有Block以维持完整前向链路
        self.blocks = nn.ModuleList([block for block in self.all_blocks])

    def forward(self, x):
        batch_size = x.shape[0]
        # 前端嵌入处理,原生支持5D输入(B, C, T, H, W)
        x = self.patch_embed(x)
        cls_tokens = self.cls_token.expand(batch_size, -1, -1)
        x = torch.cat((cls_tokens, x), dim=1)
        x = x + self.pos_embed
        
        # 遍历所有Block,收集目标位置的特征
        extracted_features = []
        for idx, block in enumerate(self.blocks):
            x = block(x)
            if idx in self.target_indices:
                extracted_features.append(x.clone())
        
        # 跳过原模型的norm和head层,直接返回特征
        return extracted_features

核心要点

  • 索引转换:TimeSformer的blocks采用0索引,代码自动将你指定的第4/8/11个Block转换为对应索引3/7/10。
  • Block访问兼容:部分pytorchvideo版本中blocks被封装在backbone子模块下,代码通过属性判断自动适配。
  • 完整前向链路:必须遍历所有Block并执行前向计算,不能跳过中间层,否则后续Block的输入状态会出错,导致特征无效。
  • 移除norm和head:完全不调用原模型的norm和head模块,自然规避这两部分的计算逻辑。

使用示例

# 加载预训练模型
pretrained_model = TimeSformer.from_pretrained("timesformer_base_16x16_160k_kinetics")

# 初始化特征提取器,指定要提取的Block编号(1-based)
extractor = TimeSformerFeatureExtractor(pretrained_model, target_blocks=[4,8,11])

# 构造5D输入:(batch_size=2, channels=3, frames=16, height=224, width=224)
test_input = torch.randn(2, 3, 16, 224, 224)

# 获取特征
features = extractor(test_input)

# 输出每个特征的形状
for i, feat in enumerate(features):
    print(f"第{extractor.target_indices[i]+1}个Block特征形状: {feat.shape}")

排障提示

如果仍无法访问blocks,先打印模型结构确认层级:

print(pretrained_model)

找到blocks所在的路径(比如pretrained_model.backbone.blocks),替换代码中self.all_blocks的赋值语句即可。

内容的提问来源于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 20:39:14