发表机构
Google Cloud(谷歌云)
机构由 AI 辅助整理,请以论文原文为准。AI 中文总结
针对序列推荐中大词表多标签BCE损失的内存爆炸问题,提出CutBCE,一种基于JAX/Pallas的精确硬件加速算子,通过融合重构、片上计算和动态内存管理,在TPU上消除OOM并大幅提升训练速度。
AI 中文摘要
工业级序列推荐系统需处理海量物品目录(如10^5至10^7个物品)。多标签推荐模型通常使用全词表上的二元交叉熵(BCE)损失进行训练,但标准BCE会在高带宽内存(HBM)中物化一个稠密的[B, N, V] logits张量,导致O(BNV)的不可承受内存开销及致命的内存溢出(OOM)错误。尽管大语言模型(LLM)中已有针对Softmax交叉熵的分块损失优化,但在深度学习生态中,大规模多标签BCE优化仍未被探索。我们提出CutBCE,一种在JAX和Pallas中实现的精确、硬件加速的BCE损失与梯度算子,专为大词表工作负载设计。CutBCE引入了:(1)一种精确融合的重构,同时评估稠密背景损失和稀疏目标修正;(2)一个自定义的向量-雅可比积(VJP)及专用的Pallas TPU反向内核,在两次前向/反向传播中均在芯片上计算logit分块,使得logits及其梯度从不驻留于HBM;(3)动态VMEM预算分配和面向分布式网格的分片感知集合提升;(4)基于计数的零开销训练指标。在单芯片TPU v5e/v6e微基准测试中,CutBCE消除了OOM错误,最高加速91.9%。在8芯片TPU切片上训练具有876k物品的多标签SASRec(Yambda-50M)时,CutBCE将峰值HBM降低65.7%(每芯片节省超过14 GiB),训练速度提升225.9%,且精度相当。CutBCE已在https URL开源。
英文摘要
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.