본문으로 건너뛰기

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

기존에는 h0ht가 개별 시퀀스에 대해 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%의 성능 향상을 보였습니다.

핵심 교훈:

  1. I/O Fusion: 커널 외부에서 수행되던 데이터 이동(Gather/Scatter)을 커널 내부의 TMA(Tensor Memory Accelerator) 연산으로 통합하면 메모리 대역폭과 커널 실행 오버헤드를 획기적으로 줄일 수 있습니다.
  2. Zero-copy: initial_state를 직접 수정하는 'Pool mode'를 도입함으로써, 대규모 시퀀스 처리 시 발생하는 메모리 복사 비용을 0으로 만들었습니다.

참고 자료

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글