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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 07:08:16