본문으로 건너뛰기

[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.pyflashinfer/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 추론과 같이 동일한 커널을 반복적으로 호출하는 워크로드에서 이러한 오버헤드는 전체 성능에 상당한 영향을 미칠 수 있습니다.

  1. 반복적인 컴파일 및 래핑 제거:

    • @functools.cache를 사용한 커널 객체 및 컴파일 옵션 캐싱은 동일한 설정의 커널이 필요할 때마다 새로 생성하고 컴파일하는 비용을 없앱니다.
    • 가장 중요한 개선은 런타임 시 cute.runtime.from_dlpack 호출 및 텐서 래핑을 제거한 것입니다. 이전에는 매번 커널을 호출할 때마다 PyTorch 텐서를 cute 텐서로 변환하는 과정이 필요했지만, 이제는 컴파일 시점에만 이 변환이 이루어지고, 이후에는 컴파일된 커널 객체에 원본 PyTorch 텐서를 직접 전달합니다. 이는 런타임 오버헤드를 크게 줄여줍니다.
  2. 성능 향상:

    • PR 설명에 따르면, 이 변경은

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글