Meta Speeds Up PyTorch Embeddings With FBTriton
Meta has introduced FBTriton to modernize Table Batched Embeddings in PyTorch, replacing legacy CUDA kernels to significantly accelerate recommendation system workloads on modern GPUs.

Meta's PyTorch team has developed FBTriton, a new Triton-based kernel design for Table Batched Embedding (TBE) forward and backward passes. TBE combines embedding lookup and pooling for multiple tables into a single GPU launch to minimize overhead. The generic forward launch uses a calculation of ceil(B / BAGS_PER_PROGRAM) programs, with the inner loop issuing four independent row loads, or eight in a tuned two-bag path. Large non-VBE, non-FP32 workloads utilize two bags per program. Additionally, TorchRec now supports int32 indices and offsets when the linearized range is under 2^31, halving storage requirements.
For specialized workloads where E is 64 or less, D is between 64 and 128, L is 64 or more, and FP16 weights are used, a dedicated path processes 16 bags per program using a histogram over the first 256 indices. On B200 GPUs, this achieves up to a 1.24x speedup on bounds checking. An optional core-module path that moves index transpose, sort, and run-length encoding into the forward pass reduced combined latency on a large B200 configuration from 79.537 ms to 66.183 ms, a 16.8% improvement, even though forward latency rose from 22.844 ms to 33.252 ms and backward fell from 56.693 ms to 32.931 ms.
The backward pass routes runs by segment length (SL). Short runs under 256 bypass atomic operations, while longer runs default to 256-lookup chunks. On Blackwell architectures, a device-scope fence allows a fused single-launch update. Blackwell also utilizes Tensor Memory Accelerator (TMA) bulk atomic reductions to merge partials. Testing 307 shard configurations across 283 distinct shapes on GB200 with exact row-wise Adagrad and FP16 weights showed a median forward speedup of 1.28x. Triton achieved 3,948 GB/s of memory throughput compared to 678 GB/s for legacy CUDA, representing a 5.8x improvement at identical occupancy, while moving 0.91x the DRAM bytes and issuing 29% more load requests. Only 11 of the 307 shards, representing shapes with runs shorter than 4, fell below parity.
For AI practitioners, FBTriton replaces complex CUDA Jinja templates with clean Python code, boosting developer velocity and simplifying the implementation of state-of-the-art embedding algorithms. The implementation maintains portability across Blackwell, Hopper, and AMD architectures without requiring separate code forks. It also paves the way for "mega sparse kernels" that fuse forward, backward, and optimizer steps to eliminate extra memory passes.
This is our own summary of reporting by PyTorch Blog



