[Liger-Kernel] Ascend NPU 성능 극대화: Liger-Kernel의 커널 최적화 분석
PR 링크: linkedin/Liger-Kernel#1426 상태: Merged | 변경: +2166 / -1169
들어가며
최근 대규모 언어 모델(LLM) 학습에서 NPU(Ascend)의 활용도가 높아지고 있습니다. 하지만 범용적인 Triton 커널은 특정 하드웨어 아키텍처에 최적화되지 않아, torch_npu와 같은 전용 라이브러리 대비 성능 저하가 발생하곤 합니다. 본 PR은 linkedin/Liger-Kernel에 Ascend NPU 전용 백엔드를 추가하고, CrossEntropy, RMSNorm, RoPE 등 핵심 연산을 최적화하여 하드웨어 가속을 극대화하는 것을 목표로 합니다.
코드 분석
1. CrossEntropy 최적화
기존의 범용 Triton 커널은 행 단위 연산 시 오버헤드가 컸습니다. 이번 변경에서는 HAS_TAIL을 constexpr로 처리하여 불필요한 마스킹 연산을 제거했습니다.
Before (범용 로직):
# 모든 로드에 mask가 포함되어 분기문이 메인 루프에 잔류함
mask=offs < n_cols
After (Ascend 최적화):
# HAS_TAIL이 constexpr이므로, V % BLOCK_SIZE == 0일 때 마스크 연산이 DCE(Dead Code Elimination)됨
if y == ignore_index:
tl.store(loss_ptr + row_i64, 0.0)
else:
# 루프 내에서 마스크 제거로 MTE2 대역폭 극대화
for row in tl.range(row_start, row_end):
# ... 연산 수행
2. RMSNorm 최적화
RMSNorm은 두 번의 메모리 로드를 피하기 위해 단일 패스(Single-pass) 융합 커널로 재작성되었습니다. 특히 UB(Unified Buffer) 크기에 맞춰 2D 융합 경로를 선택하도록 개선되었습니다.
3. RoPE 최적화
기존의 BSND 레이아웃 전치(Transpose) 비용이 전체 RoPE 시간의 40%를 차지하던 문제를 해결했습니다. BNSD 레이아웃을 네이티브로 지원하고, rotate_half 연산을 UB 내에서 수행하도록 변경했습니다.
왜 이게 좋은가
이번 최적화의 핵심은 HBM(High Bandwidth Memory) Round-trip 최소화와 하드웨어 친화적인 데이터 접근입니다.
- constexpr-DCE:
HAS_TAIL을 컴파일 타임 상수로 처리하여, 정렬된 데이터(Vocab size가 BLOCK_SIZE의 배수인 경우)에 대해 불필요한 조건 분기를 제거했습니다. 이는 대역폭 제한적인(Bandwidth-bound) 연산에서 성능 향상을 이끌어냅니다. - Software Pipelining:
tl.range를 활용하여 MTE2(메모리 로드), VEC(연산), MTE3(메모리 저장)를 파이프라이닝하여 하드웨어 유닛의 유휴 시간을 줄였습니다. - 성능 수치: 910B 칩셋 기준, 기존 2-pass 방식의 RMSNorm이 ~350 µs였던 반면, 최적화된 커널은 ~107 µs로 약 3배 이상의 성능 향상을 보였습니다.
교훈: NPU와 같은 가속기에서는 범용적인 코드보다 하드웨어의 메모리 레이아웃과 버퍼 크기를 고려한 커널 설계가 필수적입니다. 특히 constexpr을 활용한 분기 제거는 Triton 커널 최적화의 핵심 기법입니다.
참고 자료
참고 자료
- https://triton-lang.org/main/index.html
- https://www.hiascend.com/document/detail/en/canncommercial/80rc1/operatordev/tbe/tbe_10_0001.html
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [Liger-Kernel] Liger-Kernel의 Fused Linear Cross Entropy 성능 최적화: C=16 전략
- [sglang] [NPU] GLM-4.7-Flash 성능 최적화: Fused Triton 커널로 연산 병목 해결하기
- [sglang] Qwen3.5 및 Qwen3_Next 모델의 NPU 성능 향상을 위한 Triton 커널 퓨전 최적화
- [sglang] SGLang에서 NPU를 위한 LTX-2/2.3 추론 성능 최적화 및 호환성 개선
- [sglang] SGLang: LFM2-MoE 모델을 위한 SM90 커널 퓨전 최적화 분석
PR Analysis 의 다른글
- 이전글 [sglang] NVIDIA SM90 GPU를 위한 SGLang SubBlock Sparse Attention 최적화: Sage FP8 Compute 도입
- 현재글 : [Liger-Kernel] Ascend NPU 성능 극대화: Liger-Kernel의 커널 최적화 분석
- 다음글 [sglang] SGLang에서 NPU를 위한 LTX-2/2.3 추론 성능 최적화 및 호환성 개선
댓글