[sglang] SGLang: LFM2-MoE 모델을 위한 SM90 커널 퓨전 최적화 분석
PR 링크: sgl-project/sglang#37622 상태: Merged | 변경: +658 / -3
들어가며
최근 LLM 서빙 프레임워크인 SGLang에서 LFM2-MoE 모델의 성능을 극대화하기 위한 흥미로운 최적화가 진행되었습니다. 기존 LFM2 모델의 short-convolution 경로에서는 B * x를 계산하고, 이를 다시 전치(transpose)한 뒤 C * conv(B * x)를 수행하는 다단계 연산이 필요했습니다. 이러한 파편화된 커널 실행과 전치 연산은 특히 BF16 TP1 LFM2.5-8B-A1B 서빙 환경에서 프리필(prefill) 단계의 심각한 병목으로 작용했습니다. 본 글에서는 이 문제를 해결하기 위해 도입된 SM90 전용 Triton 커널 퓨전 전략을 살펴봅니다.
코드 분석
이번 최적화의 핵심은 별도로 실행되던 게이팅(gating) 연산과 컨볼루션 연산을 하나의 Triton 커널로 통합(fuse)하는 것입니다.
1. Triton 커널 구현 (lfm_short_conv.py)
새로 추가된 _lfm_short_conv_prefill_kernel은 게이트 연산과 컨볼루션을 단일 커널 내에서 처리합니다. 기존에는 중간 결과를 메모리에 쓰고 다시 읽어야 했으나, 이제는 레지스터 수준에서 연산이 이루어집니다.
# Before (Conceptual):
# 1. bx = B * x
# 2. transpose(bx)
# 3. conv(bx)
# 4. y = C * conv_result
# After (Fused in Triton):
conv = tl.zeros((BLOCK_T, BLOCK_D), dtype=tl.float32)
conv += bx0.to(tl.float32) * w0
conv += bx1.to(tl.float32) * w1
conv += bx2.to(tl.float32) * w2
# ... (중략) ...
active_y = (c * conv.to(tl.float32)).to(tl.bfloat16)
이 커널은 SM90 아키텍처(H100/H200 등)에서 최적화되어 있으며, bfloat16 연산을 유지하면서도 메모리 접근 횟수를 획기적으로 줄였습니다.
2. 디스패치 로직 (lfm2_moe.py)
모든 상황에서 이 커널을 사용하는 것은 위험하므로, 특정 조건(SM90, BF16, bias 없음 등)을 만족할 때만 동작하도록 can_dispatch_fused_lfm_short_conv 함수를 통해 가드(guard)를 설정했습니다.
def can_use_fused_lfm_short_conv(...):
return (
not _DISABLE_LFM_FUSED_CONV
and torch.cuda.get_device_capability(b.device) == (9, 0)
and b.dtype == torch.bfloat16
and bias is None
# ...
)
왜 이게 좋은가
이번 최적화는 단순히 코드 몇 줄을 합친 것이 아니라, GPU 메모리 계층 구조를 효율적으로 활용한 결과입니다.
- 성능 향상: 프리필 컨볼루션 세그먼트에서 기존 ~310us 대비 ~105us로 약 2.95배의 성능 향상을 달성했습니다. 전체 서빙 처리량 또한 3% 이상 개선되었습니다.
- 메모리 대역폭 절감: 중간 연산 결과인
B * x를 전역 메모리에 쓰지 않고 레지스터/Shared Memory에서 바로 컨볼루션으로 넘김으로써 메모리 대역폭 병목을 제거했습니다. - 정확성 보장: 리뷰 과정에서
test_lfm_short_conv.py를 통해 기존 구현과 비트 단위로 동일한 결과(bit-exact)를 내는지 검증되었으며, CUDA Graph와의 호환성도 확보했습니다.
교훈
- Kernel Fusion: 연산량이 적고 메모리 접근이 잦은 연산은 Triton을 통해 퓨전하는 것이 현대 GPU 아키텍처에서 가장 효과적인 최적화 방법입니다.
- Dispatch Guards: 최적화된 커널은 특정 하드웨어(SM90)나 데이터 타입에 종속적일 수 있으므로, 안전한 디스패치 로직을 함께 설계하는 것이 필수적입니다.
참고 자료
참고 자료
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [sglang] SGLang DFlash 최적화: 호스트-디바이스 동기화 제거를 통한 추론 성능 향상
- [sglang] SGLang의 MLA KV 캐시 쓰기 최적화: TMA Bulk-Store 도입
- [sglang] SGLang Triton 커널 최적화: libdevice.tanh 도입과 2D Strided Tensor 지원
- [sglang] SGLang의 디코드 성능 향상을 위한 Temperature 및 Softmax 커널 융합
- [sglang] FLUX.2 모델 성능 최적화: Token Concatenation과 NVFP4 양자화의 커널 융합
PR Analysis 의 다른글
- 이전글 [flashinfer] Blackwell 시대를 위한 최적화: FlashInfer의 SM120 Block-Sparse Attention 백엔드 도입기
- 현재글 : [sglang] SGLang: LFM2-MoE 모델을 위한 SM90 커널 퓨전 최적화 분석
- 다음글 없음
댓글