[flashinfer] FlashInfer GDN 커널의 SM90/SM120 비-CP 런치 오버헤드 감소 최적화 분석
PR 링크: flashinfer-ai/flashinfer#4699 상태: Merged | 변경: +267 / -175
들어가며
최근 GPU 컴퓨팅 라이브러리들은 LLM 추론 성능 향상을 위해 끊임없이 최적화를 추구하고 있습니다. 특히, Attention 메커니즘의 핵심 연산 중 하나인 Gated Linear Unit (GLU)의 변형인 Gated Dual-Linear Unit (GDN)은 Transformer 모델에서 중요한 역할을 합니다. 하지만 GDN 연산을 GPU에서 효율적으로 실행하기 위해서는 커널 런치 오버헤드를 최소화하는 것이 필수적입니다.
이번 PR(#4374 후속 작업)은 FlashInfer 라이브러리의 GDN 커널, 특히 SM90 및 SM120 아키텍처를 사용하는 GPU에서 발생하는 비-커널-프로파일링(non-CP) 런치 오버헤드를 줄이는 데 초점을 맞추고 있습니다. 기존에는 각 커널 호출 시마다 컴파일 및 래핑(wrapping) 과정이 반복되어 불필요한 오버헤드가 발생했습니다. 이 PR은 이러한 반복적인 준비 작업을 제거하고, 컴파일된 커널 객체를 재사용함으로써 성능을 크게 향상시킵니다.
본 글에서는 이 PR의 코드 변경 사항을 상세히 분석하고, 왜 이러한 변경이 성능 개선으로 이어지는지, 그리고 어떤 일반적인 교훈을 얻을 수 있는지 살펴보겠습니다.
코드 분석
이번 PR의 핵심 변경 사항은 flashinfer/gdn_kernels/delta_rule_dsl/delta_rule_sm120.py와 flashinfer/gdn_kernels/delta_rule_dsl/delta_rule_sm90.py 파일에 집중되어 있습니다. 두 파일 모두 GDN 연산을 위한 CUDA 커널을 정의하고 있으며, 이번 PR은 이 커널들의 런치 방식을 최적화합니다.
1. SM120 커널 최적화 (delta_rule_sm120.py)
1.1. 커널 컴파일 옵션 캐싱 및 재사용
기존 코드에서는 커널을 컴파일할 때마다 sm12x_compile_options를 사용하고, from_dlpack 함수를 통해 텐서를 래핑하는 과정이 반복되었습니다. 이번 PR에서는 functools.cache 데코레이터를 사용하여 커널 컴파일 옵션과 커널 객체 자체를 캐싱합니다.
Before:
# ... (기존 코드에서는 컴파일 옵션 및 텐서 래핑이 반복적으로 수행됨)
def delta_rule_prefill_dsl(...):
# ...
from_dlpack = lambda *args, **kwargs: cute.runtime.from_dlpack(
*args, **{**kwargs, "enable_tvm_ffi": True}
)
# ...
q_cute = from_dlpack(q_tma, assumed_align=16).mark_layout_dynamic(leading_dim=1)
# ... (다른 텐서들도 유사하게 래핑)
delta_rule_kernel = _FullyFusedDeltaRuleSm120(...)
kernel_args = (
q_cute, k_cute, v_cute, o_cute, alpha_cute, beta_cute, state_cute, init_state_cute, state_indices_cute, state_checkpoints_cute, checkpoint_cu_cute, tensormaps_cute, cu_cute, ...
)
compiled_delta_rule_kernel = cached_compile(
delta_rule_kernel, *kernel_args, compile_options=sm12x_compile_options(q.device),
)
compiled_delta_rule_kernel(*kernel_args)
After:
import functools
# ...
from .custom_compile_cache import (
KeyedCompileMixin,
cached_compile,
get_cached_compile, # 새로 추가
)
@functools.cache
def _sm120_compile_options(device):
return (cute.EnableTVMFFI(True),) + sm12x_compile_options(device)
@functools.cache
def _get_prefill_kernel(
needs_alpha, needs_beta, needs_init_state, needs_checkpointing, kernel_dtype, ...
):
return _FullyFusedDeltaRuleSm120(
# ... (인자들)
)
def delta_rule_prefill_dsl(
# ...
):
# ...
device = q.device
# ...
delta_rule_kernel = _get_prefill_kernel(
# ... (인자들)
)
compile_options = _sm120_compile_options(device)
compiled_delta_rule_kernel = get_cached_compile(delta_rule_kernel, compile_options)
if compiled_delta_rule_kernel is None:
from_dlpack = lambda *args, **kwargs: cute.runtime.from_dlpack(
*args, **{**kwargs, "enable_tvm_ffi": True}
)
kernel_args = (
from_dlpack(q_tma, assumed_align=16).mark_layout_dynamic(leading_dim=1),
# ... (다른 텐서들도 래핑)
from_dlpack(tensormaps_t, assumed_align=128).mark_layout_dynamic(),
from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic(),
# ... (나머지 인자들)
stream,
)
compiled_delta_rule_kernel = cached_compile(
delta_rule_kernel, *kernel_args, compile_options=compile_options,
)
compiled_delta_rule_kernel(
q_tma, k_tma, v_tma, o_tma, ... # 원본 텐서들을 직접 전달
)
설명:
_sm120_compile_options함수는functools.cache를 통해 컴파일 옵션(cute.EnableTVMFFI(True),cute.GPUArch("sm_120")등)을 캐싱합니다. 이는 동일한 디바이스에 대해 컴파일 옵션을 다시 계산하는 오버헤드를 제거합니다._get_prefill_kernel함수 역시functools.cache를 사용하여_FullyFusedDeltaRuleSm120커널 객체 생성을 캐싱합니다. 커널의 동작 방식을 결정하는 다양한 인자(데이터 타입, 체크포인팅 여부 등)가 동일하면 이전에 생성된 커널 객체를 재사용합니다.get_cached_compile함수는cached_compile과 유사하지만, 컴파일된 커널 객체를 먼저 확인하고, 없을 경우에만cached_compile을 호출하여 컴파일을 수행합니다. 이는 컴파일된 커널이 이미 존재하면 불필요한 재컴파일을 방지합니다.- 가장 중요한 변경점 중 하나는,
compiled_delta_rule_kernel호출 시from_dlpack으로 래핑된cute텐서 대신 원본 PyTorch 텐서(q_tma,k_tma등)를 직접 전달한다는 것입니다. 이전에는 매번from_dlpack을 호출하여cute텐서 래퍼를 생성했지만, 이제는 컴파일 시점에만from_dlpack을 사용하여cute텐서 래퍼를 생성하고, 이후 호출에서는 컴파일된 커널 객체에 원본 텐서를 직접 전달합니다. 이는 런타임 시from_dlpack호출 및 텐서 래핑 오버헤드를 제거합니다.
2. SM90 커널 최적화 (delta_rule_sm90.py)
SM90 커널의 최적화 방식은 SM120과 매우 유사합니다. 동일하게 functools.cache를 사용하여 컴파일 옵션(_SM90_COMPILE_OPTIONS)과 커널 객체(_get_prefill_kernel) 생성을 캐싱하고, get_cached_compile을 통해 컴파일된 커널을 효율적으로 관리합니다.
Before: (SM120과 유사하게 반복적인 래핑 및 컴파일 수행)
After:
import functools
# ...
from .custom_compile_cache import cached_compile, get_cached_compile
_SM90_COMPILE_OPTIONS = (cute.EnableTVMFFI(True), cute.GPUArch("sm_90a"))
@functools.cache
def _get_prefill_kernel(
needs_alpha, needs_beta, needs_init_state, needs_checkpointing, kernel_dtype, ...
):
return _FullyFusedDeltaRuleSm90(
# ... (인자들)
)
def delta_rule_prefill_dsl(
# ...
):
# ...
delta_rule_kernel = _get_prefill_kernel(
# ... (인자들)
)
compiled_delta_rule_kernel = get_cached_compile(delta_rule_kernel, _SM90_COMPILE_OPTIONS)
if compiled_delta_rule_kernel is None:
from_dlpack = lambda *args, **kwargs: cute.runtime.from_dlpack(
*args, **{**kwargs, "enable_tvm_ffi": True}
)
kernel_args = (
from_dlpack(q_tma, assumed_align=16).mark_layout_dynamic(leading_dim=1),
# ... (다른 텐서들도 래핑)
stream,
)
compiled_delta_rule_kernel = cached_compile(
delta_rule_kernel, *kernel_args, compile_options=_SM90_COMPILE_OPTIONS,
)
compiled_delta_rule_kernel(
q_tma, k_tma, v_tma, o_tma, ... # 원본 텐서들을 직접 전달
)
설명:
SM120과 동일한 원리로, 컴파일 옵션과 커널 객체 생성을 캐싱하고, 런타임 시 from_dlpack 호출 및 텐서 래핑 오버헤드를 제거하여 성능을 개선합니다. _SM90_COMPILE_OPTIONS는 SM90 아키텍처에 특화된 컴파일 옵션을 포함합니다.
3. 공통 개선 사항
custom_compile_cache.py수정:get_cached_compile함수가 추가되어 컴파일된 커널의 존재 여부를 먼저 확인하고, 없을 경우에만 실제 컴파일을 수행하도록 로직이 개선되었습니다. 이는 동일한 커널에 대한 반복적인 컴파일 시도를 방지합니다.KeyedCompileMixin: 이 믹스인 클래스는 컴파일된 커널을 캐싱하는 메커니즘을 제공하며, 이번 PR에서 캐싱 로직을 구현하는 데 활용되었습니다.
왜 이게 좋은가?
이번 PR의 핵심 목표는 커널 런치 오버헤드 감소입니다. 특히 LLM 추론과 같이 동일한 커널을 반복적으로 호출하는 워크로드에서 이러한 오버헤드는 전체 성능에 상당한 영향을 미칠 수 있습니다.
-
반복적인 컴파일 및 래핑 제거:
@functools.cache를 사용한 커널 객체 및 컴파일 옵션 캐싱은 동일한 설정의 커널이 필요할 때마다 새로 생성하고 컴파일하는 비용을 없앱니다.- 가장 중요한 개선은 런타임 시
cute.runtime.from_dlpack호출 및 텐서 래핑을 제거한 것입니다. 이전에는 매번 커널을 호출할 때마다 PyTorch 텐서를cute텐서로 변환하는 과정이 필요했지만, 이제는 컴파일 시점에만 이 변환이 이루어지고, 이후에는 컴파일된 커널 객체에 원본 PyTorch 텐서를 직접 전달합니다. 이는 런타임 오버헤드를 크게 줄여줍니다.
-
성능 향상:
- PR 설명에 따르면, 이 변경은
참고 자료
- https://github.com/flashinfer-ai/flashinfer/pull/4374
- https://pytorch.org/docs/stable/generated/torch.compile.html
- https://docs.nvidia.com/cuda/cuda-driver-api/group__CUDA__TYPES.html#group__CUDA__TYPES_1g332777047f571716440807189393a48b
- https://github.com/openai/cute/blob/main/cute/runtime.py#L277
- https://github.com/openai/cute/blob/main/cute/compile.py#L118
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer SM12x MoE 최적화: 정적 MoE 경로 통합 및 성능 향상
- [flashinfer] FlashInfer SM120 NVFP4 어텐션 최적화: N64 스코어-슬롯 재사용을 통한 성능 향상
- [flashinfer] FlashInfer, Blackwell 아키텍처를 위한 Recurrent KDA Prefill 최적화: Small-BH 커널 도입
- [flashinfer] FlashInfer, MoE 모델의 성능을 극적으로 향상시키는 융합 커널과 최적화된 스케줄러 도입
- [sglang] SM120 Blackwell에서 DeepSeek-V4 모델 서빙 최적화: FlashInfer MXFP4 MoE 도입 및 메모리 절감
PR Analysis 의 다른글
- 이전글 [open-webui] Open WebUI 스트리밍 성능 190배 개선: O(N^2)에서 O(N)으로의 최적화
- 현재글 : [flashinfer] FlashInfer GDN 커널의 SM90/SM120 비-CP 런치 오버헤드 감소 최적화 분석
- 다음글 [flashinfer] Blackwell 아키텍처를 위한 고성능 Paged MQA Logits 커널 도입
댓글