[sglang] SGLang SM100 CuteDSL Prefill 최적화: State I/O의 커널 퓨전
PR 링크: sgl-project/sglang#30169 상태: Merged | 변경: +299 / -37
들어가며
SGLang의 SM100 CuteDSL prefill 커널(GDN 및 KDA)은 기존에 각 레이어 호출마다 ssm_states[indices].contiguous()를 통한 gather와 index_copy_를 통한 scatter 작업을 수행하고 있었습니다. 이는 매 레이어마다 2번의 추가적인 커널 실행과 [N, HV, V, K] 크기의 중간 할당을 발생시켜, 연산 성능이 뛰어난 CuteDSL 커널의 이점을 상쇄하는 병목 현상을 유발했습니다. 본 PR은 이러한 state I/O를 커널 내부로 퓨전(Fusion)하여 성능을 최적화합니다.
코드 분석
1. 커널 로직 변경 (kernel_h.py)
기존에는 h0와 ht가 개별 시퀀스에 대해 gather된 상태를 가정했으나, 이제는 전체 state pool을 직접 참조합니다. state_indices를 통해 각 시퀀스가 자신의 state row를 직접 찾도록 변경되었습니다.
# Before
grid = (self.Hv, h0.shape[0], 1)
# After
grid = (self.Hv, cu_seqlens.shape[0] - 1, 1)
# ... 내부 로직 ...
state_slot = state_indices[seq_id]
# TMA load/store 시 state_slot을 인덱스로 사용
simple_tma_copy(H0_tma_atom, tmaH0[state_slot, head_id, None, None], sH0, h0_mbar)
2. Wrapper 인터페이스 개선 (__init__.py)
initial_state_indices를 선택적 인자로 추가하여, pool 모드일 경우 별도의 복사 없이 원본 pool을 직접 수정하도록 설계했습니다.
# Before
final_state = torch.empty_like(initial_state)
# After
if initial_state_indices is None:
final_state = torch.empty_like(initial_state)
state_indices = torch.arange(...)
else:
final_state = initial_state
state_indices = initial_state_indices
왜 이게 좋은가
이번 최적화는 불필요한 메모리 할당과 커널 런타임 오버헤드를 완전히 제거했습니다. B200 환경에서 벤치마크 결과, 동시성(N)이 높을수록 성능 향상이 두드러지며, 특히 GDN의 경우 최대 42.7%의 성능 향상을 보였습니다.
핵심 교훈:
- I/O Fusion: 커널 외부에서 수행되던 데이터 이동(Gather/Scatter)을 커널 내부의 TMA(Tensor Memory Accelerator) 연산으로 통합하면 메모리 대역폭과 커널 실행 오버헤드를 획기적으로 줄일 수 있습니다.
- Zero-copy:
initial_state를 직접 수정하는 'Pool mode'를 도입함으로써, 대규모 시퀀스 처리 시 발생하는 메모리 복사 비용을 0으로 만들었습니다.
참고 자료
- Cute (CUDA Template Engine) — SGLang에서 사용하는 고성능 커널 템플릿 엔진
참고 자료
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [sglang] [SGLang] VLM 추론 성능의 비약적 향상: Cross-request ViT Batching과 Metadata 재사용 기법
- [sglang] SGLang DFlash 최적화: 호스트-디바이스 동기화 제거를 통한 추론 성능 향상
- [sglang] DeepSeek NextN을 위한 Fused EH Norm 최적화: 커널 융합으로 성능 극대화하기
- [sglang] SGLang LTX-2.3 Diffusion 모델 최적화: Residual-Gate 연산 CUDA Fast Path 도입
- [sglang] SGLang 성능 최적화: D2H 복사 연산의 비동기 오버랩 구현
PR Analysis 의 다른글
- 이전글 [flashinfer] FlashInfer의 FP4 GEMM 최적화: 휴리스틱 개선과 Autotuning 효율화
- 현재글 : [sglang] SGLang SM100 CuteDSL Prefill 최적화: State I/O의 커널 퓨전
- 다음글 [vllm] vLLM MoE 성능 최적화: FlashInfer One-Sided Combine을 활용한 메모리 복사 제거
댓글