前途科技前途科技
  • 服务
  • 关于
  • AI
    • AI 大模型
    • 具身智能
    • 算力芯片
  • 科技
    • 智能终端
    • 软件·互联网
    • 汽车·出行
    • 科学前沿
  • 资源中心
    • 深度研究
      • AI 前沿
      • 教程
      • AI 知识库
      • 案例研究
    • 行业报告
      • 白皮书
      • 行业报告
      • 研究报告
      • 技术分享
      • 专题报告
    • 精选案例
      • 金融行业
      • 医疗行业
      • 教育行业
      • 零售行业
      • 制造行业
  • 服务
  • 关于
联系我们
教程/10.07 · 06:57/12 MIN/作者:苏晚/1 阅读

Triton 为什么跑赢 CUDA:Meta 重写推荐系统嵌入算子的关键一步

Meta 把推荐系统最重的 TBE 算子从 CUDA 模板重写成 Triton,GB200 上 307 个生产分片的反向带宽中位数领先 13%~22%。收益大头不是编译器更聪明,而是把协作内核的切换点从段长 32 挪到 256——在 Triton 里这只是改一个配置值。

2026 年 10 月 6 日,PyTorch 官方博客发布了 Meta 团队(Daohang Shi、Oleksandr Stashuk、Rupert Wu、Liangbei Xu、Rich Zhu)的《Modernizing Table Batched Embeddings with FBTriton》(原文:https://pytorch.org/blog/modernizing-table-batched-embeddings-with-fbtriton/)。下文的数据、参数和代码片段均出自原文。

这篇文章讲的是 Meta 把推荐系统里最重的一类算子——TBE(Table-Batched Embedding)的前向和反向——从 CUDA 模板重写成了 Triton,并且在生产负载上跑赢了原来的 CUDA 实现。

"Triton 跑赢手写 CUDA"这种标题很容易被读成"编译器终于比人聪明了"。原文给的解释恰恰相反:反向传播的大部分收益,来自一个切换阈值——CUDA 在段长 32 时切到协作式内核,Triton 一直用简单流式路径撑到 256。快的不是 Triton 本身,而是在 Triton 里改这个数只是改一个配置值,在模板生成的 CUDA 里却意味着重组整套内核分工。这才是这篇文章对其他团队最有参考价值的地方。

一、TBE 是什么,为什么它难

推荐模型的输入里有大量稀疏特征:用户点过的商品 ID、关注的作者 ID、看过的视频 ID……每类特征对应一张嵌入表(embedding table),每个 ID 查表取出一行向量,同一个样本里的多个 ID 再池化(通常是求和)成一个向量。生产模型里这样的表有成百上千张,分片在大量 GPU 上。

TBE 做的事,就是把许多张表的查找和池化合并进一次 GPU 启动里完成,省掉逐表启动的开销、提高显存访问效率。原文开头说,这类算子要处理跨"数千块分片 GPU"的嵌入查找。

原文用的记号先列在这里,后面会反复出现:

记号含义
T表的数量
E表的行数(嵌入表大小)
D嵌入维度(列数)
L池化因子,即一个 bag 里有几个 ID
Bbatch 大小,即每次请求的 bag 数

前向很直观:给定一个稀疏特征 [id1, id2, id3, id4](L=4,B=1),输出就是 T[id1%E] + T[id2%E] + T[id3%E] + T[id4%E]。B=100 就有 100 个输出。

反向才是麻烦所在。目标是对每个被访问过的 ID,把所有包含它的 bag 的输出梯度加起来,再做一次优化器更新(最简单的形式就是 weight = weight - grad * LR)。原文给了一组真实规模的统计:在 T×B=4.2M、B=128K 时,索引总数、去重后数量、最高频 ID 的出现次数可以达到 83M / 3M / 125K。

这组数字说明了推荐场景的本质难点:8300 万次访问只落在 300 万个不同的行上,而最热的那一行被访问了 12.5 万次。热门商品、头部创作者的 ID 天然集中,长尾 ID 则大多只出现一两次。原文的判断是,TBE 反向没有任何张量核运算,但数据搬运和归约很重,而且受负载不均困扰。它不是算力问题,是访存和调度问题。

二、前向:一个通用路径,加一条窄窄的快路径

TBE 前向的两条执行路径

TBE 前向按特征分派:默认走通用 gather 路径,符合条件的小表转成计数矩阵后走张量核。图片来源:Meta Team / PyTorch Blog

前向对每张表、每个 bag 取出索引到的行,可选地乘上逐样本权重,累加后写出一个 D 维的池化结果。Meta 做了两种实现。

通用 gather 路径的几个关键设计:

  • 网格:启动 ceil(B / BAGS_PER_PROGRAM) 个 program,每个 program 在内部循环 T 个特征,而不是开一个 B×T 的网格。
  • gather 宽度:内层循环一次发出 4 个独立的行加载;调优过的双 bag 路径发 8 个。
  • 每个 program 处理几个 bag:大规模、非 VBE、非 FP32 负载用 2 个;拆出直方图特征后,剩余的通用特征区段用 4 个;其他形状用 1 个。
苏晚
苏晚Su Wan

前途科技 · 前沿主笔

关注 AI 与前沿科技的观察者,前途科技前沿内容主笔。

查看全部文章

想了解 AI 如何助力您的企业?

免费获取企业 AI 成熟度诊断报告,发现转型机会

置顶文章

Triton 为什么跑赢 CUDA:Meta 重写推荐系统嵌入算子的关键一步
置顶

Triton 为什么跑赢 CUDA:Meta 重写推荐系统嵌入算子的关键一步

新硬件接入 PyTorch:从补丁地狱到标准化工作流
置顶

新硬件接入 PyTorch:从补丁地狱到标准化工作流

前途科技前途科技
服务关于快讯技术商业报告
微信咨询二维码

咨询微信

扫码添加微信咨询

前途科技微信公众号

微信公众号

扫码关注

Copyright © 2026 AccessPath.com, 前途国际科技咨询(北京)有限公司,版权所有。|京ICP备17045010号-1|京公网安备 11010502033860号|隐私政策|服务条款
  • 索引位宽:当线性化后的范围小于 2^31 时,TorchRec 支持通过配置使用 int32 的 indices 和 offsets。存储减半,CUB 基数排序的 key 宽度也减半。默认仍是 int64。
  • 累加精度:FP16/BF16 权重用 FP32 累加;FP32 权重用 FP64 累加,以保证大 D 时的精度。
  • 小表直方图 + 张量核路径的触发条件写得非常死:单个特征、E≤64、64≤D≤128、L≥64、FP16 权重、FP32 输出、无逐样本权重、非 VBE。满足时一个 program 处理 16 个 bag,对前 256 个索引建直方图,用 tl.dot 计算"计数 × 表",剩余索引走标量路径。

    这条路径的思路值得单独说一句:表只有 64 行以内、而每个 bag 有 64 个以上 ID 时,同一行必然被反复命中。与其逐个 gather 再相加,不如先数清每行出现几次,把稀疏查找变成一次稠密的矩阵乘——这样就能用上原本在 TBE 里完全闲置的张量核。代价是条件极窄,图里底部那行字也直说了:资格是刻意收窄的,每个特化都有安全的回退路径。

    越界检查:独立路径在 Triton 前向之前跑一个更新过的 CUDA 校验步骤,在 B200 上这一环节最多快 1.24 倍。开启 fused_bounds_check 后,符合条件的内核会在查找过程中顺带校验并修复非法输入;offset 的校验修复仍是单独的小内核。加权表、变长 batch、AMD 以及 transpose 前移等配置会回退到标准校验。这个选项通过 TorchRec 暴露,默认关闭。

    三、前向变慢、总时间变快:一笔要算总账的交易

    前向还有两项"状态复用"的设计,都是为反向服务的。

    第一项:精确的 row-wise Adagrad 可以把前向的直方图存下来,反向直接用这些计数做补偿式的 FP16 高/低位 GEMM(累加到 FP32),再执行优化器。前向和反向仍是两次独立启动,复用的只有直方图计数。

    第二项更有意思:一个可选的核心模块路径把**索引转置、排序和游程编码(RLE)**挪进前向,通过 autograd 把元数据传给反向。原文给的 B200 大配置数据:

    阶段改动前改动后
    前向22.844 ms33.252 ms
    反向56.693 ms32.931 ms
    合计79.537 ms66.183 ms(−16.8%)

    前向慢了约 10 毫秒,反向快了约 24 毫秒,总账省下 16.8%。这组数字对做性能评估的团队是个提醒:只盯前向 benchmark 的话,这项优化会被判成"负优化"直接否掉。训练场景里前向和反向总是成对出现,衡量指标应该是一次完整迭代。

    需要注意的是,这条路径默认关闭,而且目前的 TorchRec 封装没有暴露它。

    四、反向:一切围绕"段长"展开

    索引转置与游程编码

    transpose_embedding_input 把"每个样本读了哪些行"倒过来,变成"每一行被哪些样本读过"。图片来源:Meta Team / PyTorch Blog

    反向的核心逻辑用一句话概括:对 batch 中被访问过的每一个唯一的(表,行)对,把所有访问过它的位置上的上游梯度加起来,然后对这一行恰好做一次优化器更新。

    batch 原本的组织方式是"每个样本读了哪些行",反向需要的是反过来的视图。transpose_embedding_input 做的就是这件事:线性化、排序、游程编码,把 batch 变成一个个 run——每个 run 是一个唯一的行,加上所有碰过它的样本。不同 run 不会写同一行,所以每个 run 是独立的工作单元,可以用一次读改写代替原子操作。前面说的"挪进前向",挪的就是这一步,让它离开反向的关键路径。

    一个 run 里有多少样本,叫段长(SL, segment length)。原文说这是所有设计都围着转的变量:同一个 batch 里,SL 从 1 到数百万都有。回看那组 83M / 3M / 125K 的统计,就能理解这个跨度从何而来。

    三种反向内核及路由方式

    run 按段长分到三种反向内核,长 run 那一档有两种实现。图片来源:Meta Team / PyTorch Blog

    run 按 SL 被路由到三种内核:

    • short_run(SL < 256):一个 program 从头到尾负责一个 run——gather、在寄存器里累加、执行优化器、写回。独占意味着普通 store 即可,不需要原子操作。
    • grad_accum + apply(SL ≥ 256,默认):run 被切成每块 256 次查找的小块,部分和落进 workspace,再由第二个内核执行优化器。
    • fused(SL ≥ 256,Blackwell,超大 batch):同样切块,但借助设备级 fence,最后一个子 program 在同一次启动里完成更新。图中的判断条件是 Blackwell 且查找量超过 2 亿。

    加权表的区别只在于累加前把每行梯度乘上对应的逐样本权重。

    五、六个工程问题,每个都是具体的坑

    原文第四节列了一串边界情况。它们的共同点是:在小规模测试里都不会出现,一上生产形状就冒出来。

    负载不均。 一个 program 走 200 万次查找时,GPU 其他部分在空等;而一大串只有 1 次查找的 run,每个都要付全套的单 run 开销。第一个解法是把长 run 切开:达到阈值的 run 被切成固定 256 次查找的块,类似 split-K。一个 200 万次查找的 run 会变成大约 8000 个子 program,单行的工作量就足以填满整台机器。第二个解法是 Blackwell 上的 CLC(通过 TLX):内核启动后常驻,处理完当前 run 后继续"偷"空闲的 run_id,偷不到才退出。原文的说法是,这让负载均衡从软件方案(比如按频率把 run_id 分两路处理)变成了硬件方案。

    gather 宽度是一道寄存器悬崖。 每个缓冲的行都要让 BLOCK_SIZE 个 64 位地址保持活跃,因为 dout_row_start_ptr[:, None] + col_offsets[None, :] 会物化出一整块指针。代价是"宽度 × BLOCK_SIZE",在一种行宽下调好的宽度换到另一种行宽就是错的。解法是每一档都按目标单独配置宽度。B200 实测:

    档位宽度寄存器占用率效果
    短 run,不加权8 → 2184 → 6412.5% → 49.9%0.41 → 1.05
    长 run,累加8 → 2158 → 6217.6% → 44.7%整体持平率 82% → 87%
    短 run,加权4 → 2125 → 6424.8% → 49.3%加权持平率 51% → 69%

    第一行最说明问题:宽度从 8 降到 2,相对 CUDA 的比值从 0.41 变成 1.05,从明显落后变成略微领先。一次发出的加载"更少"反而更快,因为寄存器压力降下来后,占用率提高到原来的约四倍。

    BLOCK_SIZE 是编译期常量。 一次启动只能按所有表里最宽的 D 取 next_pow2(max_D),如果查找量最大的那张表比最宽的表窄得多,每次 gather 的大部分 lane 都是被 mask 掉的空转。解法是按维度分桶:短 run 在分类阶段就按 next_pow2(D) 分到不同的桶,每个桶用自己的 BLOCK_SIZE 启动。分桶折叠进了分类内核,不额外多走一遍。原文特意说明,profiling 确认它换来的是更少的 lane 浪费,不是更高的占用率。

    别把形状读回 CPU。 哪些 run 是长的取决于数据,只有 GPU 知道。用 .item() 读回计数,等于每次反向都做一次 cudaStreamSynchronize。解法是让形状留在 GPU 上:workspace 按索引数算出的上界预分配,分类在内核里用原子计数器做流压缩:

    is_long        = (run_len >= threshold) & mask
    num_long_block = tl.sum(is_long.to(tl.int32))
    long_base      = tl.atomic_add(num_long_ptr, num_long_block)
    long_local     = tl.cumsum(is_long.to(tl.int32), axis=0) - 1
    tl.store(long_run_ids_ptr + (long_base + long_local).to(tl.int64),
             offsets.to(tl.int32), mask=is_long)
    

    每个 block 只做一次原子操作而不是每个元素一次,block 内的偏移靠前缀和算出。后续内核把计数当作设备指针读取,用 while 循环自行分配工作。

    切开的 run 需要跨 program 的屏障。 run 被切开之后,优化器更新必须等所有子 program 的部分和都落地才能做。Triton 原本没有设备级 fence,所以只能多一个内核、多一次全局内存往返。TLX 在 Blackwell 上暴露了 fence,最后一个子 program 可以在同一次启动里完成更新:

    tl.atomic_add(temp_grad_buffer_ptr + temp_grad_offset + col_offsets, grad, mask=mask)
    tlx.fence("gpu")
    remaining = tl.atomic_add(grad_accum_counter_ptr + grad_buffer_id, -1)
    if remaining == 1:
        ...  # 最后一个子 program 执行优化器并写回
    

    原文强调,正确性就靠这里的顺序:fence 保证每份部分和在倒计数递减之前已经对全设备可见,所以看到 remaining == 1 的那个 program 读到的是完整的和。少了 fence,一个 program 可能在另一个的 atomic_add 还在途时就赢下倒计数——结果是静默的数值错误,不报错、不崩溃,只是梯度不对。而会被切块的恰恰是 SL ≥ 256 的热门行。

    合并部分和是一整行的原子操作。 每个子 program 都要把一整行 BLOCK_SIZE 原子地加进 run 的 workspace 槽位。一个 run 切成 8000 块,就是 8000 个 program 争抢同一行,而 tl.atomic_add 是逐元素一次次发出的。Blackwell 有更合适的指令 cp.reduce.async.bulk.tensor,TLX 以 tlx.async_descriptor_store(..., store_reduce="add") 暴露,通过 TMA 做归约。

    六、结果:阈值从 32 挪到 256

    Triton 与 CUDA TBE 反向带宽比分布

    GB200 上 307 个生产分片的反向带宽比(Triton / CUDA TBE),精确 row-wise Adagrad,FP16 权重。左为不加权,右为加权。图片来源:Meta Team / PyTorch Blog

    反向的测试覆盖 GB200 上 307 个分片配置(283 种不同形状),精确 row-wise Adagrad,FP16 权重。从图上读:不加权的 256 个分片中位数 1.13,94% 达到或超过持平;加权的 51 个分片中位数 1.22,69% 达到或超过持平。

    前向方面,原文正文写的是"中位前向加速 1.28×",但对应的图表标注的是 B200 上两个模型各 31 次测量:Model 1 中位数 1.28,Model 2 中位数 1.18,两者都是 100% 达到或超过持平。正文把这个 1.28 和 307 个 GB200 分片写在了同一句里,读的时候最好以图为准,把前向和反向的测试条件分开看。

    两种实现切换策略的位置

    CUDA 在 SL=32 切到协作式内核,Triton 在 256 之前一直走简单路径。图片来源:Meta Team / PyTorch Blog

    为什么 Triton 能赢?原文的回答是:大部分收益来自改变切换点。

    CUDA TBE 在 SL 达到 32 时升级到 cta_per_row——一个协作式的 CTA(Cooperative Thread Array)负责一行;Triton 则一直用 short_run 的简单流式处理撑到 SL=256。于是三个区间的结果分别是:SL < 32 两边都走简单路径,接近持平;SL ≥ 256 两边都走重型路径,持平;中间 32 ≤ SL < 256 这个"争夺区间",CUDA 的协作内核主要在付同步开销,Triton 的中位带宽比是 1.76 倍。图下那句话说得很直白:整个结果取决于两个阈值之间 8 倍的差距。

    Nsight Compute 对比

    GB200 上一个处于争夺区间的生产分片,用 ncu 测量承载这一形状的两个内核。图片来源:Meta Team / PyTorch Blog

    Nsight Compute 的数据把差别定位到了访存吞吐上。承载这个形状的 CUDA 内核达到 678 GB/s,Triton 内核达到 3,948 GB/s,是 5.8 倍,而两者的占用率几乎相同(49.8% 对 49.9%)。图中的内核耗时是 730.9 ms 对 182.0 ms;原文正文另写了"快 4.3 倍",与图中的 4.0 倍不一致,原文没有解释差异来源。整个反向 pass 里,Triton 搬运的 DRAM 字节是 CUDA 的 0.91 倍(777.4 GB 对 850.6 GB),发出的全局加载请求反而多 29%(13,202M 对 10,215M)。

    原文的解释是:SL 刚过 32 时,run 太短,CTA 级同步没有足够的工作量去摊薄,而 Triton 的路径里根本没有协作。换句话说,CUDA 那边不是写得差,而是在这个区间选了一个"为长 run 设计"的策略。

    Triton 仍然输的地方,原文也写了。工作全部落在 SL < 4 的 run 里的形状低于持平,原因是纯粹的依赖加载延迟——CUDA 的 warp-per-row 更擅长摊薄每个 run 的元数据开销。这样的分片有 11 个(共 307 个),每个耗时都在约 1 毫秒以下。SL ≥ 256 的形状只是持平,没有领先,团队表示等它们在实际负载上成为瓶颈再投入。

    七、真正的收益是"改得动"

    原文有一句结论很值得记:收益来自配置值。 前面那些修正,每一项在 Triton 里都表现为一次配置改动;而在模板生成的 CUDA 里,挪一个升级点或改一个 gather 宽度,意味着重新划分哪个内核处理什么。

    把这句话和第六节放在一起读:原文没有交代 CUDA TBE 的 32 是怎么定下来的,但它给出的证据指向同一个方向——争夺区间里的差别不在内核写得好坏,而在切换点放在哪里,而在 CUDA 模板里挪这个点是一次结构性重构,在 Triton 里只是改一个数。Triton 版本赢下的那 1.76 倍,可以看作迭代成本降低之后换来的回报。对维护大型手写内核的团队来说,这值得回头检查一遍:代码里那些"一直这么设"的阈值,有多少是因为改起来太贵才没人去试。

    原文在最后一节列出的几点长期价值,也都围绕这一点:

    • 开发效率:整个稀疏路径(前向加反向)现在是普通 Python 写的,体量比原来单是 CUDA 模板部分还小。相比 CUDA Jinja 模板,排序模型和基础设施工程师更容易快速实现新的嵌入算法。
    • 可移植性:在 Blackwell、Hopper、AMD 上用的是相同的内核主体,CLC、TMA 批量原子归约、设备级 fence 这些硬件特性作为附加开关接入,而不是为每种硬件分叉一份代码。
    • 通往"超级稀疏内核":用 Triton 重写之后,[前向, 反向] 和 [prologue, epilogue] 之间有大量融合空间。把优化器当作反向的 epilogue,就省掉了额外的内存遍历;原文设想最终把整条稀疏流程合进一个内核,像 FP8 momentum scaling 这样的复杂特性也只需几行累加代码。

    八、对照自己的环境:哪些能拿来用

    原文的实测全部在 B200 / GB200 上完成,下面按"现在能不能直接用"把它提到的东西分一下类,方便对照自己的环境。

    通过 TorchRec 可以直接配置的:

    • fused_bounds_check:通过 TorchRec 暴露,默认关闭。加权表、变长 batch、AMD、transpose 前移这几种配置不走融合校验,会回退到标准校验。
    • int32 indices / offsets:线性化范围小于 2^31 时可通过配置启用,索引存储和排序 key 宽度减半,默认仍是 int64。表的总行数是否越过这条线,是能不能开的前提。

    存在但默认关闭、目前 TorchRec 封装未暴露的:

    • 把转置、排序、RLE 挪进前向的路径。它在 B200 大配置上省了 16.8% 的总时间,但前向变慢约 46%。即便以后开放,评估时也要看完整迭代而不是单看前向。

    只在 Blackwell 上生效的:

    • CLC 持久化内核与 work stealing、TLX 设备级 fence 实现的单次启动 fused 反向、cp.reduce.async.bulk.tensor 的 TMA 归约。

    跟硬件无关、可以借鉴的思路:

    • 按段长分派,而不是对所有 run 用同一种策略;长 run 切块(split-K 式)来解决单点热行拖慢整体的问题。
    • 按目标调 gather 宽度,警惕寄存器压力把占用率压垮——表里那行 0.41 → 1.05 就是证据。
    • 动态形状留在 GPU 上,用"每 block 一次原子 + block 内前缀和"做流压缩,避开 .item() 带来的隐式同步。
    • 稀疏查找在小表、大池化因子时可以转成直方图乘表,借用张量核。

    如果你在 Hopper 或 AMD 上跑推荐模型,原文能提供的承诺是"同一套内核主体可以跑",但它没有给出这些硬件上的性能数据。Blackwell 专属的几项拿掉之后收益还剩多少,需要在自己的形状分布上实测。一个直接的起点是先统计自己负载的段长分布:如果大量 run 落在 32 到 256 之间,原文最主要的那份收益就有可能在你的场景里复现;如果大部分 run 的段长小于 4,那正是 Triton 目前还落后于 CUDA 的区间。

    表格预测搬进上下文:NVIDIA Kumo Tabular 不训练不调参
    置顶

    表格预测搬进上下文:NVIDIA Kumo Tabular 不训练不调参

    //

    24小时热榜

    纽约州因麻疹疫情爆发宣布进入全州灾难紧急状态
    TOP1

    纽约州因麻疹疫情爆发宣布进入全州灾难紧急状态

    Manus脱离Meta恢复独立:估值40亿美元与Agent新战局
    TOP2

    Manus脱离Meta恢复独立:估值40亿美元与Agent新战局

    3

    Mistral发布万亿参数大模型ML4,开源权重月底上线

    2小时前
    Mistral发布万亿参数大模型ML4,开源权重月底上线
    4

    美国宾夕法尼亚州麻疹确诊破千,已致5人死亡

    17小时前
    美国宾夕法尼亚州麻疹确诊破千,已致5人死亡
    5

    AMD巨资收购World Labs,芯片巨头打响世界模型争夺战

    22小时前
    AMD巨资收购World Labs,芯片巨头打响世界模型争夺战
    6

    AI让个人效率翻倍,为何团队交付还是快不起来?

    22小时前
    AI让个人效率翻倍,为何团队交付还是快不起来?
    7

    麦当劳回应AI算法操纵菜品定价指控

    22小时前
    麦当劳回应AI算法操纵菜品定价指控
    8

    DeepSeek被曝将完成超800亿元新融资

    2小时前
    DeepSeek被曝将完成超800亿元新融资
    热门标签
    大模型AgentRAG微调私有化部署Prompt EngineeringChatGPTClaudeDeepSeek智能客服知识管理内容生成代码辅助数据分析金融零售制造医疗教育AI 战略数字化转型ROI 分析OpenAIAnthropicGoogle

    关注公众号

    前途科技微信公众号

    扫码关注,获取最新 AI 资讯

    免费获取 AI 落地指南

    3 步完成企业诊断,获取专属转型建议

    已有 200+ 企业完成诊断