[flashinfer] FlashInfer FP8 Causal Attention 최적화: O(1) 디코딩과 글로벌 스케줄링의 힘
PR 링크: flashinfer-ai/flashinfer#3575 상태: Merged | 변경: +75 / -4
들어가며\n\n대규모 언어 모델(LLM)의 추론 성능을 결정짓는 핵심 요소 중 하나는 Attention 연산의 효율성입니다. 특히 Causal Attention(인과적 어텐션)은 쿼리(Query)의 위치에 따라 참조해야 하는 키-값(KV) 캐시의 길이가 달라지는 특성을 가집니다. 이는 GPU 워크로드의 불균형(Workload Imbalance)을 초래하며, 특히 배치 크기가 커질수록 스케줄링 오버헤드가 성능의 발목을 잡게 됩니다.\n\n이번에 분석할 FlashInfer의 PR은 FP8 FMHA(Fused Multi-Head Attention) v2 커널에서 두 가지 핵심적인 최적화를 도입했습니다. 첫째는 모든 배치의 시퀀스 길이가 동일할 때(Uniform Sequence Length) 타일 정보를 찾는 과정을 $O(B)$에서 $O(1)$로 단축한 것이고, 둘째는 GPU 전체의 작업 부하를 고려한 '글로벌 역순 큐 타일 스케줄링(Global Reverse Q-tile Scheduling)'입니다. 이 변경을 통해 특정 조건에서 최대 5.66배의 성능 향상을 이끌어냈습니다.\n\n## 코드 분석: 무엇이 어떻게 바뀌었나?\n\n### 1. O(1) 타일 디코딩: 루프를 산술 연산으로 대체\n\n기존 코드에서는 각 타일(Tile)이 어느 배치(Batch)와 헤드(Head)에 속하는지 찾기 위해 cu_q_seqlens(누적 시퀀스 길이) 배열을 순회하는 $O(B)$ 루프를 사용했습니다. 배치 크기가 1024에 달하는 대규모 서빙 환경에서는 이 루프 자체가 무시할 수 없는 오버헤드가 됩니다.\n\nBefore:\ncpp\n// csrc/fmha_v2/fmha/warpspec/dma.h\n#pragma unroll 1\nfor (int batch_idx = 0; batch_idx < params.b; ++batch_idx) {\n int const actual_q_seqlen = params.cu_q_seqlens[batch_idx + 1] - params.cu_q_seqlens[batch_idx];\n // ... 타일이 이 배치에 속하는지 확인하는 로직\n}\n\n\nAfter:\ncpp\n// csrc/fmha_v2/fmha/warpspec/dma.h\nif (params.is_uniform_q) {\n int const q_tiles_per_head = compute_dynamic_q_tiles_per_head(params.cu_q_seqlens[1] - params.cu_q_seqlens[0]);\n // ...\n int const tiles_per_batch = q_tiles_per_head * params.h;\n bidb = static_cast<int>(tile_id) / tiles_per_batch; // O(1) 정수 나눗셈으로 배치 ID 계산\n int const within_batch = static_cast<int>(tile_id) - bidb * tiles_per_batch;\n bidh = within_batch / q_tiles_per_head; // 헤드 ID 계산\n // ...\n return true;\n}\n\n\n이 최적화는 is_uniform_q라는 플래그를 통해 활성화됩니다. 호스트 측에서 모든 배치의 길이가 같음을 미리 확인하면, 커널 내부에서는 복잡한 루프 없이 단순한 나눗셈과 나머지 연산만으로 자신의 위치를 찾아갈 수 있습니다.\n\n### 2. 글로벌 역순 스케줄링 (Global Reverse-Q Scheduling)\n\nCausal Attention에서 가장 무거운 작업은 시퀀스의 뒷부분에 있는 타일들입니다(참조할 KV가 가장 길기 때문). 기존에는 각 헤드 내부에서만 역순으로 처리했지만, 이번 PR은 이를 배치와 헤드 전체를 가로지르는 글로벌 단위로 확장했습니다.\n\nAfter (Global Reverse Logic):\ncpp\nif (reverse && params.use_head_first_scheduling) {\n // 모든 배치/헤드의 가장 무거운 타일을 먼저 발행\n q_step_offset = (q_tiles_per_head - 1 - static_cast<int>(tile_id) / (params.b * params.h)) * NUM_COMPUTE_GROUPS;\n int const tmp = static_cast<int>(tile_id) % (params.b * params.h);\n bidh = tmp / params.b;\n bidb = tmp % params.b;\n}\n\n\n이 방식은 'Straggler(가장 늦게 끝나는 작업자)' 문제를 해결합니다. 가장 무거운 타일들을 먼저 시작함으로써, 커널 종료 시점에 일부 스레드만 일을 하고 나머지는 노는 유휴 시간을 최소화합니다.\n\n### 3. L2 캐시를 고려한 지능형 게이팅 (L2 Gate Heuristic)\n\n글로벌 역순 스케줄링은 성능에 독이 될 수도 있습니다. 여러 헤드의 KV 스트림을 동시에 읽어오기 때문에, 만약 전체 KV 데이터가 L2 캐시 크기를 초과하면 캐시 스래싱(Thrashing)이 발생하여 성능이 급격히 떨어집니다. 이를 방지하기 위해 호스트에서 L2 캐시 적중 여부를 미리 계산합니다.\n\nAfter (Launcher Logic):\ncpp\n// csrc/fmha_v2/templates/kernel_hopper_ws.jinja\nsize_t kv_tokens = launch_params.total_kv_seqlen > 0 ? static_cast<size_t>(launch_params.total_kv_seqlen) : params.b * params.s;\nsize_t kv_size_in_bytes = kv_tokens * params.h_kv * params.d * 2 * {{ bytes_per_elt }};\n// KV 사이즈가 L2 캐시의 절반 이하일 때만 글로벌 스케줄링 활성화 (2x headroom heuristic)\nparams.use_head_first_scheduling = (2 * kv_size_in_bytes <= launch_params.device_l2_cache_size);\n\n\n## 왜 이게 좋은가?\n\n### 성능 수치 (H200, Qwen3-8B 기준)\n\n1. 대규모 배치 성능: Batch Size 1024, SeqLen 1024 환경에서 5.66x의 압도적인 속도 향상을 보였습니다. 이는 $O(B)$ 루프 제거가 대규모 배치에서 얼마나 결정적인지 증명합니다.\n2. 스케줄링 효율: 작은 시퀀스 길이(1024~4096)에서 글로벌 역순 스케줄링을 통해 약 1.08x ~ 1.28x의 성능 이득을 얻었습니다. 작업 부하가 불균형할 때 무거운 작업을 먼저 배치하는 전략이 유효함을 보여줍니다.\n\n### 일반적 교훈\n\n* Host-Side Intelligence: 커널 내부에서 복잡한 조건을 판단하기보다, 호스트에서 미리 계산하여 플래그(is_uniform_q)로 넘겨주는 것이 GPU 연산 자원을 아끼는 길입니다.\n* Memory Hierarchy Awareness: 스케줄링 전략을 짤 때는 단순히 연산 순서만 고려하는 것이 아니라, 그 순서가 메모리 계층(L2 Cache)에 미칠 영향을 반드시 계산해야 합니다. 이번 PR의 L2 게이팅 로직은 그 정석을 보여줍니다.\n\n## 리뷰어 피드백 반영\n\n리뷰 과정에서 jimmyzho는 is_uniform_q를 판단하는 더 가벼운 방법을 제안했습니다. akhilg-nv는 이를 받아들여 호스트에서 total_q_tokens == b * s_q인지 확인하는 방식으로 구현했습니다. 이는 별도의 배열 순회 없이도 배치의 균일성을 보장할 수 있는 영리한 체크 방식입니다.\n\n## 마치며\n\n이번 FlashInfer의 최적화는 단순히 코드를 빠르게 만드는 것을 넘어, 하드웨어의 특성(L2 캐시)과 알고리즘의 특성(Causal Attention의 불균형)을 깊게 이해했을 때 어떤 결과가 나오는지 잘 보여줍니다. 대규모 LLM 서빙을 준비하는 엔지니어라면 이 PR의 스케줄링 전략을 반드시 참고할 가치가 있습니다.
참고 자료
- https://github.com/flashinfer-ai/flashinfer
- https://docs.nvidia.com/cuda/cuda-c-programming-guide/index.html
- https://www.nvidia.com/en-us/data-center/hopper-architecture/
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer FP8 KV-Cache Prefill 성능 최적화: Repacking 기법을 통한 오버헤드 제거
- [flashinfer] [FlashInfer] CUTLASS MoE 커널 최적화: 벡터화와 동적 스레드 할당으로 성능 한계 돌파하기
- [flashinfer] FlashInfer의 GDN 커널 런칭 오버헤드 80% 절감하기: 호스트 측 최적화 전략
- [flashinfer] [FlashInfer] Kimi K3 모델을 위한 초고속 Fused KDA Decode 커널 분석 (SM100 최적화)
- [flashinfer] FlashInfer의 FP4 GEMM 최적화: 휴리스틱 개선과 Autotuning 효율화
PR Analysis 의 다른글
- 이전글 [onnxruntime] ONNX Runtime WebGPU EP: 디바이스 없는 오프라인 컴파일 지원
- 현재글 : [flashinfer] FlashInfer FP8 Causal Attention 최적화: O(1) 디코딩과 글로벌 스케줄링의 힘
- 다음글 [sglang] SGLang Diffusion: 2-rank Ulysses를 위한 CUDA-IPC 기반 Zero-Staging All-to-All 최적화
댓글