[Liger-Kernel] Liger-Kernel: Cross-Entropy와 Total Variation Distance를 하나로 융합하여 성능을 극대화하다
PR 링크: linkedin/Liger-Kernel#1384 상태: Merged | 변경: +584 / -0
들어가며
딥러닝 모델 학습 과정에서 지식 증류(Knowledge Distillation)는 매우 중요한 기법 중 하나입니다. 특히, 모델의 예측 분포와 정답 레이블 간의 Cross-Entropy(CE) 손실과 함께, 모델의 예측 분포와 교사 모델의 예측 분포 간의 유사성을 측정하는 Total Variation Distance(TVD)를 함께 사용하는 경우가 많습니다. 하지만 기존 방식에서는 이 두 가지 손실을 계산하기 위해 각각 별도의 연산을 수행해야 했고, 이 과정에서 상당한 메모리 및 연산 자원이 소모되었습니다.
LinkedIn의 Liger-Kernel 프로젝트에서 제안된 이번 PR은 이러한 비효율성을 해결하기 위해 Cross-Entropy와 Total Variation Distance를 하나의 융합된 커널(fused kernel)로 통합했습니다. 이 글에서는 해당 PR의 코드 변경 사항을 분석하고, 이 최적화가 왜 뛰어나며 어떤 이점을 가져다주는지 상세히 살펴보겠습니다.
코드 분석
이번 PR의 핵심은 src/liger_kernel/ops/fused_ce_tvd.py 파일에 새롭게 추가된 LigerFusedCETVDFunction입니다. 이 함수는 기존의 두 가지 손실 계산을 하나의 Triton 커널로 통합하여 효율성을 높입니다.
1. src/liger_kernel/ops/__init__.py 변경 사항
가장 먼저 눈에 띄는 변경은 __init__.py 파일에서 새로운 융합 커널을 import하는 부분입니다. 이는 새로운 기능이 Liger-Kernel의 연산자(ops) 모음에 정식으로 포함되었음을 나타냅니다.
--- a/src/liger_kernel/ops/__init__.py
+++ b/src/liger_kernel/ops/__init__.py
@@ -43,6 +43,9 @@
from liger_kernel.ops.fused_add_rms_norm import LigerFusedAddRMSNormFunction # noqa: F401
from liger_kernel.ops.fused_add_rms_norm import fused_add_rms_norm_backward # noqa: F401
from liger_kernel.ops.fused_add_rms_norm import fused_add_rms_norm_forward # noqa: F401
+from liger_kernel.ops.fused_ce_tvd import LigerFusedCETVDFunction # noqa: F401
+from liger_kernel.ops.fused_ce_tvd import fused_ce_tvd_backward # noqa: F401
+from liger_kernel.ops.fused_ce_tvd import fused_ce_tvd_forward # noqa: F401
from liger_kernel.ops.fused_linear_cross_entropy import LigerFusedLinearCrossEntropyFunction # noqa: F401
from liger_kernel.ops.fused_linear_cross_entropy import fused_linear_cross_entropy_backward # noqa: F401
from liger_kernel.ops.fused_linear_cross_entropy import fused_linear_cross_entropy_forward # noqa: F401
2. src/liger_kernel/ops/fused_ce_tvd.py 신규 파일
이 파일은 융합된 Cross-Entropy와 Total Variation Distance 연산을 구현하는 핵심 로직을 담고 있습니다.
핵심 아이디어:
기존 방식에서는 CE 손실 계산을 위해 student_logits에서 softmax를 취한 확률 분포 p를 계산하고, target 레이블에 해당하는 확률을 사용하여 손실을 구합니다. TVD 계산을 위해서는 student_logits와 teacher_logits 각각에서 softmax를 취한 확률 분포 p와 q를 계산하고, 이들의 차이의 절대값 합을 이용합니다. 이 과정에서 두 번의 softmax 연산과 두 개의 (BT, V) 크기의 확률 분포 임시 버퍼가 필요합니다.
새로운 LigerFusedCETVDFunction은 이 두 연산을 하나의 Triton 커널로 통합합니다. 이를 통해 다음과 같은 이점을 얻습니다:
- 메모리 절약:
(BT, V)크기의 두 개의 softmax 확률 분포 임시 버퍼 생성을 피합니다. 대신, 역전파(backward pass)에 필요한O(BT)크기의 상태(log-sum-exp 값, 시그마 값 등)만 유지합니다. - HBM 트래픽 감소: 불필요한 데이터 로딩 및 저장을 줄여 HBM(High Bandwidth Memory) 대역폭 사용량을 크게 감소시킵니다.
Forward Pass (_fused_ce_tvd_forward_kernel)
Forward 커널은 두 단계로 나뉩니다:
- Log-Sum-Exp 계산: 각 토큰(
BT)에 대해 어휘(V)를 스트리밍하면서student_logits와teacher_logits각각의 최대값(max)과 지수 합(sum-exp)을 계산하여 log-sum-exp (LSE) 값을 구합니다. 이 LSE 값은 softmax 확률 분포를 계산하는 데 사용되며, 역전파 시에도 필요합니다.# Pass 1: online max / sum-exp for both rows at once. student_max = -float("inf") student_sumexp = 0.0 teacher_max = -float("inf") teacher_sumexp = 0.0 for i in range(0, n_cols, BLOCK_SIZE): # ... (softmax 계산 로직) ... student_lse = student_max + tl.log(student_sumexp) teacher_lse = teacher_max + tl.log(teacher_sumexp) - TVD 및 시그마 계산: 다시 어휘를 스트리밍하면서, 이전 단계에서 계산된 LSE 값을 이용하여
student와teacher의 확률 분포p와q를 계산합니다. 이후, TVD의 절반(0.5 * sum |p - q|)과 역전파에 필요한sigma = sum_v p_v * sign(p_v - q_v)값을 계산합니다.# Pass 2: total variation and the Jacobian correction scalar. abs_diff_sum = 0.0 sigma = 0.0 for i in range(0, n_cols, BLOCK_SIZE): # ... (확률 분포 p, q 계산 및 TVD, sigma 누적 로직) ... tl.store(ce_ptr + pid, student_lse - student_at_target) # CE 계산 tl.store(tvd_ptr + pid, 0.5 * abs_diff_sum) # TVD 계산 tl.store(student_lse_ptr + pid, student_lse) tl.store(teacher_lse_ptr + pid, teacher_lse) tl.store(sigma_ptr + pid, sigma)
Backward Pass (_fused_ce_tvd_backward_kernel)
Backward 커널은 Forward pass에서 저장된 LSE 값과 시그마 값을 이용하여 student_logits에 대한 그래디언트를 계산합니다. CE와 TVD 각각의 그래디언트(grad_ce, grad_tvd)와 함께, softmax 함수의 야코비안(Jacobian)을 고려하여 최종 그래디언트를 계산합니다.
# With p = softmax(student) and s_v = sign(p_v - q_v)::
#
# d(ce)/d(student_v) = p_v - 1[v == target]
# d(tvd)/d(student_v) = 0.5 * p_v * (s_v - sum_u p_u * s_u)
#
# The second line is the softmax Jacobian applied to 0.5 * s; the summed
# term is the ``sigma`` scalar the forward pass already reduced.
# ... (그래디언트 계산 로직) ...
grad = grad_ce * p + grad_tvd * 0.5 * p * (sign - sigma)
grad = grad - tl.where(offsets == target, grad_ce, 0.0) # CE 그래디언트 조정
tl.store(grad_student_ptr + offsets, grad, mask=mask)
이 커널은 student_logits와 teacher_logits를 다시 로드하여 확률 분포 p와 q를 재계산하고, 이를 바탕으로 최종 그래디언트를 계산합니다. ignore_index 처리, tie-breaking 시 0 subgradient 사용 등 세심한 부분까지 고려되어 있습니다.
왜 이게 좋은가?
이 PR의 가장 큰 장점은 성능 향상입니다. PR 설명에 따르면, 특정 환경(V=151936, bf16)에서 기존 방식 대비 다음과 같은 개선이 이루어졌습니다:
- 메모리 사용량: 약 0.87 MB/token에서 ~0 MB/token으로 감소 (1024 토큰 기준 890 MiB 절약)
- HBM 트래픽: 약 5.9배 감소
이러한 성능 향상은 다음과 같은 이유로 가능합니다:
- 연산 융합 (Kernel Fusion): 두 개의 독립적인 연산을 하나의 커널로 통합함으로써, GPU 커널 실행 오버헤드를 줄이고 데이터 재사용성을 높였습니다.
- 메모리 접근 최적화:
(BT, V)크기의 중간 확률 분포 버퍼를 생성하지 않고,O(BT)크기의 상태만 유지함으로써 GPU의 HBM 사용량을 획기적으로 줄였습니다. 이는 특히 메모리 대역폭이 병목 현상을 일으키는 대규모 모델에서 큰 이점을 제공합니다. - Triton 활용: Triton 언어를 사용하여 GPU 하드웨어에 최적화된 커널을 효율적으로 작성했습니다. Triton은 자동 튜닝 및 최적화 기능을 제공하여 다양한 하드웨어 및 입력 크기에 대해 높은 성능을 달성할 수 있도록 돕습니다.
일반적인 교훈:
- 연산 융합의 중요성: 딥러닝 모델의 성능을 극대화하기 위해서는 개별 연산을 최적화하는 것뿐만 아니라, 여러 연산을 묶어 하나의 효율적인 커널로 만드는 것이 중요합니다.
- 메모리 접근 패턴 최적화: GPU 메모리(특히 HBM)는 CPU 메모리보다 훨씬 빠르지만, 여전히 병목 현상의 주요 원인이 될 수 있습니다. 불필요한 메모리 할당 및 접근을 최소화하는 것이 성능 향상의 핵심입니다.
- Triton과 같은 DSL 활용: 복잡한 GPU 커널을 직접 CUDA로 작성하는 것은 어렵습니다. Triton과 같은 고수준 DSL(Domain-Specific Language)은 개발 생산성을 높이면서도 GPU 최적화를 가능하게 합니다.
리뷰 피드백
이번 PR은 kashif님이 기여하고 vaibhavjindal님이 리뷰하며 긍정적으로 받아들여졌습니다. 특별히 기술적인 논쟁이나 복잡한 피드백은 없었으며, 이는 제안된 최적화가 명확하고 효과적이었음을 시사합니다.
References
- Triton Language Documentation
- PyTorch CrossEntropyLoss Documentation
- PyTorch KLDivLoss Documentation (TVD와 관련하여 확률 분포 간의 차이를 계산하는 데 사용될 수 있습니다.)
참고 자료
- https://triton-lang.org/main/getting-started/tutorial.html
- https://pytorch.org/docs/stable/generated/torch.nn.CrossEntropyLoss.html
- https://pytorch.org/docs/stable/generated/torch.nn.KLDivLoss.html
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
PR Analysis 의 다른글
- 이전글 [flashinfer] FlashInfer MiniMax-H3 Attention 최적화: K/V-split을 통한 성능 향상 분석
- 현재글 : [Liger-Kernel] Liger-Kernel: Cross-Entropy와 Total Variation Distance를 하나로 융합하여 성능을 극대화하다
- 다음글 [sglang] MiniMax-H3 모델의 추론 속도 4.6배 향상: Spectrum Skip-Step 최적화
댓글