[flashinfer] FlashInfer SM110 XQA 최적화: register_mma_split 도입으로 FP16 Paged Attention 성능 향상
PR 링크: flashinfer-ai/flashinfer#5597 상태: Merged | 변경: +2958 / -21
들어가며
FlashInfer의 SM110 XQA(eXperimental Query Attention) 모듈은 NVIDIA Thor 아키텍처에서 고성능 어텐션 연산을 제공하기 위해 설계되었습니다. 기존 register_mma 커널은 D512 트리 어텐션에서 강력한 성능을 보여주었으나, 특정 캐시 모드(FP16 page128)에서 하드웨어 자원을 더 효율적으로 활용할 여지가 있었습니다. 본 PR은 register_mma_split이라는 새로운 커널 경로를 추가하여, KV 시퀀스를 두 개의 워프 그룹으로 나누어 병렬 처리하고 공유 메모리에서 병합함으로써 성능을 극대화했습니다.
코드 분석
1. 벤치마크 및 설정 추가 (benchmarks/sm110_xqa_shapes.json)
새로운 커널 경로인 register_mma_split을 벤치마크 대상에 추가하여 성능을 측정할 수 있도록 했습니다.
{"name": "tree_fp16_paged_mma_split", "batch": 1, ..., "kernel": "register_mma_split"},
2. 백엔드 로직 개선 (flashinfer/experimental/sm110_xqa/backend.py)
사용자가 kernel="register_mma_auto"를 선택하면, 시스템이 자동으로 FP16 Paged 모드에서는 register_mma_split을, 그 외에는 기존 register_mma를 선택하도록 디스패치 로직을 구현했습니다.
if kernel == "register_mma_auto":
kernel = "register_mma_split" if fp16_paged else "register_mma"
3. 커널 설계 (flashinfer/experimental/sm110_xqa/README.md)
register_mma_split은 512개의 스레드를 가진 CTA 내에서 8개의 워프로 구성된 두 개의 그룹이 KV 시퀀스의 절반씩을 담당합니다. 각 그룹은 독립적인 K/V 링을 사용하여 스테이징하고, 최종적으로 공유 메모리에서 부분합(partial)과 행 통계를 병합합니다.
왜 이게 좋은가
이 최적화의 핵심은 작업 부하의 분산(Workload Splitting)과 공유 메모리 병합(In-CTA Merge)입니다.
- 성능 수치: Thor 노드에서 측정 결과, 기존
register_mma대비 FP16 page128 모드에서 1.06배~1.08배(Cold-L2 기준)의 성능 향상을 보였습니다. 특히 업스트림 TensorRT Edge XQA 소스 대비로는 1.5배 이상의 성능을 기록했습니다. - 교훈: GPU 커널 최적화에서 단순히 하나의 큰 스케줄을 사용하는 것보다, 특정 데이터 레이아웃(여기서는 page128)에 맞춰 워프 그룹을 분할하고 공유 메모리를 전략적으로 활용하는 것이 메모리 대역폭과 연산 유닛의 병목을 해소하는 데 효과적임을 보여줍니다.
리뷰어 피드백
리뷰 과정에서 compute-sanitizer를 통한 racecheck가 수행되었습니다. 비록 53개의 경고가 발생했으나, 이는 cp.async와 mbarrier를 사용하는 비동기 스테이징 로직에서 발생하는 일반적인 현상으로, 실제 데이터 레이스 오류가 아님을 확인했습니다. CI 파이프라인에서 25/25 테스트가 통과하며 안정성을 입증했습니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.compile.html
- https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#warp-level-matrix-instructions-mma
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer MiniMax-H3 Attention 최적화: K/V-split을 통한 성능 향상 분석
- [flashinfer] FlashInfer의 실험적 NVFP4 어텐션 도입: SM103 최적화
- [flashinfer] FlashInfer: Blackwell W8A8 AlphaMoE Expert 계산 커널 퓨전으로 성능 비약적 향상
- [flashinfer] FlashInfer의 SM100/SM103 최적화: CAKE 기반 블록 희소 어텐션(VSA) 도입
- [flashinfer] FlashInfer의 PrimTS를 활용한 고성능 Block-Sparse Attention 최적화
PR Analysis 의 다른글
- 이전글 [flashinfer] FlashInfer, Qwen3-30B 모델의 성능 향상을 위한 CUDA 커널 최적화: L2 캐시 힌트 도입
- 현재글 : [flashinfer] FlashInfer SM110 XQA 최적화: register_mma_split 도입으로 FP16 Paged Attention 성능 향상
- 다음글 없음
댓글