AivexaNewsSearch
AI news for builders and product teamsChecked every hour

Modernizing Table Batched Embeddings with FBTriton

Collected Oct 6, 2026

PyTorch published a technical post describing FBTriton, a Triton implementation of the Table Batched Embedding (TBE) forward and backward operators used for embedding lookups across sharded GPUs in recommendation systems. TBE is an embedding table operation that combines lookup and pooling for many tables in a single GPU launch to cut launch overhead and improve memory efficiency.

The forward path has a generic gather and a specialized fast path selected for one feature with E up to 64, D between 64 and 128, L at least 64, FP16 weights, FP32 output, and no per-sample weights or variable batch embedding. That path uses a histogram over the first 256 indices and evaluates counts times table with tl.dot, routing other shapes to the generic kernel. FP16 and BF16 weights accumulate in FP32, while FP32 weights accumulate in FP64 to preserve accuracy at large D. TorchRec accepts config-driven int32 indices and offsets when the linearized range fits below 2^31. A bounds-check step runs before the Triton forward pass, reported at up to 1.24x faster on B200; when fused_bounds_check is enabled the kernel validates and repairs invalid tensors, with the option defaulting off.

The backward pass sums upstream gradients for every unique table and row pair and applies one optimizer update per row. Work is reorganized into runs by segment length, the number of samples touching a row, and routed to short_run below 256, grad_accum plus apply at 256 and above, or a fused Blackwell tier using a device-scope fence so the last sub-program applies the update in one launch. Runs at or above the threshold split into fixed 256-lookup chunks.

Measured on GB200 across 307 shard configurations and 283 distinct shapes with exact row-wise Adagrad and FP16 weights, median forward speedup versus CUDA TBE is 1.28x. The post attributes gains to memory throughput: at identical occupancy, a CUDA kernel reached 678 GB/s versus 3,948 GB/s for Triton, and Triton stays on simple streaming to SL 256 while CUDA escalates at SL 32. On one large B200 setup with forward and backward state reuse, forward moved from 22.844 ms to 33.252 ms, backward from 56.693 ms to 32.931 ms, and combined latency from 79.537 ms to 66.183 ms, a 16.8% reduction.

The post states Triton still trails on shapes whose work sits entirely in runs shorter than 4, listing 11 of 307 shards each under about a millisecond, and is at parity above SL 256.

Why it matters: the entire sparse forward and backward path is now written in Python and is smaller than the original CUDA templates alone, with identical kernel bodies on Blackwell, Hopper and AMD and hardware features added as flags, which the post says should let rank and infra engineers implement embedding algorithms and fusions faster.

Read at PyTorch

Based on reporting from the original publisher. Visit the source for full context and later updates.

Publisher excerpt

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.... The post Modernizing Table Batched Embeddings with FBTriton appeared first on PyTorch .