特色项目

要点速览

  • 传统的推荐系统(RecSys)推理需要为每个候选对象显式复制共享的用户嵌入(user embeddings)或序列。内核内广播优化(IKBO)通过内核-模型-系统协同设计消除了这一开销,将广播逻辑直接融合到用户-候选对象交互内核中。通过减少内存占用和 IO 利用率,IKBO 实现了更高的吞吐量。
  • IKBO 可将计算密集型网络延迟降低高达 2/3,成为支持 Meta 自适应排序模型 的以请求为中心、推理高效框架的可扩展性基石
  • IKBO 已在 GPU 和 MTIA(Meta 训练与推理加速器)上端到端部署于 Meta 的多阶段推荐流程中。
  • 经过四个阶段的渐进式协同设计,IKBO 线性压缩内核在 H100 SXM5 上实现了累计 ~4 倍 的加速,最终通过 TLX 实现了线程束专用(warp-specialized)融合。
  • IKBO 协同设计将 Flash Attention 内核从 IO 密集型转变为计算密集型(在 H100 SXM5 上达到 621 BF16 TFLOPs)。配合 TLX 的线程束专用优化,相比非协同设计的 CuTeDSL FA4 Hopper 基准(仅内核/内核+广播),吞吐量提升了 2.4 倍/6.4 倍

在这篇文章中,我们介绍了内核内广播优化(IKBO),这是一种内核-模型-系统协同设计方法,消除了推荐模型推理中冗余的用户嵌入广播。在生产级推荐系统中,给定请求的用户嵌入在所有候选对象中都是相同的,但标准方法需要显式复制,这浪费了随候选对象数量线性增长的内存带宽和计算资源。IKBO 包含一个简单的见解:广播是数据布局问题,而非计算必需。每个 IKBO 内核在接收自然、批次大小不匹配的用户和候选对象输入时,会在内部处理广播,因此不会产生任何重复的张量。我们通过两个内核深入解析:线性压缩和 Flash Attention 来展示该方法。

IKBO 已部署在 Meta 的整个推荐系统推理堆栈中——从早期到晚期的排序模型,涵盖 GPU 和 MTIA(Meta 训练与推理加速器)——在协同设计的模型上实现了高达 2/3 的计算密集型网络延迟降低。它是支撑 Meta 自适应排序模型(生产环境中服务于 LLM 规模模型)的以请求为中心、推理高效框架的可扩展性基石。在 H100 SXM5 上,我们的 IKBO 线性压缩内核通过四个渐进式协同设计阶段实现了 ~4 倍加速:矩阵乘法分解、内存对齐、广播融合,以及通过 TLX(Triton 底层扩展) 实现的线程束专用多阶段融合。对于 Flash Attention,IKBO 相比非协同设计的 CuTeDSL FA4-Hopper 提供了 2.4 倍/6.4 倍的吞吐量(仅内核/内核+广播),达到 621 BF16 TFLOPs。与绕过复制的系统级广播或网络拆分不同,IKBO 在计算原语层消除了复制,以近乎独立的成本实现了稠密的交互质量。

代码仓库:https://github.com/pytorch/FBGEMM/tree/main/fbgemm_gpu/experimental/ikbo

工作完成时作者在 Meta 任职

1. 内核内广播优化:消除内存和计算冗余

当用户打开信息流时,推荐系统必须对数百到数千个候选项目进行评分以决定显示内容。模型的输入分为两类:用户特征(例如:浏览历史、个人资料、上下文),对于请求中的每个候选对象都是相同的;以及候选特征(例如:项目 ID、类别、互动统计),对于每个项目都是唯一的。两者都通过嵌入查找和后续处理生成嵌入表示。在模型的多个点,交互层(例如:线性投影、特征交叉、目标注意力)结合用户和候选嵌入。我们将请求中所有候选对象共享的嵌入称为仅请求(RO),将每个候选对象的嵌入称为非仅请求(NRO)

图 1. 一个高度简化的推荐系统推理数据流。在交互层之前,仅请求(RO)用户嵌入必须被广播(复制)以匹配非仅请求(NRO)候选批次维度。IKBO 通过在每个内核内部处理广播消除了这种物化。

交互层要求张量具有匹配的批次维度。在由约 15 个用户服务的 1,024 个候选对象的批次中,RO 嵌入必须被广播并复制约 70 次,才能在任何交互之前匹配 NRO 批次大小(图 1)。随着架构从 DLRM [1] 和 DCN [2] 演进到像 HSTU [3] 和 X 的 Phoenix [4] 这样的序列模型,它们不断丰富了用户-候选对象的交互。但更丰富的交互是有代价的:用户特征必须在所有候选对象中广播。对于推理中 10 到 10,000+ 的批次大小,这种复制开销导致了随候选对象数量线性增长的巨大计算和内存成本。

广播是数据布局问题,而非计算必需。从这个角度审视模型和推理系统,可以在每一层进行优化:推理运行时消除了系统级广播,仅用户模型层以较小的用户批次大小运行,而混合两者的内核经过重新设计以在内部处理广播——从不产生冗余的复制张量。IKBO 已部署在 Meta 的推荐系统推理堆栈中,涵盖从早期到晚期的排序模型,跨越 GPU 和 MTIA,在协同设计的模型上实现了高达 2/3 的计算密集型网络延迟降低。

本文通过两个深入解析:线性压缩和 Flash Attention,重点关注内核层面的优化。

1.1. 内核优化类型

类型 I — 可分解操作。通过数学重构,仅请求(RO)部分可以在较小的批次大小下独立计算,仅在最后与非仅请求(NRO)部分合并。这既节省了内存带宽,又节省了计算资源。

类型 II — 仅内存优化。 在内核内部处理 RO-NRO 广播避免了冗余的数据移动,使内核不再受限于 IO 瓶颈。

1.2. E2E 系统设计

部署 IKBO 涉及基础设施堆栈的三层

  1. 内核:自定义 GPU 内核,接收不匹配的 RO/NRO 批次大小并在内部处理广播(第 2 和第 3 节)。
  2. 编译规范:ML 编译器需要每个算子的动态形状范围来选择适当形状的内核。单批次大小时这很简单;若有两个(用户和候选对象)甚至更多,要系统性地自动化解析每个算子所使用的对象——尤其是在交互掩盖批次血缘的生产模型中——则极具挑战。
  3. 推理:运行时将候选对象到用户的映射传递到模型中,而不是物化广播。

这些内核通过以下两条路径之一进入模型

  1. 直接采用:模型作者将 IKBO 内核直接集成到模型定义中。当训练时候选对象与用户比率 > 1 时,相同的内核也降低了训练成本。
  2. 推理时转换:一种转换过程在推理时自动将标准算子替换为 IKBO 等价物 — 无需更改模型代码。

最终效果:广播从推理的每个阶段消失,模型没有架构限制,除了推理运行时的映射接口外,无需任何基础设施更改。

1.3. 与其他方法的比较

现有方法在绕过广播而不是消除它。 

  1. 系统级广播在 GPU 分发前物化复制张量——简单但浪费,成本随候选对象数量线性增长。 
  2. 网络拆分 (ROO) [5] 将模型划分为 RO 和 NRO 子网络,减少了冗余工作,但限制了用户-候选对象交互的发生位置,且在较小的 RO 批次大小时仍引入额外成本。 

两者都保留了作为物化张量的广播。IKBO 在计算原语层消除了它:节省随候选对象/用户比率增长,任何交互模式都无需广播成本,且完整的 NRO 批次维度在融合内核内提供了 GPU 占用率。

IKBO 已部署在 GPU 和 MTIA 加速器上。在本博客中,我们重点讨论 H100 GPU 内核设计以说明核心优化原则。

2. 内核深入解析 I:IKBO 线性压缩

线性压缩嵌入(LCE)通过学习到的投影 (M, K) @ (B, K, N) → (B, M, N) 压缩输入嵌入 (B, K, N),并被 Meta 推荐系统模型广泛采用,如 Wukong [6]。我们进行了四个渐进式优化阶段。

2.1 矩阵乘法分解

图 2. LCE 分解:基准批次矩阵乘法(左上),沿 K 的嵌入分离和用户去重(右上),带有压缩输出广播加法的两个独立 GEMM(左下)。

基准 LCE 在所有 B 个候选对象上计算单个批次矩阵乘法。输入嵌入沿 K 连接用户和候选部分 — 但同一用户的所有候选对象的用户嵌入是相同的。

将广播推迟到矩阵乘法之后。 由于 W 与批次无关,我们通过线性分解:沿 K 分离用户和候选嵌入块,对重复的用户嵌入去重,并在其自然的批次大小下计算两个独立的 GEMM。我们不再在矩阵乘法前复制用户嵌入,而是仅广播小的压缩结果。见图 2。候选对象与用户比率约为 70(一个典型设置),用户批次从 B=1024 缩小到 B_user ≈ 15 — 用户侧计算量减少了 70 倍。该分解在标准 PyTorch 中实现。

结果。 1.944 毫秒 → 1.389 毫秒(减少 28.5%;基准测试设置见附录 1)。原始批次 GEMM(算术强度 ~ 356 FLOPs/Byte,低于 H100 的 ~495 FLOPs/Byte 机器平衡点;推导见附录 2)和两个分解后的 GEMM 都是内存受限的,因此加速是由内存成本降低驱动的。去重使内存成本降低了一半以上 — 因为用户侧 GEMM(B_user ≈ 15 对比 B = 1024)的成本变得微不足道。

注意,该分解将广播推到了矩阵乘法之后:我们不再在 GEMM 前复制完整的 K 维输入嵌入,而是仅广播小的压缩结果,这要便宜得多。在第 2.3 节中,我们将通过内核内广播融合彻底消除剩余的这一步广播。

目前的瓶颈是 L1/TEX 流水线利用率(84%)而不是 DRAM 利用率 — 我们将在下一节深入分析这种不平衡。详细配置分解见附录 3。

2.2 内存布局优化

分解 GEMM 的详细结果分析揭示了不平衡:L1/TEX 处于峰值的 84%,而 DRAM 仅达到 19%,表明内存加载过于狭窄。SASS 确认:每个 cp.async 仅复制 4 字节,而不是 128 位加载。

LDGSTS.E.LTC128B P0, [R203],      [R38.64]       // 4 bytes
LDGSTS.E.LTC128B P1, [R203+0x4],  [R38.64+0x4]   // 4 bytes  (×4 total, only 16B load in total)

cp.async 宽度受限于源指针的自然对齐。矩阵 A 为 (M, K) 行主序,步长为 K × 2 字节,因此当 K 不是 8 的倍数时,步长会破坏 128 位对齐。

模型-内核协同设计见解。 内存对齐是一种被充分理解的 GPU 优化 — 但分解使其成为模型-内核协同设计的挑战。K 由嵌入张量的 torch.cat 形成,其大小取决于许多模型配置因素。分解使得手动调整这些因素以使分解后的嵌入保持完美倍数非常困难。需要一个系统性的解决方案。

解决方案。 通过在 concat 列表中追加零,将每个分解后的 K 填充到下一个 8 的倍数。我们证明这在数学上在向前和向后传播中是等价的(见下面的证明 1),并且配合 ML 编译器的内存规划器,这可简化为廉价的常量复制。

证明 1。 零填充 K 在向前和向后传播中均保持精确的数值等价。

结果。 1.389 毫秒 → 0.798 毫秒(减少 42.5%)。填充使得 CUTLASS 能够选择基于 TMA 的内核,从而完全绕过 L1/TEX(扇区 351M → 0),并将 GEMM 延迟从 0.984 毫秒削减至 0.400 毫秒。随着 GEMM 的解决,未融合的广播和加法(0.398 毫秒)现在占总延迟的一半 — 将在下一节解决。详细结果分析见附录 5。

2.3 候选 GEMM 内核内广播融合

未融合的广播和加法受内存瓶颈限制:将候选 GEMM 结果写入 HBM,连同用户结果一起读回,相加,再次写入。我们通过将广播融合到候选 GEMM 尾声(epilogue)中消除了这一点(图 3)。在每个平铺(tile)积累后,尾声查找用户索引,加载预计算的用户结果,在寄存器中相加,并写入最终总和 — 中间张量从未被物化。我们将其实现为 Triton 内核:具有自定义积累后尾声块的标准批次 GEMM。

图 3. 内核内广播融合:GEMM 尾声通过索引查找加载预计算的用户结果并在寄存器内相加。

结果。 0.798 毫秒 → 0.580 毫秒(减少 27.4%)。融合消除了 0.87 GB 的中间 DRAM 流量,贡献了延迟提升。然而,占用率仅为 6.25%(每个调度程序 1 个线程束),使每个停顿完全暴露。除了 42% 的周期在等待全局加载外,20% 的周期花费在等待 WGMMA — 这些停顿无法被尾声隐藏,且如果没有持久性,则没有下一个平铺加载来重叠。这是一个具有挑战性的权衡:需要大的平铺和深流水线来保持张量核心的供应,但它们消耗了大部分共享内存预算,几乎没有空间通过占用率来隐藏延迟。详细结果分析见附录 6。

2.4 通过 TLX 进行线程束专用的多阶段融合

TLX(Triton 底层语言扩展) 暴露了 Hopper 的线程束专用化、TMA、mbarrier 和命名屏障,同时保留了 Triton 的 Python DSL 和自动调优基础设施。

利用 TLX,我们通过线程束专用化解决了第 2.3 节的占用率限制 — 通过功能分区而非额外的线程束来隐藏延迟。

第 2.1 – 2.3 节将原始 LCE 分解为两个独立的计算:用户 GEMM(阶段 1)和带有融合广播加法尾声的候选 GEMM(阶段 2)。我们首先优化阶段 2(主要瓶颈)内的延迟隐藏,然后将两个阶段融合为单个持久内核。

阶段内延迟重叠

候选 IKBO 内核受内存限制 — 设计目标是保持内存流水线持续供应。Triton 的软件流水线(第 2.3 节)已经将加载与 WGMMA 重叠,但尾声仍是序列化的 — 它阻塞了未来的加载并暴露了 WGMMA 等待停顿。我们通过将每个 CTA 划分为专用的线程束组来解决这两个问题:一个专用的生产者持续发出 TMA 加载(重叠 #1,类似于 Triton 的软件流水线),而两个消费者平铺乒乓操作,使一个的尾声与另一个的 WGMMA 重叠(重叠 #2)。通过持久性,平铺连续流动,没有跨平铺间隙。见图 4。

图 4. 具有两个阶段内延迟重叠和线程束组角色分配的候选 IKBO 内核结构。

多阶段融合

我们将用户 IKBO(阶段 1)和候选 IKBO(阶段 2)融合为单个超大内核,以减少波量化(wave quantization),消除内核启动开销,并提高 L2 缓存利用率。高候选对象/用户比率放大阶段 1 中的波量化。由于候选 GEMM 在其尾声前独立于用户结果,我们并发调度这两个阶段。

这种并发调度解锁了两个额外的跨阶段重叠,使总重叠达到四个。见图 5。

图 5. 并发阶段调度:没有用户平铺的 SM 立即进入阶段 2,与阶段 1 的部分波重叠。多阶段融合后的所有四个延迟重叠,显示了阶段内(#1, #2)和跨阶段(#3, #4)重叠机会。SM 0-49, 50-131 为示例数字。

线程束组专用化与同步设置

为了实现所有四个重叠,每个 CTA 被划分为一个生产者和两个消费者线程束组。关键是,两个阶段共享相同的循环缓冲区和 mbarrier 基础设施 — 在阶段边界处不会发生流水线排空或屏障重新初始化。最后一个用户 K 块和第一个候选 K 块同时存在于不同的缓冲区槽中。见图 6。

图 6. 每个 CTA 的线程束组设置和三个同步机制。

双向阶段交替平铺调度

当两个阶段的平铺计数不能被 SM 计数整除时,天真的单向分发会导致工作负载不平衡。我们反转了阶段之间的平铺分配方向:阶段 1 从 pid 开始,阶段 2 从 NUM_SM - 1 - pid 开始。见图 7。

图 7. 单向(左)对比双向阶段交替分发(右),平衡部分波上的每 SM 工作负载。

平铺粒度跨 CTA 同步

用户和平铺候选平铺可能在不同的 CTA 上执行,需要跨 CTA 同步 — 但设备级屏障会序列化所有工作并破坏重叠。我们使用三步发布-获取协议在平铺粒度上进行同步:

  1. 每个线程束组的单个线程以 ld.relaxed 在平铺标志上轮询,最大限度地减少内存流量
  2. 一旦设置,单个 ld.acquire 建立发生前边(happens-before edge)
  3. 命名屏障向线程束组中的所有 128 个线程广播就绪状态

这避免了轮询期间昂贵的围栏(fence),并允许不同用户平铺上的候选 CTA 完全独立地进行。详细信息见附录 7。

结果

结合所有优化,延迟从 0.580 毫秒改善至 0.482 毫秒(减少 16.9%)。清晰的线程束内 Proton 追踪器 时间线确认所有四个重叠在实践中均已实现。

图 8。 针对两个 CTA 的 Proton 分析器时间线,所有四个重叠均有颜色标记。内存流水线保持持续供应。

主要收益来自重叠 #2:乒乓消费者隐藏了每个平铺上的 WGMMA 和尾声停顿 — 直接解决了第 2.3 节中主要的浪费周期。重叠 #1(加载↔WGMMA)继承自 Triton 现有的软件流水线。重叠 #3 和 #4 隐藏了用户到候选阶段过渡时的空闲时间。见图 8。

NCU 确认:占用率从 6.25% 上升至 18.75%(3 个线程束组对比 1 个),DRAM 吞吐量从 39% 上升至 52%,L2(瓶颈)从 74% 上升至峰值的 84%。这不仅仅是占用率的问题:所有四个重叠中激进的延迟隐藏保持了内存流水线的饱和,这正是将 L2 推过 80% 的原因。详细 NCU 指标见附录 8。

我们在不同批次大小和候选对象/用户比率下进行基准测试,采用默认(批次=1024,比率=70)设置。见图 9。

图 9. 跨批次大小(左,比率=70)和候选对象/用户比率(右,批次=1024)的累计 IKBO 加速。

IKBO 融合在不同场景下均提供了稳健的收益:跨批次大小(左)和候选对象/用户比率(右)均有 ~4 倍加速。即使在较低的候选对象/用户比率下,内核仍实现了显著加速。

3. 内核深入解析 II:IKBO Flash Attention

随着推荐模型扩展以捕获更丰富的用户顺序行为,包括注意力机制在内的顺序架构已成为关键的计算瓶颈,在 1K 序列长度时约占推理延迟的 40%。这促使我们关注与推荐系统独特批次语义协同设计的 IKBO 感知 Flash Attention。

受 Transformers 和 Set Transformers [7, 8] 的启发,两种基础的用户历史交互模块已被推荐系统广泛采用: 

  • 目标注意力(类似于交叉注意力)捕获预测候选对象与用户历史交互之间的关系。
  • 自注意力 对用户历史内部的顺序依赖关系进行建模

由于用户历史是 RO 特征,而目标作用于不同的候选对象(非 RO)批次维度,这种架构不对称性为 IKBO 提供了提高模型可扩展性和计算效率的机会。目标注意力将是我们优化的主要焦点,通过细微的协同设计,自注意力也可以在第 3.3 节中融合到 IKBO 目标注意力中。由于我们的模型是编码器驱动的,因此应用的是没有因果掩码的完整注意力。

利用 E2E 协同设计的最终优化目标注意力版本实现了非协同设计 CuTeDSL FA4-Hopper(仅注意力内核 / 注意力内核 + 广播成本)的 2.4 倍/6.4 倍 吞吐量,分别降低了 0.320ms / 1.232ms 的延迟(表 2)。

3.1 IKBO Flash Attention 解决了推荐系统边界条件下的 IO 瓶颈问题

图 10:传统 SDPA 带候选对象-用户广播(左)对比融合的 IKBO 目标注意力(右)。 

IKBO 将 K/V 广播融合到注意力内核中,通过推理运行时提供的候选对象-用户映射张量保持数学等价性,该张量处理非均匀的候选对象/用户比率。图 10 对比了两种方法:传统 SDPA 路径在注意力之前将 K 和 V 广播到完整候选批次大小,而 IKBO 路径完全消除了这种物化 — 每个候选对象即时索引到其用户的 K/V。

通过 IKBO 协同设计将 IO 受限转变为计算受限

在推荐系统边界条件下,与用户的浏览历史相比,目标注意力使用相对较少数量的候选嵌入来表示候选属性。标准注意力的 Roofline 分析显示算术强度约为 60 FLOPs/Byte – 远低于 H100(SXM5 HBM2e 版本)的峰值约 495 FLOPs/Byte(附录 2)— 使得即使是标准 Flash Attention 也高度受限于 IO。IKBO 通过跨共享同一用户上下文的多个候选对象摊销 K/V 内存访问来解决此问题,将算术强度从约 60 FLOPs/Byte 提高到约 833 FLOPs/Byte(在 B_candidate : B_user = 70:1 时),并将内核牢牢推入计算受限区域。

为了最大化此收益,我们的实现重新排序了线程块启动网格,使 batch_size_candidate 排在 num_heads 之前。这确保了处理不同候选对象 — 但共享相同用户 K/V — 的线程块能够并发调度,提高了 L2 缓存复用率。

网格维度 Flash Attention (SDPA) IKBO 目标注意力
x num_q_seq_block num_q_seq_block
y num_heads batch_size_candidate
z batch_size_candidate num_heads

表 1:启动网格配置比较。SDPA 通过将 num_heads 置于 grid.y 来优先考虑 GQA 优化。IKBO 交换了头和候选维度,将 batch_size_candidate 置于 grid.y 以实现跨候选对象的高效 K/V 共享。

表 2 对比了我们的 IKBO Triton 实现(FA2 逻辑 + IKBO)与 Hopper 上最先进的 Flash Attention 实现(无 IKBO 协同设计)。吞吐量和 IO 仅在注意力上测量;Key 和 Value 的广播延迟甚至比注意力成本本身还要大。

吞吐量 (TFLOPs/s) IO (GB/s) 延迟 (ms)
Triton IKBO FA2 425 487 0.321(广播已融合)
TLX FA3 245 2152 0.561 + 0.912(K&V 广播)
CuTeDSL FA4 Hopper 250 2193 0.550 + 0.912(K&V 广播)
TLX IKBO FA3 持久通用化 594 681 0.230(广播已融合)

表 2:推荐系统边界条件下的注意力内核比较(B_candidate = 2048, B_u = 32, 均匀候选对象/用户比率)。没有协同设计,即使是最先进的 Hopper 实现也保持受限于 IO。

3.2 在 TLX 上采用现代内核技术(FA3, FA4)与 IKBO

随着 IKBO 将内核从 IO 受限转变为计算受限,下一步自然是采用 Flash Attention 3 (FA3 [10]) 和 Flash Attention 4 (FA4 [11]) 在 Hopper 上的最先进计算优化 — 特别是线程束专用化和流水线。然而,我们在查询嵌入数量(q_seq = 32 或 64)上的边界条件使得难以直接采用 FA3 的乒乓或协作线程束专用化。

Hopper 上的线程束专用化需要异步 WGMMA 指令,这强加了最小 BLOCK_M ≥ 64。还需要两个消费者线程束组来最大限度地减少它们之间的气泡。为了满足这些约束,我们定制了内核以在单个线程块内启动 B_candidate = i 和 B_candidate = i + 1,共享相同的 B_user。在下面的讨论中,我们假设所有用户都以 q_seq = 64 对偶数个候选对象进行排序;奇数候选对象处理随后进行。

IKBO FA3 内核的性能提升

从 FA3 的方案开始 — 线程束内流水线、线程束组专用化和乒乓调度 — 最初的 TLX IKBO FA3 内核表现与 FA2 基准类似(图 12,蓝色对比红色,附录 11),吞吐量相当。

为了诊断瓶颈,我们使用以 GPU 周期为延迟单位的 Proton 追踪器 可视化了线程束内流水线(图 10)。表 3 总结了持久化前后关键的瓶颈,通过 Proton 追踪器以 GPU 周期测量。

图 11:TLX IKBO FA3 内核的基于 Proton 的线程束内分析。显示了来自每个线程束组的代表性线程束:线程束 0(生产者)、线程束 4(消费者 1)和线程束 8(消费者 2)。softmax_PV_overlap 和纯 softmax 区域被单独标记以识别张量核心气泡。(A)持久化前(B)持久化前(2 个波)的放大视图(C)持久化后(2 个波)

瓶颈 之前 之后 关键变化
张量核心气泡(每个波的第 1 个 QKT,蓝色) ~1,300 周期(线程束调度程序切换导致 400 周期 ~1,300 周期 未改变
张量核心气泡(每个波的最后一个 PV,蓝色) ~2,000 周期 ~300 周期 异步 TMA 存储 + 与最后一个 PV 的倒数重叠
跨 CTA 停顿(橙色) ~14,000 周期 已消除 持久化彻底移除了 CTA 重新启动
初始化缓冲区和屏障(绿色) ~1,600 周期/波 ~1,600 周期(仅第 1 个波) 持久化共享缓冲区和屏障在波之间摊销 
等待第 1 个 Q/K 加载深紫色 2,100~4,000 周期/波(长度取决于 HBM 带宽争用 ~2,000 周期(仅第 1 个波) 跨波流水线;生产者预取 ~3K 周期

表 3:持久化 + 优化前后的关键瓶颈。

关键结论:在这些小的查询序列长度下,跨 CTA 停顿是主要的瓶颈 — 而不是张量核心利用率。对于此改进,持久化是必须的。持久化后,分析结果及其延迟变化显示在图 11C 和表 3 中。

HBM2e 特定优化

我们进一步针对 H100 SXM5 的 HBM2e 带宽约束对持久内核进行了调优,用共享内存容量换取了减少的加载/存储阻塞。(表 4)。

定制优化/修复 收益
将 O 的 SMEM 缓冲区与 Q/V 解耦,使用流水线 TMA 异步存储 启用 TMA 异步存储可以将 O 与 Q/V SMEM 共享解耦,从而与下一波计算重叠,将存储阻塞时间从 1,300 缩短至 400 周期/波
单独的 Q₀ 和 Q₁ 缓冲区  减少了每 Q 的加载时间,允许一个消费者组更早开始 — 当波计数大大超过 K/V 序列迭代时(推荐系统中常见)非常有利
指令缓存缺失修复 将剥离出的最后一次迭代代码路径合并回主循环,消除了由过多的线程束专用指令引起的指令缓存抖动(附录 12)

表 4: 针对 HBM2e H100 SXM5 的定制优化。在推荐系统边界条件下,这些仍符合可用的 SMEM 预算(附录 10)。

我们还实现了持久化 V2,它从 K 序列末尾向前迭代(匹配 FA3/FA4-Hopper 的方法)以简化掩码逻辑。两个持久化变体都应用了表 4 的优化。如图 12 所示,在低序列长度(512–4,096)下,TLX FA3 持久内核优于所有其他候选对象;超过 8K 时,两个持久化变体趋于一致。

图 12:IKBO 实现吞吐量对比序列长度(B_candidate = 2,048; B_candidate : B_user = 64; num_head = 2; d_head = 128)。实际推荐系统序列长度在 4K 以下 [3];更长的长度包含在内是为了与 LLM 用例进行比较。通用版本处理非偶数个每用户候选对象,且每用户有 50% 的奇数候选对象概率

为任意候选对象批次大小通用化 IKBO FA3

我们的 IKBO FA3 内核每个 CTA 处理两个候选批次,以满足 WGMMA 的 BLOCK_M ≥ 64 要求。当用户有奇数个候选对象时,一个消费者线程束组没有配对伙伴。我们通过空转逻辑处理这一点(图 13,左;算法 1)

  • 空转线程束组通过 mbarrier 信号排空 K/V 缓冲区以防止生产者死锁。
  • 活动线程束组禁用乒乓同步(其伙伴不再到达命名屏障)。

在约 70 : 1 的候选对象/用户比率下,空转路径触发时间不到 0.7%,开销可忽略不计(图 12,IKBO TLX FA3 通用化)。此方法通用化到 q_seq_len = 32,其中使用类似的空转和掩码逻辑在每个 CTA 中捆绑四个候选批次。

图 13:通用化目标注意力的 CTA 分配(左)和自 + 目标注意力融合(右)。每个 CTA 分配两个共享相同用户 K/V 的消费者线程束组。当候选对象计数为奇数时,第 2 个消费者空转并排空屏障。

算法 1:具有奇数候选对象处理的 IKBO 注意力向前传播

3.3 通过模型协同设计进行自 + 目标注意力融合

前面的章节专注于优化目标(交叉)注意力。一个自然的问题出现了:我们能将自注意力折叠到同一个内核中吗?

关键见解是这两种注意力类型共享相同的键-值源 — 用户序列。唯一的区别是查询:自注意力查询来自用户侧,而目标注意力查询来自候选对象侧。通过在两者之间共享 K/V 投影,我们实现了单个启动内的直接水平内核融合。图 13(右)说明了融合的 CTA 布局:第一个 CTA 处理自注意力查询块,而剩余的 CTA 处理目标注意力候选对象对 — 全部读取相同的流水线 K/V 流。

类似的协同设计思想已在 X 的开源推荐系统 XAI Phoenix [4] 中得到探索。

我们原型化了一个融合内核来量化融合收益,排除了 K/V 投影节省(图 13,右)

  • seq_len = 512:    6.6% 提升 (514 对比 482 TFLOPs/s)
  • seq_len = 1,024:  4.1% 提升 (581 对比 558 TFLOPs/s)
  • seq_len = 2,048:  0.3% 提升 (612 对比 610 TFLOPs/s) — 自注意力使 SM 饱和

短序列上的收益源于内核融合优势:减少启动开销、共享缓冲区分配节省、跨内核流水线机会以及波量化缓解 — 这些正是 LLM 推理中超大内核(megakernel)技术 [12] 针对的低效之处。在生产中,共享的 K/V 投影在线性投影成本上提供了额外节省,类似于 KV 缓存复用。

4. 基准测试和结果总结

我们总结了本文中提出的内核级基准测试以及端到端部署成果。以下所有内核基准测试均在 H100 SXM5 上(详细信息见附录 1)。

  • 线性压缩(第 2 节)。 四个渐进式协同设计阶段 — 矩阵乘法分解、内存对齐、广播融合以及通过 TLX 进行的线程束专用多阶段融合 — 在典型设置下产生累计 ~4 倍加速(1.944 ms → 0.482 ms)。收益在不同批次大小和候选对象/用户比率下保持稳健(图 9)。
  • Flash Attention(第 3 节)。 IKBO 将目标注意力从 IO 受限(~60 FLOPs/Byte)转变为计算受限(~833 FLOPs/Byte),实现了非协同设计 CuTeDSL FA4-Hopper(仅内核/内核+广播)吞吐量的 2.4 倍/6.4 倍,达到 621 BF16 TFLOPs。
  • 端到端部署。 IKBO 已广泛部署在 Meta 的推荐系统推理堆栈中 — 从早期到晚期的排序模型,在 GPU 和 MTIA 加速器上 — 在协同设计的模型上实现了高达 2/3 的计算密集型网络延迟降低。IKBO 已在候选对象/用户广播比率从 ~10,000 : 1 到 ~10 : 1 的范围内得到验证,证实了其在工作负载中的数值稳定性和可扩展性。

5. 结论和未来方向

IKBO 表明广播 — 长期以来被视为用户-候选对象交互不可避免的成本 — 可以通过内核-模型-系统协同设计在计算原语层被消除。通过将广播语义直接编码到内核中,不会物化冗余张量,且节省自然地随候选对象/用户比率扩展。

虽然这项工作中提出的内核实现通过 Triton 和 TLX 针对 NVIDIA Hopper,但核心思想 — 用索引驱动的内核内查找替换物化的广播 — 是与硬件供应商无关的。将 IKBO 内核适应到 CuTeDSL(用于先进的 NVIDIA 后端支持)并完成 AMD CK 支持是自然的下一步。

除了此处提出的两级用户-候选对象层次结构外,一些推荐系统场景涉及更深的层次 — 例如,用户 → 广告商 → 广告项目,其中每个用户看到多个广告商,每个广告商提供多个项目。这引入了两个具有独立、非均匀比率的嵌套广播关系。IKBO 可以优雅地处理这一点,将其应用于多级工作负载是进一步减少生产推荐系统架构中物化开销的自然方向。

致谢

我们感谢 Hongtao Yu, Yuanwei (Kevin) Fang, Daohang Shi, Yueming Hao, Srivatsan RameshManman Ren 对 Triton 和 TLX 基础、强大的 Triton 分析工具的支持,以及在整个工作中及时解决与 Triton 相关的问题。

感谢 Chris Gottbrath 的深刻反馈,显著提高了本文的清晰度。我们也衷心感谢他协助促进顺利的评审过程。

感谢 Santanu Kolay, Sandeep Pandey, Matt Steiner, GP Musumeci, Ashwin Kumar, Ian Barber, Aparna Ramani, CQ Tang 的领导支持。

参考文献

[1] Naumov, M., et al. “Deep Learning Recommendation Model for Personalization and Recommendation Systems,” arXiv:1906.00091, 2019.

[2] Wang, R., et al. “Deep & Cross Network for Ad Click Predictions,” ADKDD, 2017.

[3] Zhai, J., et al. “Actions Speak Louder than Words: Trillion-Parameter Sequential Transducers for Generative Recommendations,” ICML, 2024.

[4] xAI. “Phoenix: Recommendation System,” GitHub, 2026. https://github.com/xai-org/x-algorithm

[5] Guo, L., et al. “Request-Only Optimization for Recommendation Systems,” arXiv:2508.05640, 2025.

[6] Zhang, B., et al. “Wukong: Towards a Scaling Law for Large-Scale Recommendation,” ICML, 2024.

[7] Vaswani, A., et al. “Attention Is All You Need,” NeurIPS, 2017.

[8] Lee, J., et al. “Set Transformer: A Framework for Attention-based Permutation-Invariant Input,” ICML, 2019.

[9] Dao, T. “FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning,” ICLR, 2024.

[10] Shah, J., et al. “FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision,” NeurIPS, 2024.

[11] Zadouri, T., et al. “FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling,” arXiv:2603.05451, 2026.

[12] Spector, B., et al. “Look Ma, No Bubbles! Designing a Low-Latency Megakernel for Llama-1B,” Hazy Research Blog, 2025. https://hazyresearch.stanford.edu/blog/2025-05-27-no-bubbles

附录

附录 1. 基准测试设置

所有实验均在单个 NVIDIA H100 SXM5 GPU(700 W TDP, 96 GB HBM2e)上进行,具有以下软件堆栈

  • CUDA: 12.4
  • PyTorch: 2.11.0a0+fb(内部构建)
  • Triton: facebookexperimental/triton@4059e79bf (#831)

附录 2. 算术强度分析

2.1 H100 SXM5 的机器平衡点(700 W TDP, 96 GB HBM2E)

2.2 基准 LCE 的算术强度

对于 FP16 中的批次矩阵乘法 (M, K) @ (B, K, N) → (B, M, N),其中 B=1024, M=433, K=2044, N=256:

附录 3. 第 2.1 节的详细结果分析

设置: H100 SXM5(附录 1),PyTorch 饥饿模式(无内核融合),推理。来自典型配置的形状。

版本 总计 

(ms)

内核 延迟

(ms)

DRAM

(GB)

L1/TEX 扇区

(M)

计算

(GFLOPs)*

瓶颈

基准 1.944 1 CUTLASS GEMM 1.944 1.31 798 460 L1/TEX (89%) 
分解 1.389 2 CUTLASS GEMM(用户 + 候选矩阵乘法) 0.984 0.68 351 200 L1/TEX (84%)
1 ATen Gather + 1 ATen add 0.405 0.87 36 0.11 DRAM (92%)

*执行的总 FLOPs,非吞吐量。
†瓶颈通过 NCU 速度之光(Speed of Light)分析识别;方法见附录 4。

去重消除了 >98% 的用户侧工作(批次 1024 → ~15),将 L1/TEX 扇区从 798M 削减至 351M,将 GEMM 延迟从 1.944 毫秒削减至 0.984 毫秒。GEMM 后广播和相加成本为 0.405 毫秒(DRAM 受限),产生 0.555 毫秒的净节省。

精度说明。 基准在单个 FP32/TF32 缩减中累积所有 K 乘积。分解分别累积 K_user 和 K_cand,然后在 BF16/FP16 中对部分结果求和。训练使用相同的分解,因此数值端到端匹配。为了精确的推理一致性,融合内核(第 2.4 节)可以在 FP32 中执行最终求和。

附录 4. 瓶颈分析方法

在 Roofline 分析后进行更仔细的查看时,我们使用 NCU 的速度之光分析来识别硬件子系统瓶颈。瓶颈是利用率相对于其峰值持续吞吐量最高的子系统。对于第 2.1 节中的分析,我们监控三个指标

计算是峰值 SM 流水线利用率,直接由 NCU(Compute (SM) Throughput)报告。它测量最活跃的执行流水线(用于 GEMM 的张量核心)相对于其峰值指令速率有多忙。

L1/TEX 利用率根据 L1/TEX 单元必须处理的总扇区数导出,如下所示,其中 num_L1_tex_sectorsl1tex__t_sectors_pipe_lsu_mem_global_op_ld.sum_st.sum 计数器,SM_active_cycles sm__cycles_active.avg 计数器,num_SM 为 132,num_sustained_peak_sectors_per_sm_per_cycle 在 H100 上为 2.0。

DRAM 利用率根据传输的总 HBM 字节导出,如下所示,其中 dram_bytes_read_and_write 是 dram__bytes_read.sum 和 dram__bytes_write.sum 计数器。peak_bandwidth 在测试 GPU 服务器上为 2TB/s。

附录 5. 第 2.2 节的详细结果分析

结果。 1.389 毫秒 → 0.798 毫秒(减少 42.5%)。

版本 总延迟

(ms)

内核 延迟

(ms)

DRAM 流量

(GB)

计算

(GFLOPs)

*非速度

L1/TEX 扇区

(M)

瓶颈

分解

(未填充)

1.386 2 CUTLASS GEMM – 用户 & 候选矩阵乘法 0.984 0.68 200 351 L1/TEX (84%)
1 ATen Gather – 广播

1 ATen Elementwise – 添加

0.402 0.87 0.11 36 DRAM (92%)
分解

(填充 K)

0.798 2 CUTLASS GEMM – 用户 & 候选矩阵乘法 0.400 0.69 200 0 平衡
1 ATen Gather – 广播

1 ATen Elementwise – 添加

0.398 0.87 0.11 36 DRAM (92%)

大加速背后的两个因素。

  • TMA。 使用对齐矩阵,CUTLASS 选择基于 TMA 的内核,完全绕过 L1/TEX(扇区 → 0)。未填充的内核还对矩阵 B 进行了不必要的惩罚:它对两个矩阵应用了 4 字节加载,尽管 B(具有对齐的 N)本可以使用 128 位加载。
  • Bank 冲突。 未填充的内核还使用了 sm80 MMA 路径,其混合模式(swizzle pattern)不能保护 4 字节 cp.async 写入,导致许多共享内存 bank 冲突。填充的内核没有这个问题。

附录 6. 第 2.3 节的详细结果分析

结果。 延迟:0.798 毫秒 → 0.580 毫秒(减少 27.4%)。

版本 总延迟

(ms)

内核 延迟

(ms)

DRAM 流量

(GB)

分解

(填充 K)

0.798 2 CUTLASS GEMM – 用户 & 候选矩阵乘法 0.400 0.68
1 ATen Gather – 广播

1 ATen Elementwise – 添加

0.398 0.87
iKBO 融合 0.580 用户 GEMM & 候选 iKBO 内核 0.580 0.68

0.87 GB 的中间 DRAM 流量按预期被消除。NCU 分析揭示了进一步的机会:占用率仅为 6.25%,每个调度程序 1 个线程束,PC 采样显示仅 23% 的周期是生产性的

停顿原因 百分比 内核中主要指什么
停顿长记分板 41.8% 全局内存加载
已选(正在执行) 23.1% 生产性工作(好) – 实际发出的指令
停顿等待 20.1% 等待 WGMMA
停顿屏障 5.7% 软件流水线阶段之间的 bar.sync

每个调度程序 1 个线程束时,每个停顿都完全暴露:没有其他线程束可以切换到。通过减少流水线深度来增加占用率会牺牲 K 循环延迟隐藏。对于此内核来说,这是一个具有挑战性的情况:需要大的平铺和深流水线来保持张量核心吞吐量,但它们消耗了大部分共享内存预算,几乎没有空间通过占用率来隐藏延迟。

附录 7. 发布-获取同步协议

生产者(用户 CTA)。 将用户平铺存储到全局内存后,CTA 设置具有发布语义的平铺标志,确保标志写入前的数据可见性:

tl.atomic_add(user_tile_flag_ptr, 1, sem="release", scope="gpu")

消费者(候选对象 CTA)。 每个线程束组的单个线程以 ld.relaxed 轮询标志,以最大限度地减少自旋期间的内存流量。一旦标志转换,单个 ld.acquire 建立发生前边,命名屏障向线程束组中的所有 128 个线程广播就绪状态:

if tlx.thread_id(axis=0) % 128 == 0:  # 1 thread per warp group (4 warps)
    ready = tl.inline_asm_elementwise(
        "ld.relaxed.gpu.global.b32 $0, [$1];", "=r,l",
        [user_tile_flag_ptr], dtype=tl.int32, is_pure=False, pack=1)
    while ready == 0:
        ready = tl.inline_asm_elementwise(
            "nanosleep.u32 50; ld.relaxed.gpu.global.b32 $0, [$1];", "=r,l",
            [user_tile_flag_ptr], dtype=tl.int32, is_pure=False, pack=1)
    tl.inline_asm_elementwise(
        "ld.acquire.gpu.global.b32 $0, [$1];", "=r,l",
        [user_tile_flag_ptr], dtype=tl.int32, is_pure=False, pack=1)
tlx.named_barrier_wait(12, 128)

附录 8. TLX 对比 Triton 的 NCU 分析指标

指标 Triton TLX 注意事项
理论占用率 6.25% 18.75% 每个 CTA 3 个线程束组对比 1 个
DRAM 吞吐量

(dram__cycles_active.avg.pct_of_peak_sustained_elapsed)

38.51% 52.39% 连续 TMA 加载带来的更高利用率
L2 缓存吞吐量

(lts__throughput.avg.pct_of_peak_sustained_elapsed)

73.69% 83.86% 瓶颈。TLX 接近峰值

附录 9. 普通 Flash Attention 对比 IKBO Flash Attention 的 Roofline 分析

算术强度 (AI) 计算基于 FP16/BF16 精度,user_seq_len = 1024, n_seed = 64, B_candidate (等式中的 B) : B_user (等式中的 B/num_cand_user) = 70: 1。

附录 10. IKBO TLX FA3 的 SMEM 消耗

SMEM 缓冲区 计数 块维度 总大小
查询 2(每个消费者组 1 个) 64 * 128 (2字节) 32KB
2 128 * 128 (2字节) 64KB
2 128 * 128 (2字节) 64KB
输出 2(每个消费者组 1 个) 64 * 128 (2字节) 32KB
总计 192KB

附录 11. 在推荐系统边界条件下对 IKBO FA 对比 CuTeDSL FA4 Hopper 和 TLX FA3 Hopper 内核进行基准测试

IKBO 内核基本启用了用户-候选对象交互映射逻辑,该逻辑共享与 GQA 相似的 IO 和计算模式。在基准测试期间,为 IKBO 内核应用了稳定的 B_candidate : B_user = 64 : 1,并为 CuTeDSL FA4 Hopper GQA 版本应用了相似的计算模式(Q_seq_len = 128 以确保 2 消费者线程束组完美工作)。值得额外提及的是,IKBO 内核仍需额外消耗候选对象-用户映射张量,以处理实时要排序的各种候选对象数量。

内核类型 吞吐量 (TFLOPs/s) IO (GB/s)
Triton IKBO FA2 425 519
TLX IKBO FA3 418 510
TLX IKBO FA3 持久化 592 723
TLX IKBO FA3 持久化 V2 (反转 k,v 顺序) 537 655
CuTeDSL FA4 Hopper GQA 518 633
TLX FA3 GQA 576 703

IKBO FA 对比开源 GQA 内核进行基准测试。IKBO 内核的 Q, K, V 形状按 [Batch size, num head, seq, d_head] 序列:Q_ikbo [2048, 2, 64, 128], K/V_ikbo [32, 2, 1024, 128]。GQA 内核的 Q, K, V 形状:Q_gqa [1024, 2, 128, 128], K/V_gqa [32, 2, 1024, 128] 

内核类型 吞吐量 (TFLOPs/s) IO (GB/s)
Triton IKBO FA2 449 329
TLX IKBO FA3 470 345
TLX IKBO FA3 持久化 621 455
TLX IKBO FA3 持久化 V2 (反转 k,v 顺序) 587 430
CuTeDSL FA4 Hopper GQA 608 445
TLX FA3 GQA 628 460

IKBO FA 对比开源 GQA 内核进行基准测试。Q_ikbo [2048, 2, 64, 128], K/V_ikbo [32, 2, 2048, 128]。GQA 内核的 Q, K, V 形状:Q_gqa [1024, 2, 128, 128], K/V_gqa [32, 2, 2048, 128] 

注意:由于标准 Flash Attention 内核不包含 IKBO 逻辑,我们使用具有相似 IO 成本和 FLOPs 消耗的 GQA 配置来模拟 cuteDSL 版本的吞吐量结果。

附录 12:指令缓存缺失导致消费者-2 线程束组出现显著延迟

图 A1 

修复前后的指令缓存缺失结果

Before instruction cache miss fix:
    ---------------------------------------------------- ----------- ------------
    Metric Name                                          Metric Unit Metric Value
    ---------------------------------------------------- ----------- ------------
    gcc__cache_requests_type_instruction.sum                              319,394
    gcc__cache_requests_type_instruction_lookup_miss.sum                    7,234
    sm__icc_requests.sum                                       cycle    6,049,376
    sm__icc_requests_lookup_hit.sum                            cycle    5,438,421
    sm__icc_requests_lookup_miss.sum                           cycle      610,955
    ---------------------------------------------------- ----------- ------------

After instruction cache miss fix:
    ---------------------------------------------------- ----------- ------------
    Metric Name                                          Metric Unit Metric Value
    ---------------------------------------------------- ----------- ------------
    gcc__cache_requests_type_instruction.sum                               33,008
    gcc__cache_requests_type_instruction_lookup_miss.sum                      769
    sm__icc_requests.sum                                       cycle      792,437
    sm__icc_requests_lookup_hit.sum                            cycle      722,244
    sm__icc_requests_lookup_miss.sum                           cycle       70,193
    ---------------------------------------------------- ----------- ------------