PyTorch MultiheadAttention与Torchvision MViT_v2注意力模块运行速度差异问询
问题:MViT_v2_s骨干+自定义跨注意力模块的性能差异分析
背景与实现
基于Torchvision的MViT_v2_s作为骨干网络,添加了自定义跨注意力模块FusionModule,模块代码如下:
class FusionModule(nn.Module): def __init__(self, embed_dim: int, num_heads: int, source_a_input_channels: int, source_b_input_channels: int): super().__init__() # embed_dim = source_a_input_channels = source_b_input_channels self.attn = nn.MultiheadAttention(embed_dim=embed_dim, num_heads=num_heads, batch_first=True) self.source_a_pool = nn.LayerNorm() self.source_b_pool = nn.LayerNorm() self.proj_norm = nn.LayerNorm() self.mlp = MLP() # a two layer mlp def forward(self, source_a: torch.Tensor, source_b: torch.Tensor): # source_a takes the output of a MViT multiscale block source_a = self.source_a_pool(source_a) # reshape source_b input to (b, thw, c) source_b = self.source_b_pool(source_b.flatten(2).transpose(1, 2)) # after reshape, source_b has almost the same shape as source_a # except source_b has one less token fused = self.attn(source_a, source_b, source_b)[0] mid_prod = source_a + fused mid_prod = self.proj_norm(mid_prod) out = self.mlp(mid_prod) out = mid_prod + out return out
模块被添加在MViT骨干网络的每个stage之后(即第1、3、14、16个多尺度模块之后)。
性能异常现象
输入形状为[1, 3, 16, 224, 224]的张量时,第一个FusionModule的跨注意力计算耗时约2.7秒,而MViT骨干的第一阶段(包含自注意力模块及其他组件)仅耗时0.2秒。
从理论FLOPs来看,两者应相近:
- MViT骨干模块对
[1, 25089, 96]张量计算自注意力 - 自定义
FusionModule对[1, 25089, 96]查询张量与[1, 25088, 96]键/值张量计算跨注意力
已确认:
- MLP计算耗时占比极小
- 所有计算在CPU上完成
- 纯MViT_v2_s模型处理输入耗时约1.2秒,与带自定义模块的模型中骨干网络耗时一致,排除骨干本身问题
核心疑问
是MViT_v2的自注意力实现确实比PyTorch原生MultiheadAttention高效得多,还是自定义模块存在性能瓶颈?
原因分析与优化建议
1. MViT_v2自注意力的针对性优化
Torchvision的MViT_v2并非直接使用PyTorch原生MultiheadAttention,而是针对视觉Transformer的多尺度场景做了深度优化:
- 针对多尺度token的稀疏性设计了稀疏注意力逻辑,减少无效计算
- CPU端做了内存布局优化(如保证张量连续、对齐),并利用了更高效的向量化矩阵乘法实现
- 原生
MultiheadAttention是通用实现,未针对视觉Transformer的张量形状、计算场景做定制化优化,CPU上的并行性利用效率较低
2. 自定义模块的潜在性能瓶颈
- 张量内存不连续:
source_b.flatten(2).transpose(1,2)会导致张量内存碎片化,CPU上对非连续张量的矩阵乘法效率会大幅下降。建议修改为:source_b = source_b.flatten(2).transpose(1, 2).contiguous() source_b = self.source_b_pool(source_b) - LayerNorm参数缺失:代码中
nn.LayerNorm()未传入归一化维度,默认会做隐式维度推断,可能带来额外计算开销。需明确传入维度,比如nn.LayerNorm(embed_dim) - 非对称Q/KV形状:source_b比source_a少一个token,原生
MultiheadAttention处理非对称形状时无针对性优化,而MViT自注意力是对称形状,计算更高效
3. 验证方法
用相同形状的张量分别测试torch.nn.MultiheadAttention和MViT内部自注意力模块的耗时,直接对比两者的效率差异,即可验证核心原因。
内容的提问来源于stack exchange,提问作者whz
相关产品推荐
相关产品推荐

