본문으로 건너뛰기

[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에 명시적인 설치 가이드를 추가하여 사용자 경험을 개선한 점도 주목할 만한 부분입니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글