본문으로 건너뛰기

[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 테스트가 통과하며 안정성을 입증했습니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글