PyTorch 官方最新动态:Modernizing Table Batched Embeddings with FBTriton
来源:PyTorch 官方动态 | 发布日期:2026-10-1
核心更新概览
This post explores the FBTriton kernel design for Table Batched Embedding (TBE) forward and backward passes. These core operators handle embedding lookups across thousands of sharded GPUs within recom
详细内容记录
This post explores the FBTriton kernel design for Table Batched Embedding (TBE) forward and backward passes. These core operators handle embedding lookups across thousands of sharded GPUs within recommendation systems. Our Triton implementation successfully outperforms the legacy CUDA kernels on these workloads. Here, we explain the architectural design, detail the measured performance gains, and highlight future optimization opportunities. TBE (Table-Batched Embedding) kernel efficiently performs embedding lookups and pooling across many tables in one GPU operation. TBE combines embedding lookup and pooling for many tables in a single GPU launch, reducing launch overhead and improving memory efficiency. You can read more general info 2. Implementation of Triton TBE Forward For each table and bag, forward gathers the indexed rows, optionally multiplies them by per-sample weights, accumulates them in FP32 (FP64 for FP32 weights), and writes one D-wide pooled output. We built two implementations: a generic gather and a fast path implementation with a small-table histogram. programs. Each program loops over T features instead of launching a B×T grid. The inner loop issues four independent row loads. The tuned two-bag path issues eight Large non-VBE, non-FP32 workloads use two bags per program. When a histogram feature is split out, the remaining generic feature ranges use four. Other shapes use one TorchRec accepts config-driven int32 indices and offsets when the linearized range fits below 2^31. This halves index/offset storage and the CUB radix-sort key width while keeping int64 as the default
更多技术细节可访问官方原文:https://pytorch.org/blog/modernizing-table-batched-embeddings-with-fbtriton/。