본문으로 건너뛰기

[flashinfer] Blackwell 아키텍처에서 FlashInfer Ragged Prefill 성능 3.3배 향상시키기

PR 링크: flashinfer-ai/flashinfer#5133 상태: Merged | 변경: +1324 / -73

들어가며

최신 NVIDIA Blackwell(SM100) 아키텍처에서 FlashInfer의 backend="auto"는 기존에 fa2로만 고정되어 있었습니다. 이는 Blackwell의 강력한 cuDNN 및 CUTLASS 커널을 활용하지 못하게 하여, 특히 Ragged Prefill 작업에서 성능 병목을 유발했습니다. 본 PR은 BatchPrefillWithRaggedKVCacheWrapper.plan() 단계에서 지능적인 백엔드 선택 로직을 도입하여, auto 모드가 cuDNN과 CUTLASS를 우선적으로 시도하도록 개선했습니다.

코드 분석

1. flashinfer/prefill.py: 백엔드 선택 로직 개선

auto 모드에서 단순히 fa2를 선택하던 기존 방식 대신, Blackwell 환경에서 cudnncutlass를 순차적으로 검증하여 최적의 커널을 선택하도록 로직을 변경했습니다.

Before:

# 기존에는 fa2로 고정되거나 특정 조건에서만 분기
if backend == "auto":
    backend = determine_attention_backend(...)

After:

# Blackwell에서 cuDNN -> CUTLASS -> FA2 순으로 우선순위 적용
_BLACKWELL_RAGGED_AUTO_PREFERENCE = ("cudnn", "cutlass")

def _blackwell_ragged_auto_upgrade(self, ...):
    for backend in _BLACKWELL_RAGGED_AUTO_PREFERENCE:
        if self._is_eligible(backend): 
            return backend
    return "fa2"

2. flashinfer/cudnn/prefill.py: 토큰 단위 오프셋 정규화

cuDNN 백엔드가 토큰 단위 인덱스(token-unit indptrs)를 직접 처리할 수 있도록 하여, 별도의 변환 오버헤드 없이 고성능 커널을 호출하도록 했습니다.

Before:

# 요소 단위(element-unit) 오프셋만 지원
cudnn_q.set_ragged_offset_multiplier(h_qo * d_qk)

After:

# 토큰 단위 오프셋을 사용하여 유연한 스트라이드 처리
cudnn_q.set_ragged_offset_multiplier(q_token_stride)

왜 이게 좋은가

이번 최적화는 Blackwell GPU에서 최대 3.3배 이상의 성능 향상을 보여줍니다. 특히 d256/256과 같은 새로운 헤드 차원 구성에서도 cuDNN 백엔드가 성공적으로 작동하며, 기존 FA2 대비 압도적인 처리량을 제공합니다.

핵심 교훈:

  1. 백엔드 우선순위 전략: 단순히 하나의 커널을 고집하기보다, 하드웨어 가속기(cuDNN 등)의 지원 여부를 런타임에 검사하여 유연하게 전환하는 것이 중요합니다.
  2. 호출 규약 정규화: qo_indptr을 토큰 단위로 통일함으로써, 여러 백엔드 간의 인터페이스를 단순화하고 불필요한 메모리 복사나 변환 커널을 제거했습니다.

리뷰 피드백 반영

리뷰 과정에서 비연속적인 텐서(T3HD)에 대한 대응과 CUDA Graph 모드에서의 안정성 문제가 제기되었습니다. 이를 위해 plan() 단계에서 int32 인덱스 검증과 o_data_type 일치 여부를 사전에 체크하여, 런타임 오류를 방지하고 안정적인 백엔드 전환이 가능하도록 보완했습니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글