[sglang] SGLang의 DeepSeek DSA 모델 최적화: Skip-TopK 레이어의 KV 캐시 효율화
PR 링크: sgl-project/sglang#30531 상태: Merged | 변경: +325 / -85
들어가며
최근 LLM 서빙 프레임워크인 SGLang에서 DeepSeek의 DSA(DeepSeek Sparse Attention) 모델을 위한 중요한 최적화가 이루어졌습니다. DSA 모델은 특정 레이어에서 이전 레이어의 Top-K 인덱스를 재사용하는 구조를 가지고 있는데, 기존 SGLang 구현체는 이러한 'Skip-TopK' 레이어에 대해서도 불필요하게 Indexer KV 캐시 슬롯을 할당하고 있었습니다. 본 PR은 이러한 낭비를 제거하여 메모리 점유율을 낮추고, 가용 KV 캐시 용량을 확보하는 것을 목표로 합니다.
코드 분석
1. python/sglang/srt/configs/model_config.py
dsa_layer_skips_topk 함수를 수정하여 cli_factor에 따른 레이어 스킵 여부를 명확히 판단하도록 했습니다.
# Before
pattern = getattr(config, "index_topk_pattern", None)
if pattern is not None:
return layer_id < len(pattern) and pattern[layer_id] == "S"
# After
cli_factor = getattr(config, "cli_factor", 1)
if cli_factor > 1:
return layer_id % cli_factor != 0
2. python/sglang/srt/mem_cache/index_key_cache.py
Skip-TopK 레이어는 인덱스 K를 기록하지 않으므로, 버퍼 할당 시 0-row를 반환하여 메모리 할당을 방지하고, 데이터 이동 로직에서 이를 건너뛰도록 최적화했습니다.
# After: 0-row placeholder를 통한 정렬 유지
def _layer_num_pages(self, layer_idx: int, num_pages: int) -> int:
return 0 if self.pool.skip_topk_layers[layer_idx] else num_pages
# After: 데이터 이동 시 빈 버퍼 스킵
for index_k in self.buffer:
if index_k.shape[0] == 0:
continue
index_k[tgt_loc_flat] = index_k[src_loc_flat]
왜 이게 좋은가
이번 최적화의 핵심은 '불필요한 리소스 할당의 제거'입니다. GLM-5.2 모델 기준으로 전체 78개 레이어 중 57개가 Skip-TopK 레이어인데, 이들에 대한 Indexer 모듈 생성을 방지함으로써 다음과 같은 성능 향상을 얻었습니다.
- 메모리 효율성:
max_total_num_tokens가 기존 2,574,848에서 3,002,752로 약 16.62% 증가했습니다. - 가용 자원 확보: Indexer 모듈을 생성하지 않음으로써 랭크당 약 0.98GB의 가중치 메모리를 절약하여 KV 캐시 풀로 환원했습니다.
- 처리량: GSM8K 벤치마크 기준, 정확도 손실 없이 처리량이 5,326 tok/s에서 5,432 tok/s로 소폭 상승했습니다.
이러한 최적화는 모델의 아키텍처적 특성(Sparse, Re-use)을 이해하고, 프레임워크의 메모리 관리 계층(Memory Pool)에서 이를 반영할 때 얼마나 큰 이득을 볼 수 있는지 보여주는 좋은 사례입니다.
리뷰어 피드백
초기 구현에서는 hiCache와의 호환성 문제로 인해 CUDA illegal memory access 오류가 보고되었습니다. 이후 PR 작성자는 이를 해결하기 위해 _should_elide_dsa_index_k와 같은 조건부 로직을 추가하고, 메모리 풀 구성 시 skip_topk_layers 정보를 명시적으로 전달하도록 개선하여 안정성을 확보했습니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.Tensor.data_ptr.html
- https://github.com/sgl-project/sglang
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [sglang] DeepSeek-V3.2를 위한 Native FP8 Sparse MLA 최적화: SGLang DSA 백엔드 통합 분석
- [sglang] DeepSeek NextN을 위한 Fused EH Norm 최적화: 커널 융합으로 성능 극대화하기
- [sglang] SGLang의 KV-Canary JIT 커널 도입: 효율적인 KV 캐시 검증 최적화
- [sglang] SGLang: LFM2-MoE 모델을 위한 SM90 커널 퓨전 최적화 분석
- [sglang] SGLang 성능 최적화: RTX 5090 32GB 환경에서의 CUDA Graph 및 Chunked Prefill 개선
PR Analysis 의 다른글
- 이전글 [sglang] [AMD gfx950] GLM-5.2 MLA 최적화: FP8 양자화와 Zero-Copy 레이아웃 전환
- 현재글 : [sglang] SGLang의 DeepSeek DSA 모델 최적화: Skip-TopK 레이어의 KV 캐시 효율화
- 다음글 [cpython] Python 문자열 split/splitlines 성능 개선: _PyList_AppendTakeRef 도입
댓글