[sglang] MiniMax-M3 모델을 위한 FP8 Attention GEMM 최적화 및 성능 개선
PR 링크: sgl-project/sglang#30971 상태: Merged | 변경: +1768 / -147
들어가며
최신 LLM 추론 환경에서 메모리 대역폭과 연산 효율은 성능의 핵심입니다. MiniMax-M3 모델은 기존에 bf16 기반의 Attention GEMM을 사용하고 있었으며, fp8_e4m3 KV 캐시를 사용할 때도 로드 시점에 bf16으로 확장(widening)하는 과정을 거쳐야 했습니다. 본 PR은 NVIDIA Blackwell(SM100) 아키텍처의 FP8 연산 능력을 십분 활용하여, Sparse/Dense Attention GEMM을 fp8_e4m3로 엔드투엔드 처리하도록 최적화했습니다. 이를 통해 프리필(prefill) 처리량을 크게 높이고 메모리 사용량을 대폭 절감했습니다.
코드 분석
1. minimax_decode_topk: 정렬된 블록 ID 출력
MSA(Multi-Scale Attention) 커널은 kv_block_indexes가 오름차순으로 정렬되어야 한다는 엄격한 제약이 있습니다. 기존의 비정렬(unordered) 출력 방식을 warp-wise 랭크 정렬 방식으로 변경했습니다.
// Before: 비정렬 출력
TopKTrait::forward(row, static_cast<uint32_t>(num_blocks), out, static_cast<uint32_t>(topk), &smem);
// After: 오름차순 정렬을 위한 warp-wise rank sort 적용
__shared__ int32_t s_topk[TopKTrait::kMaxTopK];
TopKTrait::forward(row, static_cast<uint32_t>(num_blocks), s_topk, static_cast<uint32_t>(topk), &smem);
__syncthreads();
// ... (ballot + popc를 이용한 랭크 계산 및 정렬 로직 추가)
2. trtllm_mha: 페이지 사이즈 128 지원
trtllm-gen의 동적 토큰-페이지 커널을 활용하기 위해 page_size == 128을 지원하도록 백엔드 제약 조건을 완화했습니다. 이는 특히 Sparse 블록 처리가 필요한 M3 모델의 Dense 백엔드에서 필수적입니다.
3. FP8 Attention GEMM 모드 도입
kv_cache_dtype == fp8_e4m3이고 SM100 아키텍처인 경우, 별도의 플래그 없이 자동으로 FP8 모드가 활성화됩니다. tl.dot 연산이 fp8x8로 수행되어 텐서 코어 효율을 극대화합니다.
# python/sglang/kernels/ops/attention/minimax_sparse/common/utils.py
# FP8 Q/K/V 연산을 위한 스케일 정규화 및 dtype 검증 로직
def unit_scale(scale: Optional[float]) -> float:
return 1.0 if scale is None else scale
왜 이게 좋은가
이번 최적화의 핵심은 연산 정밀도와 메모리 대역폭의 최적화입니다.
- 성능 향상: 32k 입력 길이에서 프리필 처리량이 기존 대비 +66% 향상되었습니다. 이는 Attention 연산이 병목인 긴 컨텍스트 처리에서 매우 유의미한 수치입니다.
- 메모리 절감: 인덱서 캐시를
fp8로 유지함으로써 GPU당 약 20GB의 메모리를 추가로 확보했습니다. - 정확도 유지: GSM8K 벤치마크 결과, FP8 모드 사용 시 0.970으로 기존 bf16 방식(0.968)과 대등한 수준을 유지했습니다.
일반적 교훈:
- 특정 커널(MSA 등)이 요구하는 데이터 정렬(Ascending order) 제약은 성능 최적화의 선결 조건입니다.
warp-wise정렬(ballot+popc)은smem을 거치는k^2정렬보다 훨씬 효율적입니다.- 하드웨어 가속기(SM100)의 네이티브 FP8 지원을 활용하려면, 데이터 로드 시점의 widening을 제거하고 연산 파이프라인 전체를 FP8로 유지하는 것이 중요합니다.
참고 자료
- Triton tl.dot — FP8 텐서 코어 연산의 핵심 함수
- SGLang Contribution Guide — 코드 스타일 및 테스트 가이드
참고 자료
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
PR Analysis 의 다른글
- 이전글 [sglang] DeepSeek-V4 Flash 공식 모델 최적화: H200 및 B200을 위한 SGLang 서빙 전략
- 현재글 : [sglang] MiniMax-M3 모델을 위한 FP8 Attention GEMM 최적화 및 성능 개선
- 다음글 [sglang] SGLang, KV VMM 할당자 스텁 최적화를 통한 시작 시간 및 라이브러리 크기 대폭 개선
댓글