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 里却意味着重组整套内核分工。这才是这篇文章对其他团队最有参考价值的地方。
推荐模型的输入里有大量稀疏特征:用户点过的商品 ID、关注的作者 ID、看过的视频 ID……每类特征对应一张嵌入表(embedding table),每个 ID 查表取出一行向量,同一个样本里的多个 ID 再池化(通常是求和)成一个向量。生产模型里这样的表有成百上千张,分片在大量 GPU 上。
TBE 做的事,就是把许多张表的查找和池化合并进一次 GPU 启动里完成,省掉逐表启动的开销、提高显存访问效率。原文开头说,这类算子要处理跨"数千块分片 GPU"的嵌入查找。
原文用的记号先列在这里,后面会反复出现:
| 记号 | 含义 |
|---|---|
| T | 表的数量 |
| E | 表的行数(嵌入表大小) |
| D | 嵌入维度(列数) |
| L | 池化因子,即一个 bag 里有几个 ID |
| B | batch 大小,即每次请求的 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 前向按特征分派:默认走通用 gather 路径,符合条件的小表转成计数矩阵后走张量核。图片来源:Meta Team / PyTorch Blog
前向对每张表、每个 bag 取出索引到的行,可选地乘上逐样本权重,累加后写出一个 D 维的池化结果。Meta 做了两种实现。
通用 gather 路径的几个关键设计:
ceil(B / BAGS_PER_PROGRAM) 个 program,每个 program 在内部循环 T 个特征,而不是开一个 B×T 的网格。免费获取企业 AI 成熟度诊断报告,发现转型机会
小表直方图 + 张量核路径的触发条件写得非常死:单个特征、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 ms | 33.252 ms |
| 反向 | 56.693 ms | 32.931 ms |
| 合计 | 79.537 ms | 66.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 被路由到三种内核:
加权表的区别只在于累加前把每行梯度乘上对应的逐样本权重。
原文第四节列了一串边界情况。它们的共同点是:在小规模测试里都不会出现,一上生产形状就冒出来。
负载不均。 一个 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 → 2 | 184 → 64 | 12.5% → 49.9% | 0.41 → 1.05 |
| 长 run,累加 | 8 → 2 | 158 → 62 | 17.6% → 44.7% | 整体持平率 82% → 87% |
| 短 run,加权 | 4 → 2 | 125 → 64 | 24.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 做归约。

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 倍的差距。

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 倍,可以看作迭代成本降低之后换来的回报。对维护大型手写内核的团队来说,这值得回头检查一遍:代码里那些"一直这么设"的阈值,有多少是因为改起来太贵才没人去试。
原文在最后一节列出的几点长期价值,也都围绕这一点:
原文的实测全部在 B200 / GB200 上完成,下面按"现在能不能直接用"把它提到的东西分一下类,方便对照自己的环境。
通过 TorchRec 可以直接配置的:
fused_bounds_check:通过 TorchRec 暴露,默认关闭。加权表、变长 batch、AMD、transpose 前移这几种配置不走融合校验,会回退到标准校验。存在但默认关闭、目前 TorchRec 封装未暴露的:
只在 Blackwell 上生效的:
cp.reduce.async.bulk.tensor 的 TMA 归约。跟硬件无关、可以借鉴的思路:
.item() 带来的隐式同步。如果你在 Hopper 或 AMD 上跑推荐模型,原文能提供的承诺是"同一套内核主体可以跑",但它没有给出这些硬件上的性能数据。Blackwell 专属的几项拿掉之后收益还剩多少,需要在自己的形状分布上实测。一个直接的起点是先统计自己负载的段长分布:如果大量 run 落在 32 到 256 之间,原文最主要的那份收益就有可能在你的场景里复现;如果大部分 run 的段长小于 4,那正是 Triton 目前还落后于 CUDA 的区间。
关注公众号

扫码关注,获取最新 AI 资讯
3 步完成企业诊断,获取专属转型建议
已有 200+ 企业完成诊断