[sglang] SGLang에서 SM120 GPU를 위한 SubBlock Sage FP8 어텐션 최적화
PR 링크: sgl-project/sglang#40116 상태: Merged | 변경: +277 / -44
들어가며
최근 대규모 언어 모델 및 멀티모달 모델의 추론 성능을 극대화하기 위해 하드웨어 가속기를 최대한 활용하는 것이 중요해졌습니다. SGLang은 MiniMax-H3 모델의 SubBlock sparse attention을 위해 SM90(H100) 아키텍처에서 Sage FP8 연산을 지원해 왔습니다. 이번 PR은 이 최적화 범위를 SM120(Blackwell) 아키텍처로 확장하여, FlashInfer의 CuTe-DSL 백엔드를 통해 추론 시간을 유의미하게 단축하는 것을 목표로 합니다.
코드 분석
1. 설정 및 런타임 검증 (minimax_h3.py)
기존에는 SM90만 지원하던 sage_fp8 모드를 SM120에서도 사용할 수 있도록 검증 로직을 수정했습니다.
# Before
if capability is None or capability.to_int() != 90:
# After
if capability is None or capability.to_int() not in (90, 120):
또한, 아키텍처별로 적절한 커널 로더를 동적으로 호출하도록 분기 처리하여 런타임 오버헤드를 최소화했습니다.
2. SM120 Sage 연산 구현 (subblock_sparse_attn.py)
SM120을 위한 새로운 Sage 어텐션 경로를 추가했습니다. 핵심은 quantize_sage_qkv_sm120을 통해 Q, K, V를 온라인으로 양자화하고, bsa_attn_sm120_blk64_sage_fwd 커널을 호출하는 것입니다.
# SM120 Sage FP8 Sparse Attention 구현
def _sm120_sage_fp8_sparse_attention(q, k, v, q2k_block_index, topk, softmax_scale, block_counts=None):
quantize, attention = _load_sm120_sage_ops()
q_hnd, k_hnd, v_hnd = (x.transpose(1, 2).contiguous() for x in (q, k, v))
quantized = quantize(q_hnd, k_hnd, v_hnd)
out = attention(*quantized, q2k_block_index.contiguous(), topk, ...)
return out.transpose(1, 2).contiguous()
왜 이게 좋은가
이번 최적화의 핵심은 온라인 FP8 양자화(Online FP8 Quantization)와 FlashInfer의 CuTe-DSL 기반 커널을 결합한 점입니다.
- 성능 향상: 8개의 RTX PRO 5000(SM120) GPU 환경에서 BF16 SubBlock 대비 3.86%에서 최대 14.86%의 추론 시간 단축을 달성했습니다.
- 범용성: SM90의 Q64xK128 블록 구조와 달리 SM120에서는 Q64xK64 블록을 사용하여 하드웨어 아키텍처에 최적화된 메모리 접근 패턴을 확보했습니다.
- 교훈: 하드웨어 가속기마다 최적화된 커널(FlashInfer 등)을 선택적으로 로드하고, 아키텍처별로 블록 크기를 유연하게 조정하는 설계가 대규모 모델 추론 엔진의 핵심임을 보여줍니다.
리뷰 과정에서 언급되었듯이, 최신 FlashInfer 의존성(PR #4691)이 필요하므로 README에 명시적인 설치 가이드를 추가하여 사용자 경험을 개선한 점도 주목할 만한 부분입니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.compile.html
- https://github.com/flashinfer-ai/flashinfer/pull/4691
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer: Blackwell W8A8 AlphaMoE Expert 계산 커널 퓨전으로 성능 비약적 향상
- [sglang] ROCm 환경에서 BF16 All-Reduce의 수치 안정성 확보하기: QuickReduce의 FP16 Saturation 이슈 해결
- [sglang] MiniMax-H3 모델을 24GB GPU에서 가속화하는 INT8 양자화 및 플러그형 어텐션 최적화
- [sglang] AMD MI355X 환경에서 Triton 3.7 레지스터 스필링 최적화
- [sglang] SGLang의 Wan2.2-TI2V 최적화: Triton 커널을 통한 메모리 트래픽 병목 해결
PR Analysis 의 다른글
- 이전글 [sglang] H200 GPU에서 GLM-5.2 MoE를 위한 W4A8 GEMM 커널 최적화 분석
- 현재글 : [sglang] SGLang에서 SM120 GPU를 위한 SubBlock Sage FP8 어텐션 최적화
- 다음글 [vllm] [ROCm] DeepSeek V4 성능 극대화: FP8 WO_A 출력 프로젝션 최적화 분석
댓글