본문으로 건너뛰기

[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

왜 이게 좋은가

이번 최적화의 핵심은 연산 정밀도와 메모리 대역폭의 최적화입니다.

  1. 성능 향상: 32k 입력 길이에서 프리필 처리량이 기존 대비 +66% 향상되었습니다. 이는 Attention 연산이 병목인 긴 컨텍스트 처리에서 매우 유의미한 수치입니다.
  2. 메모리 절감: 인덱서 캐시를 fp8로 유지함으로써 GPU당 약 20GB의 메모리를 추가로 확보했습니다.
  3. 정확도 유지: GSM8K 벤치마크 결과, FP8 모드 사용 시 0.970으로 기존 bf16 방식(0.968)과 대등한 수준을 유지했습니다.

일반적 교훈:

  • 특정 커널(MSA 등)이 요구하는 데이터 정렬(Ascending order) 제약은 성능 최적화의 선결 조건입니다.
  • warp-wise 정렬(ballot+popc)은 smem을 거치는 k^2 정렬보다 훨씬 효율적입니다.
  • 하드웨어 가속기(SM100)의 네이티브 FP8 지원을 활용하려면, 데이터 로드 시점의 widening을 제거하고 연산 파이프라인 전체를 FP8로 유지하는 것이 중요합니다.

참고 자료

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글