본문으로 건너뛰기

[flashinfer] FlashInfer, MTP/Speculative 디코딩 성능 1.23배 향상: Packed-Row 최적화 분석

PR 링크: flashinfer-ai/flashinfer#5490 상태: Merged | 변경: +14184 / -150

들어가며

최근 대규모 언어 모델(LLM)의 추론 성능은 모델의 크기뿐만 아니라 얼마나 효율적으로 KV 캐시를 처리하고 Attention 메커니즘을 계산하는지에 따라 크게 좌우됩니다. 특히, Multi-Turn Prediction (MTP) 또는 Speculative Decoding과 같이 여러 개의 토큰을 순차적으로 예측하는 시나리오에서는 효율적인 KV 캐시 접근 및 처리가 필수적입니다. FlashInfer는 이러한 LLM 추론 가속을 위한 고성능 라이브러리로, 이번 PR(#5474 후속 작업)에서는 실험적인 BF16 Paged GQA 디코딩(flashinfer.decode.prepare_balanced_batch_decode_with_kv_cache(..., backend="cake"))에서 MTP/Speculative 디코딩(q_len_per_req 3~8)의 성능을 획기적으로 개선하는 Packed-Row MTP 프로그램을 도입했습니다.

이전 구현에서는 각 요청의 각 draft row를 별도의 작업 항목으로 처리하여, q_len_per_req = 7인 요청의 경우 KV 청크를 7번 스트리밍해야 했습니다. 이는 trtllm-gen의 MTP 커널 대비 약 0.4배의 성능을 보였습니다. 본 PR은 이러한 비효율성을 해결하고, MTP 디코딩 성능을 최대 1.9배까지 향상시키는 Packed-Row MTP 프로그램을 도입하여 LLM 추론 속도를 크게 개선했습니다.

이번 글에서는 이 PR의 핵심 변경 사항을 분석하고, Packed-Row MTP 프로그램이 왜 성능 향상을 가져오는지, 그리고 이 최적화가 주는 일반적인 교훈은 무엇인지 살펴보겠습니다.

코드 분석

이번 PR의 핵심은 MTP/Speculative 디코딩을 위한 새로운 "Packed-Row MTP 프로그램"의 도입과 기존 "Single-Row 프로그램"의 스케줄링 개선입니다.

1. Packed-Row MTP 프로그램 도입 (q_len_per_req 3~8)

기존에는 각 draft token의 8개 쿼리 헤드를 개별적으로 처리했지만, 새로운 Packed-Row MTP 프로그램은 각 요청의 모든 draft token에 대한 8개 쿼리 헤드를 하나의 N = 8 * q_len_per_req 크기의 MMA 타일로 묶습니다. 예를 들어, q_len_per_req가 34이면 32행, 58이면 64행의 MMA 타일을 사용합니다. 이를 통해 각 KV 블록은 요청당 한 번만 로드하면 됩니다.

Before (Single-Row Program의 비효율성):

- #5474 treated every draft row of a request as its own work item, so a request with
- #`q_len_per_req = 7` streamed each KV chunk seven times; the MTP rows ran at roughly 0.4x of the
- #trtllm-gen MTP kernel.

After (Packed-Row MTP Program):

+ # This PR adds a second generated program, the **packed-row MTP program**: the
+ # eight query heads of every draft token of a request are packed into one `N = 8 * q_len_per_req`
+ # MMA tile (a 32-row instance for `q_len_per_req` 3..4 and a 64-row instance for 5..8), so each KV
+ # block is loaded once per `(request, kv head, KV block range)` work item.

스케줄러는 이전과 동일하게 커널 자체에서 관리되며, 요청 길이를 기반으로 최적의 청크 길이를 선택하여 런치 makespan을 최소화합니다. 또한, 두 개의 청크 타일은 제자리에서 병합되고, 더 긴 분할 타일은 1, 2, q_len_per_req 또는 2 * q_len_per_req개의 병합 티켓을 사용하여 병합됩니다. 이를 통해 각 병합 워프는 최대 8개의 행 청크를 처리할 수 있습니다.

호스트 측에서는 program_kind(q_len_per_req) 함수가 q_len_per_req 값에 따라 적절한 프로그램을 선택합니다. q_len_per_req가 1 또는 2인 경우에는 기존의 Single-Row 프로그램이 유지됩니다. cake_jit.MODULES 레코드에는 kind 필드 (row, mtp32, mtp64)가 추가되었으며, Shape에 독립적인 워크스페이스 크기는 더 큰 MTP 슬롯을 수용하도록 증가했습니다 (약 85MB).

2. Single-Row 프로그램 스케줄링 개선 (q_len_per_req 1~2)

Single-Row 프로그램(q_len_per_req 1 및 2)에서도 개선이 이루어졌습니다. 스케줄러 워프는 이제 런치 비용 모델을 사용하여 "Multi-Wave Near-Uniform Batches"를 결정합니다. 이는 전체 타일보다 더 많은 CTA를 사용하고, KV 길이가 가장 긴 길이의 1/8 이내에 있도록 하는 방식입니다.

Before (기존 스케줄링):

- # its scheduler warp now decides **multi-wave near-uniform batches** (more whole tiles than CTAs, KV lengths within 1/8 of the longest) with a launch-cost model instead of always keeping whole tiles.

After (개선된 스케줄링):

+ # Full chunks run in complete waves, the partial wave carries the remainder chunks on its otherwise idle CTAs, and a wave with `n` streaming CTAs costs `max(96, n)` CTA-slots per unit of work: below about 96 streaming CTAs the aggregate HBM rate scales with the active CTAs, above it the idle CTAs are free.
+ # Lanes 0..7 of the scheduler warp evaluate `ceil(total_work / (k * CTAs))` for `k = 1..8`, lanes 8..14 `ceil(p_max / n)` for `n = 2..8`, and the cheapest candidate replaces whole tiles when it is at least 2 % cheaper.

이 개선은 균일한 배치에서 성능 향상을 가져옵니다. 예를 들어, 64 요청 x 32768 토큰 x 8 KV 헤드 배치에서 4.0% ~ 3.7%의 성능 향상을 보였습니다. 반면, 마지막 웨이브에서 이미 HBM 대역폭을 포화시키는 배치에서는 오히려 성능이 저하될 수 있어, 이러한 경우에는 기존의 "whole tiles" 방식을 유지합니다.

3. 워크스페이스 및 메모리 관리

Packed-Row MTP 프로그램은 더 큰 MTP 슬롯을 사용하므로, balanced_gqa_decode_workspace_size 함수는 더 큰 워크스페이스 크기를 반환하도록 업데이트되었습니다. 예를 들어, 160 SM GPU에서는 약 85MB로 증가했습니다.

Before:

- balanced_gqa_decode_workspace_size(q.device), dtype=torch.uint8, device="cuda"
+ balanced_gqa_decode_workspace_size(q.device), dtype=torch.uint8, device="cuda"

After:

- the experimental path: #4832 (owner: @yyihuang; graduation plan: promote once the
- scheduler has been exercised from a serving stack on the ragged regime and the head-ratio / page-size
- coverage has been widened).
+ the workspace for any batch on that device (about 10.7 MB on a 160-SM GPU);
+ balanced_gqa_decode_workspace_size(q.device), dtype=torch.uint8, device="cuda"
+ # the workspace for any batch on that device (about 85.2 MB on a 160-SM GPU,
+ # sized for the packed-row MTP program's 64-row partial slots);

또한, _carve 함수는 워크스페이스 영역이 필요한 크기보다 클 경우에도 올바르게 작동하도록 수정되었습니다. 이는 Packed-Row MTP 프로그램이 사용하는 더 큰 슬롯을 고려한 것입니다.

왜 이게 좋은가?

1. MTP/Speculative 디코딩 성능의 비약적 향상

이 PR의 가장 큰 성과는 MTP/Speculative 디코딩 시나리오에서의 성능 향상입니다. Packed-Row MTP 프로그램은 KV 캐시 로딩을 요청당 한 번으로 줄여 데이터 이동을 최소화하고, MMA 연산의 효율성을 극대화합니다. 그 결과, 다양한 하드웨어 및 워크로드에서 trtllm-gen 대비 상당한 성능 향상을 보여줍니다.

B200 (SM 10.0) 성능:

  • AgentX q_len_per_req = 7 시나리오: 1.02x 향상 (이전 0.4x)
  • 균일 8-KV-head 배치 (q_len_per_req = 7): 1.33x 향상
  • 균일 8-KV-head 배치 (q_len_per_req = 3): 1.36x 향상
  • 단일 긴 요청: 1.40x 향상
  • Ragged 행: 1.29x 향상

GB300 (SM 10.3) 성능:

  • AgentX q_len_per_req = 7 시나리오: 1.12x 향상
  • 균일 8-KV-head 배치 (q_len_per_req = 7): 1.59x 향상
  • 균일 8-KV-head 배치 (q_len_per_req = 3): 1.68x 향상
  • 단일 긴 요청: 1.64x 향상
  • Ragged 행: 1.90x 향상

전체 18개 테스트 케이스에 대한 기하 평균 성능 향상은 B200에서 1.15x, GB300에서 1.23x입니다.

2. Single-Row 프로그램의 효율적인 스케줄링

q_len_per_req가 1 또는 2인 경우에도 스케줄링 로직이 개선되었습니다. 런치 비용 모델을 도입하여 CTA 사용률을 최적화하고, 특히 HBM 대역폭이 포화되지 않는 시나리오에서 성능을 향상시켰습니다. 이는 모든 시나리오에서 최적의 성능을 추구하는 FlashInfer의 철학을 보여줍니다.

3. 일반화된 워크스페이스 관리

Packed-Row MTP 프로그램은 더 많은 내부 상태를 저장해야 하므로 더 큰 워크스페이스를 요구합니다. workspace_bounds 및 _carve 함수의 업데이트는 이러한 요구 사항을 충족하며, 향후 더 복잡한 커널에서도 유연하게 대처할 수 있는 기반을 마련했습니다.

4. CUDA Graph 호환성 유지

이러한 성능 개선에도 불구하고, 호스트 측에서 seq_lens를 복사하지 않고 커널 자체에서 길이를 관리하는 기존의 방식을 유지했습니다. 이는 CUDA Graph 캡처 및 재생 기능을 그대로 지원하여, 서빙 환경에서의 안정성을 보장합니다.

일반적인 교훈

  1. Work Item Packing의 중요성: MTP/Speculative Decoding과 같이 반복적인 연산에서는 관련 데이터(여기서는 쿼리 헤드)를 가능한 한 많이 묶어서 처리하는 것이 메모리 대역폭 병목 현상을 줄이고 연산 효율성을 높이는 데 매우 효과적입니다. "Packed-Row" 전략은 이러한 원칙을 잘 보여줍니다.
  2. 동적 스케줄링 및 런치 비용 모델: GPU 커널의 성능은 단순히 연산량을 줄이는 것뿐만 아니라, CTA/Wave 할당, 메모리 접근 패턴 등 동적인 요소를 얼마나 잘 관리하느냐에 달려있습니다. 런치 비용 모델을 사용하여 최적의 스케줄링 결정을 내리는 것은 복잡한 GPU 워크로드에서 성능을 극대화하는 핵심입니다.
  3. 호스트-커널 분리: 가능한 한 많은 로직(예: 스케줄링 결정, 카운터 리셋)을 호스트가 아닌 커널 내에서 처리하면 CUDA Graph와 같은 기능을 활용하기 용이해집니다. 이는 서빙 환경에서 중요한 안정성과 유연성을 제공합니다.
  4. 점진적 성능 개선: 기존의 Single-Row 프로그램도 스케줄링 로직을 개선하여 성능을 향상시켰습니다. 이는 새로운 최적화 기법을 도입하는 동시에 기존 코드베이스의 효율성도 꾸준히 개선해 나가는 것이 중요하다는 것을 보여줍니다.

References

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글