PyTorch数据并行下Hook提取TimeSformer特征报错解决
解决DataParallel包装TimeSformer后访问模型层报错的问题
问题根源
用nn.DataParallel包装模型后,原TimeSformer实例会被封装到module属性下,而非你代码中写的model属性,所以直接访问self.featureExtractor.model.blocks会触发AttributeError。
具体解决办法
方法1:修正属性访问路径
直接将访问路径中的model替换为module,比如:
# 原错误代码 # target_blocks = self.featureExtractor.model.blocks # 修正后 target_blocks = self.featureExtractor.module.blocks
后续注册Hook时,直接针对module下的对应层操作即可,示例:
def feature_hook(module, input, output): self.cached_features = output.detach() # 比如要提取第3个block的特征(索引从0开始) self.featureExtractor.module.blocks[2].register_forward_hook(feature_hook)
方法2:保留原模型引用(更便捷)
在包装DataParallel时,单独保留原TimeSformer模型的引用,后续直接通过原引用访问层和注册Hook,无需关心DataParallel的包装:
# 初始化预训练TimeSformer original_model = TimeSformer.from_pretrained("your-pretrained-path") # 用DataParallel包装 self.featureExtractor = nn.DataParallel(original_model) # 保留原模型引用 self.original_tsformer = original_model # 后续访问blocks或注册Hook直接用原模型 self.original_tsformer.blocks[2].register_forward_hook(feature_hook)
这种方式更清晰,而且Hook注册在原模型上后,DataParallel的每个GPU副本都会继承该Hook,不影响多GPU下的特征提取。
注意事项
- 多GPU环境下,Hook捕获的特征会来自每个GPU的前向传播,若需要汇总特征,可在Hook函数中对
output做torch.cat或其他聚合操作(根据你的需求)。 - 若使用
torch.nn.parallel.DistributedDataParallel,同样遵循module属性访问原模型的规则,解决逻辑一致。
内容的提问来源于stack exchange,提问作者dtr43
相关产品推荐
相关产品推荐

