[flashinfer] FlashInfer의 PrimTS를 활용한 고성능 Block-Sparse Attention 최적화
PR 링크: flashinfer-ai/flashinfer#4474 상태: Merged | 변경: +17918 / -1218
들어가며
최신 LLM 추론 환경에서 Attention 연산은 병목의 핵심입니다. 특히 긴 문맥(Long-context) 처리와 효율적인 메모리 관리를 위해 Paged KV Cache와 Block-sparse Attention 기법이 필수적입니다. 이번 FlashInfer 업데이트에서는 NVIDIA Blackwell 아키텍처를 타겟으로 하는 PrimTS(Task-Scheduled Attention) 프레임워크를 확장하여, Q64 x KV256 실행 프로파일과 범용적인 Block-sparse Attention 기능을 도입했습니다. 본 글에서는 이 최적화가 왜 강력한 성능 향상을 가져오는지 분석합니다.
코드 분석
1. Q64/KV256 실행 프로파일 도입
기존의 Q64/KV128 프로파일을 넘어, 더 큰 KV 타일 사이즈인 KV256을 활용하여 메모리 대역폭 효율을 높였습니다. fmha_decode_config.py 등에서 볼 수 있듯이, 단순히 타일 사이즈를 키우는 것이 아니라 비용 모델(Cost Model)을 통해 Q 타일을 먼저 선택하고, 특정 조건에서만 KV256으로 승격(Promotion)시키는 전략을 취합니다.
# Before: 고정된 KV128 타일 사이즈 사용
# After: Q64 기반의 동적 KV256 프로파일 선택 로직 도입
max_q_tile_size = min(
q_block_size * heads_q_per_kv,
32 if kv_block_size < 64 else 128,
)
2. Block-Sparse API 및 Wrapper 구조화
block_sparse.py를 통해 BlockSparseTSWrapper와 BlockSparsePagedTSWrapper가 추가되었습니다. 이는 Contiguous 및 Paged K/V 스토리지 모두에서 동일한 워크플로우를 공유하도록 설계되었습니다.
# Block-sparse API 예시
block_sparse_attention_with_paged_kv_cache(
q, paged_kv_cache, page_table, ...
)
왜 이게 좋은가
이번 최적화의 핵심은 하드웨어 친화적인 타일링(Tiling)과 유연한 스케줄링입니다.
- 성능 수치: NVIDIA B200 환경에서 Dense Q64/KV256 프로파일은 기존 대비 최대 2.11배(GQA 64/8 기준)의 속도 향상을 보였습니다. 또한, Block-sparse 환경에서도 FlashInfer의 기존 FA2 백엔드 대비 2~4배 이상의 성능 개선을 달성했습니다.
- 일반적 교훈:
- 계층적 타일링: 고정된 타일 사이즈보다는 하드웨어의 연산 유닛(MMA)과 메모리 접근 패턴에 최적화된 타일 사이즈를 동적으로 선택하는 것이 중요합니다.
- Paged/Contiguous 통합: 스토리지 방식에 관계없이 동일한 커널 로직을 공유함으로써 유지보수성을 높이고, Paged KV Cache의 오버헤드를 12~17% 수준으로 최소화했습니다.
리뷰어 피드백 반영
리뷰 과정에서 kv_tile_size=256을 무조건 사용하는 것이 특정 케이스에서 성능 회귀(Regression)를 유발한다는 점이 발견되었습니다. 이를 해결하기 위해 Q 타일을 먼저 결정하고, 그 결과에 따라 KV256을 선택적으로 적용하는 로직으로 개선되었습니다. 이는 성능 최적화 시 비용 모델의 중요성을 다시 한번 보여줍니다.
결론
이번 업데이트는 Blackwell 아키텍처의 잠재력을 최대한 끌어내기 위한 정교한 엔지니어링의 결과물입니다. 특히 PrimTS를 통한 추상화와 세밀한 타일링 제어는 향후 더 복잡한 Attention 패턴을 지원하는 데 강력한 기반이 될 것입니다.
참고 자료
- https://github.com/flashinfer-ai/flashinfer
- https://docs.nvidia.com/cuda/parallel-thread-execution/index.html
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer, 초저병렬성 환경에서의 CP 델타 규칙 사전 계산 최적화
- [flashinfer] Blackwell 시대를 위한 최적화: FlashInfer의 SM120 Block-Sparse Attention 백엔드 도입기
- [flashinfer] Blackwell GPU를 위한 고성능 Recurrent-KDA 커널 최적화 및 통합
- [vllm] vLLM의 PLE 메타데이터 전송 최적화: 비동기 전송으로 성능 향상
- [flashinfer] FlashInfer의 NVFP4 KV 캐시 성능 최적화: FP4 연산의 병목 현상 해소
PR Analysis 의 다른글
- 이전글 [loki] Grafana Loki 성능 최적화: 파일 핸들 풀링을 통한 I/O 오버헤드 개선
- 현재글 : [flashinfer] FlashInfer의 PrimTS를 활용한 고성능 Block-Sparse Attention 최적화
- 다음글 [flashinfer] FlashInfer, Blackwell 아키텍처를 위한 Recurrent KDA Prefill 최적화: Small-BH 커널 도입
댓글