본문으로 건너뛰기

[flashinfer] Blackwell 아키텍처를 위한 고성능 Paged MQA Logits 커널 도입

PR 링크: flashinfer-ai/flashinfer#4365 상태: Merged | 변경: +11083 / -1

들어가며

최근 LLM 추론에서 DeepSeek의 MLA(Multi-Head Latent Attention) 구조가 주목받고 있습니다. 이 구조의 핵심은 sparse attention indexer가 KV 토큰을 선택하는 과정인데, 이를 위해 per-head weighted sum of rectified scores를 계산해야 합니다. 이번 PR은 NVIDIA의 최신 Blackwell(SM100/SM103) 아키텍처에서 이 연산을 가속화하기 위해 fp8_paged_mqa_logitsfp4_paged_mqa_logits 커널을 FlashInfer에 도입했습니다.

코드 분석

1. 커널 최적화 및 Divergence

기존 TensorRT-LLM의 구현을 포팅하면서, 성능 향상을 위해 다음과 같은 세 가지 의도적인 변경을 적용했습니다.

  • Zero-length row 스킵: Before (Upstream): 모든 행에 대해 TMA(Tensor Memory Accelerator) 및 UMMA 연산을 수행. After: task iterator에서 길이가 0인 행을 명시적으로 건너뛰어 불필요한 연산을 제거.

  • FP8 Epilogue 레지스터 최적화: Before: 고정된 슬롯 수 기반의 캐시 제한. After: 레지스터 footprint 기반의 동적 제한을 적용하여 num_heads가 작을 때 캐시 효율을 극대화.

  • 스케줄링 알고리즘 개선: Before: 선형 스캔 방식의 파티션 탐색. After: 이진 탐색(Binary Search)을 도입하여 컴파일 속도를 173배, 런타임 런치 비용을 10배 개선.

2. 스케줄링 및 CUDA Graph 지원

persistent kernel은 per-call 작업 할당이 필요합니다. 이를 GPU 상에서 직접 계산하도록 하여 호스트-디바이스 간 왕복을 제거했습니다.

# flashinfer/attn_scores/attn_scores.py
def compute_paged_mqa_logits_schedule(context_lens, ...):
    # GPU 상에서 스케줄을 계산하여 CUDA Graph 캡처 가능
    # ...

왜 이게 좋은가

이번 최적화는 단순히 기능을 추가하는 것을 넘어, Blackwell 아키텍처의 특성을 최대한 활용했습니다. 특히 스케줄링 로직을 선형에서 이진 탐색으로 변경한 것은 배치 사이즈가 2048인 환경에서 10배의 런치 비용 절감을 가져왔습니다. 또한, FLASHINFER_VALIDATE_INPUTS 플래그를 통해 성능 저하 없이 안전성을 확보할 수 있는 옵션을 제공합니다.

일반적 교훈:

  1. Persistent Kernel의 스케줄링: 정적 할당이 가능한 경우, 스케줄링 메타데이터를 재사용하되 ceil(context_lens / 256) 변경 시 반드시 갱신해야 하는 정합성 규칙을 명확히 해야 합니다.
  2. Lookahead와 Task Advancement의 동기화: 제로 길이 행을 스킵할 때, producer lookahead와 task iterator가 일치하지 않으면 파이프라인이 꼬일 수 있습니다. 모든 Warp 역할이 동일한 비제로 행을 선택하도록 강제하는 것이 중요합니다.

리뷰어 피드백 반영

리뷰 과정에서 leejnau님은 제로 길이 행 처리 시 발생할 수 있는 파이프라인 불일치 문제와, CUDA Graph 재사용 시 스케줄링 메타데이터가 stale해질 수 있는 위험을 지적했습니다. 이에 따라 context_lens가 256 토큰 경계를 넘을 때마다 스케줄을 재계산하도록 가이드를 수정하고, API 경계에서 입력 유효성 검사를 강화했습니다.

References

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글