中文

面向剪枝视觉 Transformer 的调度感知参差注意力

机器学习 2026-05-13 v2 人工智能

摘要

针对视觉 Transformer (ViTs) 的 Token 剪枝方法通过丢弃信息量少的图像块,有望实现注意力 FLOPs 的平方级减少。然而,标准的变长注意力 API——包括 FlashAttention-2 的 varlen 和 PyTorch 的 NestedTensor SDPA——未能将这些节省转化为 ViT 典型的剪枝后短序列长度(\leq197 个令牌)下成比例的挂钟时间增益。我们识别出一个调度开销瓶颈:在这些长度下,无论工作负载如何,主机端内核调度消耗约 {\sim}50\,μ\mus,在中到高剪枝率下超过了实际的 GPU 计算时间。我们提出了一个轻量级的双向 Triton 注意力内核,其调度下限为 {\sim}24\,μ\mus——大约比 FlashAttention-2 varlen 低 2.17×\times——使得剪枝带来的节省在墙钟时间上变得可见。集成到一个完整的打包-关注-解包流水线中,并在 NVIDIA RTX 4000 Ada Generation GPU 上进行评估,我们的系统在标准 224×\times224 输入上实现了比填充的 PyTorch SDPA 高 1.88×\times 的端到端吞吐量,在 384×\times384 输入上扩展到 2.51×\times。与最强的基线 FlashAttention-2 varlen 相比,我们的内核在服务批次大小 (BS=1-4) 下提供了 9-12\% 更高的吞吐量,并在 80\% 令牌剪枝率下实现了 2.17×\times 更低的内核延迟。数值正确性通过最大绝对 logit 差异 <0.004 和比特精确的 top-1 预测得到验证。

关键词

引用

@article{arxiv.2604.15408,
  title  = {Dispatch-Aware Ragged Attention for Pruned Vision Transformers},
  author = {Seifeldin Abdellatif and Ahmad Almasri},
  journal= {arXiv preprint arXiv:2604.15408},
  year   = {2026}
}