[Liger-Kernel] Liger-Kernel의 Fused Linear Cross Entropy 성능 최적화: C=16 전략
PR 링크: linkedin/Liger-Kernel#1414 상태: Merged | 변경: +19 / -1
들어가며
LLM 학습 시 Fused Linear Cross Entropy는 메모리 효율성을 위해 토큰을 청크(chunk) 단위로 나누어 처리합니다. 하지만 기존 구현에서는 메모리 사용량을 최소화하기 위해 C=1이라는 매우 보수적인 메모리 버짓 상수를 사용하고 있었습니다. 이로 인해 거대한 어휘 사전(Vocabulary)을 가진 모델(예: Llama-3)에서는 너무 많은 작은 청크가 생성되어, 커널 실행 오버헤드(Launch-bound)가 성능의 병목이 되는 문제가 발생했습니다. 본 글에서는 _CHUNK_MEM_CONST 상수를 16으로 상향 조정하여 성능을 최대 12.9배까지 개선한 최적화 전략을 살펴봅니다.
코드 분석
src/liger_kernel/ops/fused_linear_cross_entropy.py
핵심 변경 사항은 청크 크기를 결정하는 inc_factor 계산식의 수정입니다. 기존 코드는 메모리 점유를 최소화하는 데 집중했으나, 새로운 코드는 약간의 메모리 사용을 허용하는 대신 루프 반복 횟수를 대폭 줄였습니다.
Before:
inc_factor = triton.cdiv(V, H) # C=1: 메모리 점유 최소화
chunk_size = triton.next_power_of_2(triton.cdiv(BT, inc_factor))
After:
_CHUNK_MEM_CONST = 16
# ...
inc_factor = triton.cdiv(V, _CHUNK_MEM_CONST * H) # C=16으로 버짓 확대
chunk_size = triton.next_power_of_2(triton.cdiv(BT, inc_factor))
chunk_size = min(chunk_size, BT) # 단일 청크가 BT를 커버하도록 제한
기존의 inc_factor = triton.cdiv(V, H)는 C=1 수준의 메모리 버짓을 의미했습니다. 이를 C=16으로 확장함으로써, chunk_size가 커지고 결과적으로 num_chunks가 줄어들어 커널 호출 횟수가 감소하게 됩니다. 또한 min(chunk_size, BT)를 추가하여 청크가 전체 토큰 수(BT)를 초과하지 않도록 안전장치를 마련했습니다.
왜 이게 좋은가
이 최적화는 'Launch-bound' 문제를 해결하는 데 탁월합니다. Triton 커널은 GPU에서 실행될 때 커널 호출 자체에 오버헤드가 존재하는데, 청크가 너무 잘게 쪼개지면 연산 시간보다 커널을 실행하고 관리하는 시간이 더 길어집니다.
성능 개선 수치 (B200, bf16)
- BT=8192: 123.4ms → 24.7ms (5.0x)
- BT=4096: 119.3ms → 15.1ms (7.9x)
- BT=1024: 116.9ms → 9.07ms (12.9x)
교훈
- Memory vs. Compute Trade-off: 메모리 사용량을 최소화하는 것이 항상 최선의 성능을 보장하지 않습니다. 특히 GPU 연산에서는 커널 실행 오버헤드를 줄이기 위해 적절한 메모리 버짓을 할당하는 것이 중요합니다.
- Launch-bound 식별: 커널의 반복 횟수가 너무 많다면, 메모리 사용량을 조금 늘리더라도 청크를 병합하여 커널 호출 횟수를 줄이는 것이 효과적입니다.
- 수치적 안정성: 이 변경은 연산 순서의 미세한 차이(GEMM reduction-order)를 제외하면 결과값의 비트 단위 일치성을 유지하므로, 모델 학습의 수렴성에 영향을 주지 않는 안전한 최적화입니다.
참고 자료
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [vllm] vLLM Triton 커널 최적화: tl.constexpr 제거를 통한 JIT 컴파일 오버헤드 해결
- [sglang] SGLang: LFM2-MoE 모델을 위한 SM90 커널 퓨전 최적화 분석
- [vllm] vLLM의 작은 배치 사이즈를 위한 Triton 기반 Split-row Top-p 샘플링 최적화
- [sglang] ERNIE-Image의 RoPE와 GELU-mul 융합 및 RoPE cos/sin 호이스팅을 통한 성능 최적화
- [flashinfer] FlashInfer FP8 Causal Attention 최적화: O(1) 디코딩과 글로벌 스케줄링의 힘
PR Analysis 의 다른글
- 이전글 [vllm] vLLM, FlashInfer BF16 CuTeDSL GEMM 통합으로 저지연 추론 성능 향상
- 현재글 : [Liger-Kernel] Liger-Kernel의 Fused Linear Cross Entropy 성능 최적화: C=16 전략
- 다음글 [triton] Triton FP4 레이아웃 변환 최적화: 불필요한 메모리 할당 및 복사 제거
댓글