Cut Binary Cross Entropy: 대규모 어휘를 위한 효율적인 손실 및 기울기 커널을 갖춘 순차 추천
요약
본 논문은 방대한 아이템 카탈로그를 사용하는 순차 추천 시스템의 메모리 문제를 해결하기 위해 CutBCE라는 새로운 BCE 손실 및 기울기 연산자를 제안합니다. 기존 표준 BCE는 대규모 로짓 텐서로 인해 OOM 오류를 유발하며, CutBCE는 온칩 계산과 효율적인 분산 처리를 통해 이를 극복하고 학습 속도를 크게 향상시킵니다.
핵심 포인트
- CutBCE는 대규모 어휘 워크로드의 메모리 문제를 해결합니다.
- 온칩(on-chip) 로짓 타일 계산을 위한 전용 Pallas TPU 커널을 사용합니다.
- HBM 피크를 65.7% 감소시키고 학습 속도를 225.9% 증가시킵니다.
- JAX와 Pallas에 구현되어 오픈 소스로 공개되었습니다.
산업용 순차 추천 시스템은 방대한 아이템 카탈로그(예: $10^5$--$10^7$개 아이템) 위에서 작동합니다. 다중 레이블 추천 모델은 전체 어휘에 대해 Binary Cross-Entropy (BCE) 손실로 학습되지만, 표준 BCE는 High Bandwidth Memory (HBM)에 밀집된 $[B, N, V]$ 로짓 텐서를 만들어내며, 이는 엄청난 $O(BNV)$ 메모리 사용량과 치명적인 Out-Of-Memory (OOM) 오류를 초래합니다. LLM에서 Softmax Cross-Entropy를 위한 청크 기반 손실 최적화는 존재하지만, 대규모 다중 레이블 BCE 최적화는 딥러닝 생태계 전반에 걸쳐 여전히 탐구되지 않은 영역입니다. 우리는 대규모 어휘 워크로드를 위해 JAX와 Pallas에 구현된 정확하고 하드웨어 가속화된 BCE 손실 및 기울기 연산자인 CutBCE를 제안합니다. CutBCE는 (1) 밀집 배경 손실 평가와 희소 타겟 보정을 계산하는 정확한 융합 재정의; (2) 로짓과 그 기울기가 HBM에 절대 머무르지 않도록 양방향 패스 모두에서 온칩(on-chip)으로 로짓 타일을 계산하는 전용 Pallas TPU 역전파 커널을 갖춘 사용자 정의 Vector-Jacobian Product (VJP); (3) 분산 메시를 위한 동적 VMEM 예산 책정 및 샤딩 인식 집합체 호이스팅; 그리고 (4) 카운트 기반 제로 오버헤드 학습 메트릭을 도입합니다. 단일 칩 TPU v5e/v6e 미니 벤치마크에서 CutBCE는 최대 91.9%의 속도 향상으로 OOM 오류를 제거합니다. 8개 칩 TPU 슬라이스에서 다중 레이블 SASRec을 876k 아이템(Yambda-50M)에 대해 학습했을 때, CutBCE는 피크 HBM을 65.7% 감소시키고 (>14 GiB/칩 절약) 유사한 정확도로 학습 속도를 225.9% 증가시킵니다. CutBCE는 https://github.com/AI-Hypercomputer/RecML/blob/main/recml/core/ops/binary_cross_entropy_ops.py에서 오픈 소스되었습니다.
AI 자동 생성 콘텐츠
본 콘텐츠는 arXiv cs.AR의 원문을 AI가 자동으로 요약·번역·분석한 것입니다. 원 저작권은 원저작자에게 있으며, 정확한 내용은 반드시 원문을 확인해 주세요.
원문 바로가기