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 优化在各深度学习框架中仍是空白。
技术要点:
- 精确的融合重构:稠密背景损失 + 稀疏目标修正
- 自定义 VJP 及专用 Pallas TPU 反向 kernel,logit 分片在片上计算,logits 及梯度从不落入 HBM
- 动态 VMEM 预算与分片感知的集合通信提升(collective hoisting)
- 基于计数的零开销训练指标
结果:单芯片 TPU v5e/v6e 小基准上消除 OOM,最高提速 91.9%;8 芯片 TPU 训练 87.6 万商品的多标签 SASRec(Yambda-50M)时,峰值 HBM 降低 65.7%。
「Infra」频道最新
- Strata v0.1.40 发布:支持 Strix Halo 与多 GPU 批处理改进 — lxfater · 2026-10-06
- QNX 营收创纪录达 8030 万美元,黑莓转型 AI 嵌入式软件平台 — alysha_lobo · 2026-10-06
- 跑了几十小时端到端基准,Strata 推理服务器在完整构建任务上翻车 — julianharris · 2026-10-06
- R9700 32GB 本地跑 i2v 提速难:用户指 Windows 下 ROCm 表现糟糕 — vladomkd · 2026-10-06
- 8-16GB 显存新思路:通用模型编排多个领域专家小模型 — CyberExplore · 2026-10-06
- 纯浏览器跑 37GB Qwen MoE:专家权重从磁盘流式加载,输出对齐 llama.cpp — naklitechie · 2026-10-06