Modernizing Table Batched Embeddings with FBTriton

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…

Annons
Annons
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. 1. What’s TBE 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 here 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. The general gather path Grid. The generic launch uses ceil(B / BAGS_PER_PROGRAM) programs. Each program loops over T features instead of launching a B×T grid. Gather width. The inner loop issues four independent row loads. The tuned two-bag path issues eight Bags per program. 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 Index and offset width. 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 Accumulation. FP16/BF16 weights accumulate in FP32. FP32 weights accumulate in FP64 to preserve accuracy at large D The small-table histogram and tensor-core path The specialized path is selected for one feature with E≤64, 64≤D≤128, L≥64, FP16 weights, FP32 output, no per-sample weights, and no VBE. One program handles 16 bags. It builds a histogram over the first 256 indices and evaluates counts × table with tl.dot; a scalar path handles the remaining indices. Other shapes use the generic kernel.  Bounds checking The standalone path uses an updated CUDA validation step before the Triton forward pass. On B200, this achieves up to 1.24x speedup on the bounds-check component across workloads. When fused_bounds_check is enabled, the eligible kernel validates and repairs invalid input tensors. Offset validation and repair remain a separate small kernel. Other configurations, such as weighted, variable batched, AMD, and transpose-hoist cases, fall back to standard validation. The option is exposed through TorchRec and defaults off. Forward-state reuse and preprocessing Exact row-wise Adagrad can save the forward histogram. Backward uses those counts for the compensated FP16 high/low GEMM into FP32, then applies the optimizer. Forward and backward remain separate launches; only the histogram counts are reused. Another optional core-module path moves index transpose, sort, and run-length encoding into forward and returns the metadata through autograd. It defaults off and is not exposed by the current TorchRec wrapper. On a large B200 configuration, forward moves 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 (−16.8%). 3. Implementation of Triton TBE Backward Kernel The core logic of the TBE backward propagation: for every unique (table, row) pair touched anywhere in the batch, sum the upstream gradient rows from every batch position that touched it, then apply exactly one optimizer update to that row. Annotations: T is number of tables; E is embedding size (i.e. number of rows); D is embedding dimension (i.e. number of cols); L is pooling factor (i.e. embedding bag size); B is batch size (i.e. number of embedding bags per request) Forward-pass. An embedding table T is a 2-D tensor (ExD) located in the GPU.  Given: a sparse feature (id-list or id-score-list) such as [id1, id2, id3, id4] (here L=4 and B=1) or multiple sparse features (B>1). Goal: By computing the sum T[id1%E]+T[id2%E]+T[id3%E]+T[id4%E], we get its forward output. If B=100, we will have 100 forward outputs. Backward-pass  Given: forward output and output grads (output.backward(grads)) Goal: calculate weights.grad and update weights Calc grads: for each index “id”, compute the sum of all the output grads from the embedding bags containing “id”. (With TxB=4.2M, B=128K, the indice stats can be total/dedup/highest_freq = 83M/3M/125K) Update weights: as simple as weight = weight – grad * LR TBE backward has no tensor core ops but it requires heavy data move/reduction and suffers from load imbalance issues.  transpose_embedding_input inverts the batch into runs: one unique row paired with the samples that touched it. This operation is hoisted into the forward pass, off the backward critical path. Segment length (SL), the number of samples touching a row, is the variable everything keys on, and it spans one to millions inside a single batch. Runs are routed by SL to one of three kernels. short_run (SL < 256): one program owns the run start to finish: gather, accumulate in registers, optimizer, store. Exclusive ownership means a plain store, no atomics. grad_accum + apply (SL >= 256, default): the run is split into 256-lookup chunks, partials land in a workspace, a second kernel applies the optimizer. fused (SL >= 256, Blackwell, very large batches): same split, but a device-scope fence lets the last sub-program apply the update in one launch. Weighted tables differ only in scaling each gradient row by its per-sample weight before accumulating. 4. Edge Cases and performance improvements Problem 1. imbalanced programs A single program walking two million lookups leaves the GPU idle, while a huge tail of one-lookup runs pays full per-run cost for almost no work. Solution 1.1: split long runs. Runs at or above the threshold become fixed 256-lookup chunks, split-K style. A two-million-lookup run turns into roughly eight thousand sub-programs, enough to fill the machine from one row’s work. Solution 1.2: CLC (TLX Blackwell) A kernel will persist once started initially and keep stealing ‘free’ run_id after finishing the current ‘run_id’. If the kernel steals a workload successfully, it will process it in the same block. Otherwise, the kernel exits. To address the load imbalance issues, we no longer need a software-based solution like bifurcating run_id processing by frequency. Instead CLC provides a hardware-based solution. Problem 2. the gather width is a register cliff Each buffered row keeps BLOCK_SIZE 64-bit addresses live, because dout_row_start_ptr[:, None] + col_offsets[None, :] materializes a whole pointer tile. Cost is width × BLOCK_SIZE, so a width tuned at one row width is wrong at another. Solution: per-target config width, in every tier. Measured on B200: tier width registers occupancy effect short run, unweighted 8 → 2 184 → 64 12.5% → 49.9% 0.41 → 1.05 long run, accumulate 8 → 2 158 → 62 17.6% → 44.7% fleet parity 82% → 87% short run, weighted 4 → 2 125 → 64 24.8% → 49.3% weighted parity 51% → 69% Problem 3. BLOCK_SIZE is a constexpr One launch must size BLOCK_SIZE to next_pow2(max_D) across all tables, so when the lookup-dominant table is much narrower than the widest, most of every gathered row is masked-off lanes. To overcome that, we bucket by dimension. Short runs are routed to one bucket per next_pow2(D) during classification, each launching with its own BLOCK_SIZE. It folds into the classification kernel, so it costs no extra pass; profiling confirms it buys reduced lane waste, not occupancy. Which runs are long is data-dependent and known only on the GPU, and reading those counts back with .item() is a cudaStreamSynchronize every backward. Solution: keep the shapes on the GPU. Workspace is preallocated to a bound computed from the index count, and classification runs in a kernel using atomic counters for stream compaction: 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) One atomic per block rather than one per element, intra-block offsets from a prefix sum. Kernels then consume the counts as device pointers and self-distribute with while-loops. Problem 5. split runs need a cross-program barrier Once a run is split, the optimizer update can only run after every sub-program’s partial has landed. Triton had no device-scope fence, so this cost a second kernel and a global-memory round trip. Solution: fence, then countdown (TLX Blackwell). TLX exposes the fence, letting the last sub-program apply the update in the same launch: 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: ... # last sub-program applies the optimizer and stores The ordering is the correctness argument: the fence makes each partial visible device-wide before the countdown decrements, so the program that sees remaining == 1 reads a complete sum. Without it one could win the countdown while another’s atomic_add was in flight, a silent numerical error. Problem 6. merging partials is a row-wide atomic Every sub-program atomically adds an entire BLOCK_SIZE row into the run’s workspace slot. A run split into eight thousand chunks means eight thousand programs contending on the same row, and tl.atomic_add issues that merge one element-wise operation at a time. Solution: reduce through TMA (TLX Blackwell). Blackwell has a better instruction than tl.atomic_add for merging partials cp.reduce.async.bulk.tensor, exposed as tlx.async_descriptor_store(..., store_reduce="add"). 5. Results and Analysis Results 307 shard configurations (283 distinct shapes) on GB200, exact row-wise Adagrad, FP16 weights. Metric: Triton performance divided by CUDA TBE. Median forward speedup is 1.28×. For backward data: Analysis: why does Triton beat CUDA here? We gain most of the win through changing the run-length. CUDA escalates to a cooperative CTA(Cooperative Thread Array)-per-row kernel at SL = 32; Triton stays on simple streaming to SL = 256. The speedup is from better memory throughputput. Deep in that band Triton runs 4.3x faster while moving the same DRAM bytes across the pass (0.91x) and issuing 29% more load requests. Nsight Compute shows where the difference lives: the CUDA kernel carrying this shape achieves 678 GB/s, versus 3,948 GB/s for the Triton kernel (5.8x) at identical occupancy. Just above SL = 32 there is nothing to amortize CTA-wide synchronization against, and Triton’s path has no cooperation at all. The wins came from config values. Every fix in Section 2 is exposed as a config change. Moving an escalation point or gather width in template-generated CUDA means restructuring which kernel handles what. Where Triton still loses. Shapes whose work is entirely in runs shorter than 4 sit below parity because of pure dependent-load latency; CUDA’s warp-per-row amortizes per-run metadata better. That is 11 of 307 shards, each under about a millisecond. Shapes above SL = 256 are at parity rather than ahead. We will invest more if they become a bottleneck on workloads. 6. Beyond Performance: What FBTriton Unlocks The most durable outcome of FBTriton is that the entire sparse path (both forward and backward) is now written in ordinary Python and is smaller than the original CUDA templates alone. As a result, we can apply more possible fusions in the future like demonstrated in forward and unlock these benefits:  Rapid Developer Velocity: FBTriton delivers high machine efficiency and developer efficiency simultaneously. Compared to CUDA Jinja templates, Triton TBE makes it much easier for rank/infra engineers to rapidly implement SOTA embedding algorithms for better model performance. Simultaneous Portability and Agility: FBTriton maintains identical kernel bodies whether running on Blackwell, Hopper, or AMD architectures. Instead of creating separate code forks for different hardware, specific features (like CLC, TMA bulk atomic reductions or device-scope fences) are integrated as additive flags. Pathway to a “Mega Sparse Kernel”: Rewriting CUDA kernels in Triton unlocks massive [forward, backward] and [prologue, epilogue] fusion opportunities. By treating the optimizer as a backward epilogue, FBTriton eliminates extra memory passes. This fusion collapses the entire sparse sequence into one kernel, reducing complex features like FP8 momentum scaling to just a few lines of accumulation. Code: link

Source: PyTorch — Published — Category: Open Source

🔗 Read full article on PyTorch →
Annons
Annons