본문으로 건너뛰기

[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)

교훈

  1. Memory vs. Compute Trade-off: 메모리 사용량을 최소화하는 것이 항상 최선의 성능을 보장하지 않습니다. 특히 GPU 연산에서는 커널 실행 오버헤드를 줄이기 위해 적절한 메모리 버짓을 할당하는 것이 중요합니다.
  2. Launch-bound 식별: 커널의 반복 횟수가 너무 많다면, 메모리 사용량을 조금 늘리더라도 청크를 병합하여 커널 호출 횟수를 줄이는 것이 효과적입니다.
  3. 수치적 안정성: 이 변경은 연산 순서의 미세한 차이(GEMM reduction-order)를 제외하면 결과값의 비트 단위 일치성을 유지하므로, 모델 학습의 수렴성에 영향을 주지 않는 안전한 최적화입니다.

참고 자료

⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.

댓글

관련 포스트

PR Analysis 의 다른글