본문으로 건너뛰기

[flashinfer] FlashInfer: Blackwell 아키텍처를 위한 Recurrent-KDA Prefill 최적화

PR 링크: flashinfer-ai/flashinfer#4675 상태: Merged | 변경: +29373 / -3168

들어가며

최신 LLM 추론 환경에서 긴 컨텍스트 처리는 성능의 핵심입니다. 특히 Blackwell(SM100/SM103) 아키텍처에서 recurrent_kda 연산은 고성능을 요구합니다. 이번 PR은 기존 backend="cake" 경로에 2단계 BT16(Block-Tensor 16) recurrent-KDA prefill 포트폴리오를 추가하여, 다양한 입력 형태와 디바이스 특성에 최적화된 스케줄링을 제공함으로써 성능을 비약적으로 개선했습니다.

코드 분석

1. 벤치마크 인프라 확장 (benchmarks/bench_recurrent_kda_prefill.py)

가장 큰 변화는 29개의 다양한 시나리오를 포함하는 PRODUCTION_CASES의 도입입니다. 이는 고정 레이아웃, 패킹된 시퀀스, 불규칙한 꼬리(tail) 길이 등 실제 프로덕션 환경을 모사합니다.

PRODUCTION_CASES = (
    Case("h96_fixed_8192", 96, (8192,), False, 10000),
    Case("h96_mixed_varlen", 96, (1300, 547, 2048, 963, 271, 3063), True, 10001),
    Case("h1_packed_524288_524288", 1, (524288, 524288), True, 11022),
)

또한, _resolve_recorded_cake_route 함수를 통해 1단계 또는 2단계(BT16 prepare + chain) 경로를 동적으로 결정하도록 로직이 개선되었습니다.

2. 동적 스케줄링 및 리소스 관리

_timing_iteration_budget 함수를 도입하여 GPU 메모리 제약 내에서 최적의 벤치마크 샘플링을 수행하도록 했습니다.

def _timing_iteration_budget(
    *, state_rotation_capacity: int, requested_dry_run_iters: int, requested_repeat_iters: int,
) -> tuple[int, int]:
    # ... (생략) ...
    available = state_rotation_capacity - _CUPTI_ESTIMATE_CALLS_PER_BLOCK
    # 가용 메모리에 맞춰 반복 횟수를 동적으로 조정
    return dry_run_iters, repeat_iters

왜 이게 좋은가

이번 최적화의 핵심은 Blackwell 아키텍처에 특화된 2단계 파이프라인(BT16 prepare + chain)입니다.

  • 성능 향상: B200 및 GB200 GPU에서 기존 FlashKDA 대비 기하 평균 2.672배의 속도 향상을 기록했습니다.
  • 유연성: backend="cake"를 통해 사용자가 명시적으로 최적화된 경로를 선택할 수 있게 했으며, 다양한 시퀀스 길이와 패킹 전략에 대응하는 범용성을 확보했습니다.
  • 교훈: 하드웨어 아키텍처(SM100/103)의 특성을 반영한 커스텀 스케줄링과, 실제 프로덕션 워크로드를 반영한 벤치마크 포트폴리오 구축이 고성능 라이브러리 개발에 필수적임을 보여줍니다.

리뷰어 피드백 반영

초기 CI 테스트에서 특정 M64 변형이 N=1, H=1 케이스에서 실패하는 문제가 있었으나, _resolve_recorded_cake_route를 통한 경로 선택 로직의 정교화와 테스트 케이스 보완을 통해 16/16 모든 테스트를 통과하도록 수정되었습니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글