[sglang] SGLang aiter 백엔드의 Sliding Window Attention(SWA) 최적화 및 안정성 개선
PR 링크: sgl-project/sglang#38756 상태: Merged | 변경: +145 / -15
들어가며
SGLang의 aiter 백엔드는 AMD ROCm 환경에서 고성능 추론을 제공하기 위해 설계되었습니다. 하지만 기존 구현에서는 Speculative Decoding(특히 EAGLE이나 FROZEN_KV MTP 방식) 사용 시 Sliding Window Attention(SWA)을 위한 KV 풀 매핑이 올바르게 해석되지 않는 문제가 있었습니다. 또한, paged_attention_ragged 함수가 SWA를 지원하지 않음에도 불구하고, SWA 레이어가 해당 경로로 라우팅될 경우 오류를 발생시키지 않고 잘못된 결과를 반환하는 치명적인 안정성 이슈가 존재했습니다. 본 PR은 이러한 문제를 해결하기 위해 SWA KV 풀 해석 로직을 개선하고, 안전한 실행을 위한 가드 레일을 도입했습니다.
코드 분석
1. SWA KV 풀 해석 로직 개선 (aiter_backend.py)
기존에는 model_runner.token_to_kv_pool이 SWAKVPool인지 직접 확인하여 SWA 사용 여부를 결정했습니다. 하지만 draft worker는 자체적인 draft pool을 가지므로, 타겟 모델의 SWA 매핑을 올바르게 참조하지 못하는 문제가 있었습니다.
Before:
self.use_sliding_window_kv_pool = (
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
and model_runner.token_to_kv_pool.swa_layer_nums > 0
)
After:
self.swa_kv_pool = self._resolve_swa_kv_pool(model_runner)
self.use_sliding_window_kv_pool = (
self.swa_kv_pool is not None and self.swa_kv_pool.swa_layer_nums > 0
)
새로 추가된 _resolve_swa_kv_pool 메서드는 draft worker의 특성(EAGLE vs FROZEN_KV MTP)을 고려하여 올바른 KV 풀을 반환하도록 로직을 분리했습니다.
2. Paged-decode 가드 추가
paged_attention_ragged는 SWA 인자를 받지 않습니다. 이를 방지하기 위해 명시적인 검증 로직을 추가했습니다.
After (추가된 가드):
@staticmethod
def _reject_paged_decode_sliding_window(layer):
if layer.sliding_window_size is not None and layer.sliding_window_size > -1:
raise ValueError(
"aiter paged decode cannot honor sliding-window attention..."
)
이 코드는 forward_decode 경로에서 호출되어, SWA 레이어가 지원되지 않는 경로로 진입하는 것을 원천 차단합니다.
왜 이게 좋은가
- 정확도 확보: Speculative Decoding 환경에서 SWA 매핑이 올바르게 해석되지 않아 발생하던 추론 오류를 해결했습니다. 실제 테스트 결과, Gemma-4 FROZEN_KV MTP 모델에서 0.820의 정확도를 기록하며 비-speculative 모델과 동등한 수준의 정확도를 달성했습니다.
- 안정성 강화:
ValueError를 통해 잘못된 연산 경로를 명확히 차단함으로써, 디버깅이 어려운 '조용한 실패(silent failure)'를 방지했습니다. - 일반적 교훈: 복잡한 추론 엔진에서 'Draft worker'와 'Target worker'가 KV 풀을 공유하거나 분리하는 구조를 가질 때, 단순히
isinstance체크만으로는 부족할 수 있습니다. 각 워커의 상태를 명확히 해석하는 별도의 Resolver 패턴을 도입하는 것이 유지보수와 확장성 측면에서 유리합니다.
참고 자료
참고 자료
- https://github.com/sgl-project/sglang
- https://pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [sglang] ROCm 환경에서 BF16 All-Reduce의 수치 안정성 확보하기: QuickReduce의 FP16 Saturation 이슈 해결
- [sglang] SGLang의 AMD AITER AllReduce 최적화: 하드코딩된 제약 제거 및 성능 개선
- [sglang] LLM 서빙 최적화: Gumbel-max 트릭으로 CPU 병목 제거하기 (SGLang 사례)
- [sglang] SGLang: LFM2-MoE 모델을 위한 SM90 커널 퓨전 최적화 분석
- [sglang] FLUX.2 모델 성능 최적화: Token Concatenation과 NVFP4 양자화의 커널 융합
PR Analysis 의 다른글
- 이전글 [vllm] vLLM, ROCm 환경에서 FP8 GEMM 최적화로 성능 4-9% 향상
- 현재글 : [sglang] SGLang aiter 백엔드의 Sliding Window Attention(SWA) 최적화 및 안정성 개선
- 다음글 [vllm] vLLM의 MLA KV 캐시 최적화: 커널 통합을 통한 성능 극대화
댓글