CutBCE: TPU Kernel Eliminates OOM in Large-Vocabulary Recommendation Training, 91.9% Faster
_reachsumit · x · 2026-10-06
A new paper introduces CutBCE, an exact, hardware-accelerated binary cross-entropy (BCE) loss and gradient operator built in JAX/Pallas for industrial sequential recommender systems with massive item catalogs (10^5–10^7 items).
The problem: Full-vocabulary BCE training materializes a dense [B, N, V] logits tensor in HBM, incurring O(BNV) memory and fatal OOM errors. While chunked Softmax loss optimizations exist for LLMs, large-scale multi-label BCE remained unexplored.
Techniques:
- Exact fused reformulation: dense background loss + sparse target corrections
- Custom VJP with a dedicated Pallas TPU backward kernel computing logit tiles on-chip—logits and gradients never reside in HBM
- Dynamic VMEM budgeting and sharding-aware collective hoisting for distributed meshes
- Count-based zero-overhead training metrics
Results: Eliminates OOM with up to 91.9% speedup on single-chip TPU v5e/v6e; on 8-chip TPU training of multi-label SASRec with 876k items (Yambda-50M), peak HBM drops 65.7%.
More from Infra
- Strata v0.1.40 adds Strix Halo support, multi-GPU batching and decode improvements — lxfater · 2026-10-06
- BlackBerry's QNX hits record $80.3M revenue, up 27%, betting on AI-era safety platform — alysha_lobo · 2026-10-06
- Long-running benchmarks find Strata inference server failing full-build scenarios — julianharris · 2026-10-06
- AMD R9700 owner: ROCm on Windows cripples local i2v, Vulkan runs fine — vladomkd · 2026-10-06
- Domain specialists orchestrated by a general model: a local-LLM architecture pitch for 8-16GB GPUs — CyberExplore · 2026-10-06
- Run a 37GB Qwen MoE in a browser tab: LocalMind streams expert weights from disk, matches llama.cpp output — naklitechie · 2026-10-06