Industrial sequential recommender systems operate over massive item catalogs (e.g., 10^5--10^7 items). Multi-label recommendation models are trained with Binary Cross-Entropy (BCE) loss over the full vocabulary, but standard BCE materializes a dense [B, N, V] logits tensor in High Bandwidth Memory (HBM), incurring prohibitive O(BNV) memory and fatal Out-Of-Memory (OOM) errors. While chunked loss optimizations exist for Softmax Cross-Entropy in LLMs, large-scale multi-label BCE optimization remains unexplored across deep learning ecosystems. We propose CutBCE, an exact, hardware-accelerated BCE loss and gradient operator implemented in JAX and Pallas for large-vocabulary workloads. CutBCE introduces (1) an exact fused reformulation evaluating dense background loss and sparse target corrections; (2) a custom Vector-Jacobian Product (VJP) with a dedicated Pallas TPU backward kernel computing logit tiles on-chip in both passes so logits and their gradients never reside in HBM; (3) dynamic VMEM budgeting and sharding-aware collective hoisting for distributed meshes; and (4) count-based zero-overhead training metrics. On single-chip TPU v5e/v6e mini-benchmarks, CutBCE eliminates OOM errors with up to 91.9% speedup. On 8-chip TPU slice training for multi-label SASRec with 876k items (Yambda-50M), CutBCE reduces peak HBM by 65.7% (>14 GiB saved per chip) and increases training speed by 225.9% with comparable accuracy. CutBCE is open-sourced at https://github.com/AI-Hypercomputer/RecML/blob/main/recml/core/ops/binary_cross_entropy_ops.py.
Figures & tables
Memory Footprint
Standard BCE
CutBCE
Logits
[B,N,V]
[BN,BV] on-chip tile
Logit Gradient
[B,N,V]
[BN,BV] on-chip tile
Complexity
O(BNV)
O(BNBV)
TABLE I: Logit Memory Comparison.
Batch × SeqLen
Vocab ( V )
Keras (ms)
CutBCE (ms)
Speedup
Google Cloud TPU v6e (Trillium, 32 GiB HBM, Single Chip)
128×128
100,000
OOM
6.93
–
32×32
100,000
1.43
0.93
+53.8%
32×32
200,000
2.86
1.49
+91.9%
Google Cloud TPU v5e (16 GiB HBM, Single Chip)
128×128
100,000
OOM
12.79
–
TABLE II: Mini-Benchmark: Standard Keras BCE vs. CutBCE on a Single TPU v6e/v5e ( D=128,L=8 ). Step times averaged over 500 iterations.
Batch × SeqLen
Method
HBM (GiB)
Speed (steps/s)
Hit@8
64×256
SASRec-BCE
22.37
6.14
0.186
64×256
SASRec-CutBCE
7.67 ( -65.7% )
20.01 ( +225.9% )
0.183
128×256
SASRec-BCE
OOM
–
–
128×256
SASRec-CutBCE
7.89
15.82
0.166
64×512
SASRec-BCE
OOM
–
–
64×512
SASRec-CutBCE
7.89
15.61
0.192
TABLE III: End-to-End Multi-Label SASRec on Yambda-50M ( V=876,939 ) on an 8-chip TPU v6e slice ( 2×4 Topology, 32 GiB HBM/chip).