CutBCE:Google JAX TPU 核消除大规模推荐 BCE 训练 OOM,提速 91.9%

_reachsumit · x · 2026-10-06

论文提出 CutBCE,一个精确的、硬件加速的大词表二值交叉熵(BCE)损失与梯度算子,用 JAX + Pallas 实现,面向工业级序列推荐场景(商品目录达 10^5–10^7)。

问题:多标签推荐模型用全词表 BCE 训练时,标准实现要在 HBM 中物化稠密 [B, N, V] logits 张量,内存开销 O(BNV),极易 OOM。LLM 领域已有 Softmax 交叉熵的分块优化,但大规模多标签 BCE 优化在各深度学习框架中仍是空白。

技术要点:

结果:单芯片 TPU v5e/v6e 小基准上消除 OOM,最高提速 91.9%;8 芯片 TPU 训练 87.6 万商品的多标签 SASRec(Yambda-50M)时,峰值 HBM 降低 65.7%。

原文链接 →

「Infra」频道最新

更多「Infra」频道 AI 资讯 →