본문으로 건너뛰기

[flashinfer] [FlashInfer] Kimi K3 모델을 위한 초고속 Fused KDA Decode 커널 분석 (SM100 최적화)

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

들어가며

LLM(Large Language Model)의 추론 속도를 결정짓는 핵심 요소 중 하나는 커널 오버헤드메모리 대역폭(Memory Bandwidth)입니다. 특히 Kimi K3와 같은 최신 모델 아키텍처는 단순한 Attention을 넘어 Convolution, Recurrent Update, Gating, 그리고 Normalization이 복잡하게 얽힌 구조를 가집니다.

기존 방식에서는 이러한 연산들을 각각 별도의 커널로 실행했으나, 이는 각 단계마다 중간 결과값을 HBM(High Bandwidth Memory)에 썼다 읽어야 하는 비효율을 초래합니다. 이번 FlashInfer의 PR은 NVIDIA의 최신 아키텍처인 SM100(Blackwell)을 타겟으로, Kimi K3의 디코딩 과정을 하나의 거대한 커널로 통합(Fusion)하여 성능을 극대화한 사례입니다.

코드 분석: API 계층의 변화

먼저 사용자가 호출하는 Python API 레벨에서의 변화를 살펴보겠습니다. 기존에는 개별적인 recurrent_kda 연산만 존재했으나, 이제는 모든 과정을 한 번에 처리하는 fused_kda_decode가 추가되었습니다.

Before: 개별 연산 중심의 구조

기존 flashinfer/kda_decode.py에는 통합된 디코딩 로직이 없었으며, 사용자는 각 단계를 수동으로 관리해야 했습니다.

# flashinfer/kda_decode.py (기존 상태)
from .kda_kernels import run_recurrent_kda as _run_recurrent_kda

def recurrent_kda(...):
    # Recurrent KDA 연산만 수행
    return _run_recurrent_kda(...)

After: 통합된 Fused API 도입

새로운 API는 Convolution weight, State, Gate, Norm weight 등 모든 파라미터를 한 번에 입력받아 내부적으로 최적화된 커널을 호출합니다.

# flashinfer/kda_decode.py (변경 후)
@flashinfer_api(trace=fused_kda_decode_trace)
def fused_kda_decode(
    x: torch.Tensor,
    weight: torch.Tensor,
    conv_state: torch.Tensor,
    raw_gate: torch.Tensor,
    # ... 중략 ...
    norm_weight: torch.Tensor,
    lower_bound: Optional[float] = -5.0,
    norm_eps: float = 1e-5,
) -> torch.Tensor:
    # SM100 전용 fused 커널 호출
    if _run_fused_kda_decode is None:
        raise NotImplementedError("fused KDA decode backend is unavailable")
    return _run_fused_kda_decode(
        x=x, weight=weight, conv_state=conv_state, # ...
    )

핵심 변경사항: CuTe DSL을 이용한 커널 Fusion

이번 PR의 정수는 flashinfer/kda_kernels/fused_kda_decode.py에 구현된 CuTe DSL(Data Parallel C++ Template Library) 기반의 커널입니다. 이 커널은 SM100의 하드웨어 특성을 최대한 활용하도록 설계되었습니다.

1. 공유 메모리(Shared Memory) 최적화

커널 내부에서 중간 결과물을 저장하기 위해 SmemAllocator를 사용하여 효율적으로 메모리를 할당합니다. 이는 HBM 접근을 차단하고 SRAM 내에서 데이터를 전달하게 해줍니다.

# flashinfer/kda_kernels/fused_kda_decode.py
smem = SmemAllocator()
mixed = smem.allocate_tensor(
    F32, cute.make_layout((3 * _HEAD_DIM,)), byte_alignment=16
)
recurrence_output = smem.allocate_tensor(
    F32, cute.make_layout((_HEAD_DIM,)), byte_alignment=16
)

2. 연산의 수직 통합 (Fusion Logic)

커널 내부에서는 다음과 같은 순서로 연산이 진행됩니다:

  1. Depthwise Causal Convolution: conv_state를 업데이트하며 4-width 컨볼루션을 수행합니다.
  2. SiLU Activation: 컨볼루션 결과에 활성화 함수를 적용합니다.
  3. Recurrent KDA Update: 이전 상태(state)와 현재 입력을 결합하여 새로운 상태를 계산합니다.
  4. Gated RMSNorm: 최종 출력에 게이팅과 정규화를 적용합니다.

이 모든 과정이 단 하나의 CUDA 커널 안에서 RegisterShared Memory를 통해 데이터가 흐르도록 구현되었습니다.

# 커널 내부 로직 (의사 코드)
# Stage 1: Convolution + SiLU
if thread_idx < _CONV_THREADS:
    # ... conv_state 업데이트 및 계산 ...
    
# Stage 2: Recurrence (KDA)
# ... state 업데이트 및 KDA 로직 수행 ...

# Stage 3: Gated RMSNorm
# ... 최종 결과 산출 ...

왜 이게 좋은가?

1. 압도적인 성능 향상 (Speedup)

NVIDIA B200(Blackwell) 환경에서 벤치마크 결과, vLLM의 기존 fused 커널 대비 지오메트릭 평균(Geomean) 1.13배의 성능 향상을 보여주었습니다. 특히 Batch Size가 작은 상황(Rows=1)에서는 최대 1.33배까지 빨라지는 모습을 보였는데, 이는 디코딩 단계의 지연 시간(Latency)을 줄이는 데 결정적인 역할을 합니다.

2. 메모리 대역폭 절약

기존에는 conv_staterecurrent state를 각각 업데이트하기 위해 여러 번의 메모리 쓰기/읽기가 발생했습니다. 이번 구현은 In-place update를 지원하며, 모든 중간 연산을 온칩(On-chip) 메모리에서 처리함으로써 HBM 트래픽을 최소화했습니다.

3. 유연한 Paged Cache 지원

실제 프로덕션 환경(vLLM 등)에서 사용하는 Paged-cache 레이아웃을 그대로 유지하면서도 최적화를 달성했습니다. state_indices를 통해 각 요청이 할당된 슬롯을 정확히 찾아가며, slot 0을 null slot으로 예약하여 예외 처리를 효율적으로 수행합니다.

리뷰어 피드백 반영

코드 리뷰 과정에서 kahyunnam 리뷰어는 다음과 같은 유의미한 피드백을 남겼습니다:

  • 검증 로직 강화: conv_state.shape[0] == state.shape[0]와 같은 텐서 크기 일치 여부를 명시적으로 확인하고 문서화할 것을 권장했습니다.
  • 캐시 설계 정렬: FlashInfer의 공식 cute_dsl_kernel_cache 디자인 가이드라인에 맞춰 커널 캐싱 로직을 개선하도록 유도했습니다.

이러한 피드백은 단순히 기능 구현을 넘어, 라이브러리의 유지보수성과 안정성을 높이는 데 기여했습니다.

마치며

이번 FlashInfer의 Kimi K3 Fused 커널 추가는 최신 GPU 아키텍처인 SM100의 잠재력을 어떻게 끌어올릴 수 있는지 보여주는 좋은 사례입니다. CuTe DSL을 활용한 정교한 메모리 관리와 연산 통합은 LLM 서빙 엔진의 효율을 한 단계 더 진화시켰습니다. 앞으로 Blackwell 기반의 추론 환경에서 FlashInfer의 입지는 더욱 공고해질 것으로 보입니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글