多模态 DiT 跑起来之后最先让人头疼的不是扩散步数也不是 VAE 解码而是 Transformer 主干里的 Attention 计算。尤其当序列里同时塞了图像 patch、文本 token甚至视频帧之后注意力的二次复杂度会直接变成训练和推理的瓶颈。我在这类模型上做了几个月的稀疏化改造最后落地的方案就是块稀疏 Attention不是去逐 token 裁剪而是按块决定哪些注意力区域可以不计算。下面把设计思路和工程细节完整讲一遍内容包括掩码设计、块划分、内核落地以及训练推理时容易踩的坑希望对正在做多模态 DiT 或类似长序列 Transformer 的朋友有参考价值。1. 多模态 DiT 的注意力瓶颈到底卡在哪里1.1 注意力二次复杂度如何被多模态长序列放大DiT 的核心是用 Transformer 替换 U-Net 这类经典扩散主干用 Attention 机制在 token 序列上做全局建模。单模态情况下比如纯文本或纯图像序列长度还能勉强控制。但多模态场景里要把不同类型的输入统一成 token 序列长度往往直接涨一个量级。拿文生图来说一张 1024x1024 的图切成 16x16 的 patch就是 4096 个图像 token再拼接上几十到几百个文本 token整体序列长度轻松超过 4200。这个长度带来的计算量是平方级别的。标准 Attention 对每个 head 要做 Q 乘 K 的矩阵乘法得到的是一个 N×N 的注意力分数矩阵。N 等于 4200 的时候单个 head 的分数矩阵就有 1764 万个元素。听起来还行但别忘了还有多头、有 batch、有扩散模型里的几十步去噪迭代以及多个不同模态特征之间的交叉注意力。序列一旦突破一万比如加入视频帧或高分辨率特征图单次前向的 Attention FLOPs 就会占据整个模型总计算量的六成以上。更麻烦的是显存。N×N 的分数矩阵即使不显式保存Flash Attention 这种分块算法也会在 SRAM 与 HBM 之间反复搬运 tile。密集 Attention 的 HBM 访问量虽然被 Flash Attention 压下来了但 FLOPs 依旧和 N² 成正比算力时间省不掉。去噪过程又要跑几十上百步每一步都在重复这些计算整体延迟就非常难看。1.2 Flash Attention 解决的是什么解决不了的是什么很多人一听到长序列 Attention 开销大第一反应就是上 Flash Attention。这个方向没有错但要分清楚它优化了什么。Flash Attention 做的事情是 IO 感知的 kernel 融合把 Q、K、V 分块后一次加载到 SRAM 里计算避免把 N×N 的分数矩阵完整写回显存并用在线 softmax 解决分块归一化的问题。它大幅降低了显存占用和访存延迟但没有改变 Attention 本身的 FLOPs 数量。算力层面的 O(N²) 复杂度依然存在。多模态 DiT 场景里这个 FLOPs 瓶颈尤其突出因为不同模态之间的信息冗余非常大。比如图像 patch 相邻位置的特征高度相关文本 token 之间也有大量局部依赖如果对所有 token 对都做完整 Attention很大一部分计算产出的权重接近于零本质上是在浪费算力。这也就是我转向块稀疏 Attention 的直接原因。目标很明确在保证生成质量不塌的前提下把那些低价值、甚至完全没有必要计算的注意力块跳过让 FLOPs 从 O(N²) 下降到 O(N² × 稀疏率)。这里稀疏率指实际参与计算的块占全部注意力块的比例比如只算 15% 的块理论上 FLOPs 就能接近原来的 15%效果非常直接。2. 块稀疏注意力从“精确计算一切”到“按块决定算不算”2.1 块稀疏的本质块掩码代替逐点掩码块稀疏 Attention 的核心思路不难理解。普通 Attention 有一个 N×N 的分数矩阵块稀疏做的就是把这个矩阵划分成若干个小方块比如每个块 64×64然后给每个块一个开关要算就整个块参与矩阵运算不算就整个跳过。为什么不直接在 token 级别做稀疏而是非要以块为单位原因在 GPU 的并行执行特性。你如果只是把少数几个无关的 token 去掉剩下的 token 分布不规则计算时仍然要把 Q、K、V 完整地加载进来通过 mask 把某些位置置零。这种逐点稀疏实际上没有减少矩阵乘法的规模CPU 或 GPU 依旧会遍历所有元素甚至因为 mask 操作带来额外开销。块稀疏就不同了它跳过的是一个连续的内存区域加载和计算只针对被选中的块内存带宽和算力都能真正省下来。以序列长度 4096、块大小 64 为例注意力矩阵会划分成 64×64 个块对。如果只保留其中 10% 的块对参与计算实际执行的矩阵乘规模就只是密集版本的十分之一。对依赖大规模并行计算的 GPU 来说这种按块跳过的方式既保留了规整的计算模式又真正做到了算力下降。2.2 块大小的选择计算效率与稀疏粒度的天平块大小是第一个要拍板的超参。常见的选项是 32、64、128。块越大GPU 上的矩阵乘越规整kernel 的执行效率越高但稀疏粒度会变粗可能把一些需要保留的注意力细节连同大块一起丢掉。块越小掩码越精细该省的地方省得更精确但计算模式开始变得零碎甚至出现部分 warp 空转的问题。我实测下来序列长度在 4096 到 8192 之间时64×64 是一个比较平衡的起点。128×128 在部分硬件上吞吐更高但遇到某些需要精细跨模态对齐的场景生成质量下降明显。32×32 的粒度最好但 kernel 开销明显变大实际收益反而不如 64。块大小的选择还和硬件平台相关A100/H100 的 SRAM 更大可以容纳更大的块消费级显卡上 64 往往更稳。建议初期直接用 64 跑通全流程后续把块大小作为可配置参数在目标硬件上做一组小扫描。2.3 静态掩码与动态路由两种主流做法块掩码怎么确定大致分两条路线。静态掩码在训练开始前就固定下来不依赖输入内容。它利用的是任务本身的先验结构比如图像 patch 只和邻近 patch 做注意力、文本 token 之间做全连接、跨模态只保留某些对齐区域。这种做法的优势是稳定可控训练和推理的稀疏模式完全一致缺点是遇到了分布外数据固定掩码可能丢掉关键信息。动态掩码则通过一个小型路由网络对每条样本预测应该激活哪些注意力块类似 MoE 的路由思想。动态掩码的表达能力更强理论上可以针对样本定制稀疏模式但训练难度也更高。路由网络需要额外的训练目标来保证稀疏率稳定还要防止训练和推理时的路由不一致导致质量崩坏。我在实际项目里的做法是两者结合模态间的粗粒度结构用静态掩码模态内部的一些长距离交互用动态路由去发现。这个组合在多模态 DiT 上表现比纯静态或纯动态都要好后面第 3 章会展开。3. 多模态下的块划分策略模态异质性就是天然优势3.1 模态间连接的先验设计多模态场景做块稀疏有一个很有利的条件不同的模态天然自带连接结构注意力矩阵根本不需要从头学起我们可以直接按模态边界把它划分为几个功能区。典型的多模态 DiT 输入包含文本 token 和图像 patch注意力矩阵可以分成四大块文本-文本、图像-图像、文本-图像作为 Query、图像-文本作为 Query。这四块的性质各不相同。文本-文本区域通常 token 数量少全局语义密度高应该保持密集或接近密集。图像-图像区域局部性强可以设计成局部窗口加少量全局锚点的稀疏模式。跨模态区域则要看任务文生图时文本作为全局条件应该允许所有图像 patch 都看到文本反过来图像 patch 之间如果要做区域对齐则可以根据空间位置裁剪。我建议先用一个简单的模态连接矩阵表达这些约束比如用V [V_tt, V_ii, V_ti, V_it]表示各个分区的激活密度。设计的时候不要一刀切文本对视觉的注意力可以给较高密度视觉内部反而可以压得更狠。这样既保证了生成质量又让稀疏率集中贡献在计算量最大的视觉区域上。3.2 模态内局部性的利用图像 patch 构成的 token 序列天然带有空间结构。按 patch 切分之后相邻 token 往往对应图片上的相邻区域。对这种序列做全局 Attention 浪费很大局部窗口 Attention 是更合理的选择。把图像区域的注意力矩阵设计成带状每个 patch 只和附近 k 个 patch 计算注意力再配合少量全局 token 做长距离信息交换就能把图像部分的计算量压到很低。视频场景还可以进一步扩展成三维块注意力很多工作里把这叫做 cuboid attention。视频 token 是时间、高度、宽度三个维度组成的张量块划分时把这三个维度一起切块每个 cuboid 内做密集 Attentioncuboid 之间跳过。这个思路和块的稀疏模式很契合掩码生成时直接按三维坐标判断两个 token 是否属于同一个 cuboid 即可。实际操作时要注意局部窗口会导致每个位置最终看到的 token 集合不一致attention mask 在实现时不能简单地用固定偏移量处理而是要生成一个上三角或下三角偏移的 band mask。好在块稀疏的内核天然支持任意块掩码只要把 band mask 落到块粒度上代码逻辑不用额外改动。3.3 模态不均衡怎么处理多模态输入最典型的毛病是模态不均衡文本 token 只有几十个图像 patch 却有几千个。如果直接按块稀疏把所有区域都压到同一个稀疏率文本区域的块数量本来就少再被跳掉一部分语义信息就会显著丢失。我的经验是按模态单独设置稀疏率而不是全局统一。具体做法是先把注意力矩阵按模态分组然后对每个组设定一个权重权重大的组保持较高的激活密度。这个权重可以和当前训练阶段挂钩比如扩散模型的前期去噪步侧重全局结构那时多给文本到图像的跨模态块一些密度后期精细纹理修复步则多保留图像局部块。利用这个阶段性的策略能在整体稀疏率不变的情况下把计算量用在更关键的去噪阶段。4. 块稀疏内核怎么落地Flash Attention 加一层稀疏索引4.1 把稀疏索引叠加到 Flash Attention 的 tile 循环上块稀疏 Attention 的落地方式不是把 Flash Attention 推翻重来而是在它原有的 tile 循环外面套一层稀疏索引。Flash Attention 原本的逻辑是把 Q、K、V 都切块在一个双重循环里遍历所有的 Q-block 和 KV-block 组合。我们只需要把内层遍历的 KV-block 列表替换成稀疏掩码指定的活跃块集合就能实现 block-sparse Flash Attention。参考实现逻辑大致是for i in range(num_q_blocks): q_block load(Q, i) # 加载当前 Q 块 acc zeros_like(q_block) m_i -inf for j in active_kv_blocks[i]: # 只遍历稀疏掩码选中的 KV 块 k_block load(K, j) v_block load(V, j) s_block q_block k_block.T # 得到注意力分数块 s_block apply_mask(s_block) # 处理填充 padding 或因果掩码 m_new, acc online_softmax_update(acc, s_block, v_block, m_i) m_i m_new out[i] normalize(acc, m_i)这里active_kv_blocks[i]是关键的稀疏索引。它可以从预先计算好的稀疏掩码构建对每个 Q 块维护一个活跃 KV 块的列表推荐用 CSR 或位图这样的紧凑格式存到显存里kernel 运行的时候直接从全局内存读取。选择位图格式的话每个块对应一个 bit读取后做一次 popcount 判断就可以跳过逻辑简单且对 GPU 比较友好。4.2 需要重写内核吗成熟框架和自写的取舍如果从零写 CUDA kernel工作量不小因为要同时处理分块装载、在线 softmax、块稀疏跳转这三件事。更务实的选择是看成熟的稀疏注意力框架。xFormers 里的 memory-efficient attention 支持 block-sparse 模式给定一个 block mask 就能用上 fused kernel。FlashAttention 官方后续版本也加入了对 block-sparse mask 的实验支持。在 PyTorch 生态下这两个库足够覆盖大多数 DiT 训练场景。Triton 也是一个好方向它比起 CUDA 门槛低很多而且对不规则循环的支持比手写 CUDA 容易控制写出来的 block-sparse Flash Attention kernel 性能通常能达到官方 FlashAttention 的八成以上。我建议的使用方式先在 xFormers 的 block-sparse 模式上把整个算法流程跑通用它的 kernel 验证正确性和收益如果后续对性能有更高要求再针对自己的掩码模式在 Triton 里实现定制化的 kernel。没有必要一上来就手撕 CUDA除非你是想发布一个通用的稀疏算子否则工程迭代效率会很低。4.3 在线 softmax 与稀疏掩码的配合块稀疏状态下的在线 softmax 需要额外小心。Flash Attention 的分块 softmax 之所以正确是因为最终会把所有块的分值统一归一化。但是在块稀疏模式下被跳过的块相当于注意力分数为负无穷。在线 softmax 的实现必须意识到这一点否则会把未覆盖块的分数误当成零处理导致最终的 softmax 分布错误。实现上有两个做法。一是把被跳过的块实际填充为很大的负数比如-inf在 softmax 前参与 max 和 exp 的计算这样能保持数值语义正确但会带来额外开销。二是把每个 Q 块需要覆盖的 KV 块数量记录下来在线 softmax 更新时只对活跃块做局部归并最后再整体除以真实的总和项。后者更高效但要求内核内部能感知活跃块的数量和范围属于定制化实现。我踩过的坑是某些框架的 block mask 只负责跳过计算不会自动处理 softmax 的归一化分母导致最后注意力分布整体偏小。结果模型输出能量偏低生成图像偏暗偏灰。排查时发现是分母项漏了被跳过的块补上之后一切恢复正常。所以验证 block-sparse 内核时第一件事就是用一个小模型对比密集 Attention 的输出确认分块 softmax 的数值一致性。5. 稀疏率、掩码设计与训练推理中的实战坑5.1 稀疏率应该设在什么范围一组实测参考稀疏率不是越低越好它和任务、序列长度、模型容量都有关。我在文生图类 DiT 上的实测结果大致如下可作为初始参考配置平均激活块占比推理加速比生成质量主观评价密集 Attention100%1.0x基准局部窗口 全局锚点25%-30%约 2.8x接近基准细节略有损失激进稀疏10%-12%约 6x明显失真局部结构模糊动态路由辅助15%-18%约 4.5x质量接近 25% 静态方案可以看出静态局部窗口方案能稳定压在 25% 到 30% 的激活比例主观质量损失很小。继续往下压到 10% 附近图像细节会明显崩坏尤其是高频纹理和跨模态边缘区域。因此我的建议是第一版别追求极致稀疏率先做到 20% 到 30%把质量保住后面再逐步加入动态路由或课程式稀疏率。这些数字当然和 token 长度强相关。序列越长、冗余越大能压的稀疏率就越狠。如果你的任务主要是 512 分辨率以下的图像序列长度短块稀疏的绝对收益有限甚至不如直接优化 kernel 省事。做之前先用 profiler 看一眼 Attention 占总时长的比例如果不到 30% 就先别折腾稀疏。5.2 掩码在训练和推理之间的对齐问题块稀疏最让人头疼的问题之一是动态掩码在训练和推理之间的行为不一致。训练时路由网络学会了针对某个 batch 的稀疏模式但如果推理时换了输入分布或改了采样策略路由给出的掩码会出现偏移导致某些本该激活的块被漏掉生成质量突然下降。解决这个问题的思路主要有三个。第一在训练过程中加入 mask 平滑让路由网络的输出带有一定噪声或 dropout避免它对特定样本过拟合。第二推理阶段固定使用训练后期阶段的平均 mask 统计不再依赖路由实时预测这样虽然牺牲了一点动态性但稳定性大幅提升。第三周期性对整个验证集统计掩码激活分布如果发现某些块从未被激活说明路由网络对其没有建模能力需要在训练数据中增强对应模态的样本。5.3 几个容易踩的数值和工程坑数值方面除了上一节提到的 softmax 分母问题还有一个常见坑是-inf掩码的传播。用-inf做 mask 后注意力分数矩阵里可能出现 NaN原因是在某些 GPU 上-inf乘以 0 或者-inf减去-inf会产生异常。规避方案是用一个很大的负数代替-inf比如-10000同时保证经过 softmax 后概率足够小。这个技巧虽然朴素但能省掉很多调试时间。工程方面GPU 负载不均衡要重点关注。块稀疏后每个 Q 块的活跃 KV 块数量可能差异很大如果简单地按 Q 块做并行分派某些线程块忙死、某些线程块空转。解决手段是在 kernel 里做一次排序或合并尽量让同一个 tile 内的活跃块数量接近或者把稀疏率定义成每个 Q 块独立的目标激活数而不是全局平均值。这样能明显提高实际吞吐而不是只停留在理论 FLOPs 的下降。还有一个小坑是掩码生成本身也可能成为瓶颈。如果每步去噪都要在 CPU 上重新生成掩码再拷到 GPU那么掩码生成的开销可能直接吃掉节省下来的计算时间。建议预生成静态掩码并在训练循环里把动态掩码的路由网络放到 GPU 上实时推理避免 CPU-GPU 之间反复搬运索引数据。6. 我在实际使用中的几点体会6.1 什么时候块稀疏收益最大并不是所有多模态模型都适合做块稀疏。我个人的经验是序列长度超过 4096 的 DiT 系列收益最明显尤其是需要处理视频或多帧图像的任务。短片视频里 Attention 区域的冗余度极高块稀疏的加速比可以达到十倍量级。反过来如果 input 长度只有几百密集 Attention 的实现已经被优化得很好了强行稀疏反而因为索引和跳块的额外开销出现负优化。另外一个体会是块稀疏要想真正好用必须跟推理框架的图优化结合起来。比如在 TensorRT 或自研推理引擎里把块稀疏的索引和 DiT 的采样循环融合避免每一步都重新构建稀疏块索引。训练时可以依赖 PyTorch 的动态图但部署时就需要把掩码静态化提前编译成固定推理图这才是生产环境能稳定加速的关键。6.2 后续可以沿着什么方向扩展我在现有项目上已经验证了局部窗口和交通路动态块的组合方案。下一步想试的是把稀疏掩码的设计交给一个小型可微搜索模块直接用验证集上的生成质量作为损失自动搜索每个注意力层的块稀疏结构。这样就不用手工设计掩码还能在不同数据集上自适应调整。另一个方向是在多模态对齐上用块稀疏做文章。热词里提到的多模态情感分析、多模态时序对齐这类任务本身就需要不同模态特征之间做细粒度的交互建模。块稀疏 Attention 完全可以作为一种结构先验让模型只对语义上对齐的模态块做注意力从架构层面减少模态间的噪声干扰。这个思路不止适用于 DiT也可以迁移到其他多模态 Transformer 模型上。