본문으로 건너뛰기

[flashinfer] FlashInfer, 최신 GPU 아키텍처를 위한 커널 튜닝으로 성능 극대화

PR 링크: flashinfer-ai/flashinfer#5305 상태: Merged | 변경: +38 / -6

들어가며

딥러닝 모델의 성능은 하드웨어의 잠재력을 얼마나 잘 끌어내는지에 달려있습니다. 특히 대규모 언어 모델(LLM)과 같이 연산 집약적인 작업에서는 GPU 커널 수준의 최적화가 전체 성능에 지대한 영향을 미칩니다. 이번 PR은 FlashInfer 라이브러리에서 최신 NVIDIA GPU 아키텍처(Blackwell 및 Rubin 시리즈, SM100, SM103, SM107)에 특화된 커널 튜닝을 통해 qk_rmsnormfused_add_rmsnorm_quant 연산의 성능을 향상시키는 것을 목표로 합니다.

기존에는 모든 아키텍처에 대해 동일한 커널 튜닝 전략이 적용되었지만, 이 PR에서는 특정 아키텍처 리스트(_LATENCY_BOUND_SMS)에 해당하는 GPU에서 발생하는 성능 병목 현상을 해결하기 위한 세밀한 조정을 수행합니다. 이는 커널의 스레드 구성 및 메모리 접근 패턴을 최적화하여, 특히 작은 head_dim이나 큰 시퀀스 길이를 처리할 때 발생하는 지연 시간을 줄이는 데 중점을 둡니다.

코드 분석

이번 PR의 핵심은 flashinfer/norm/kernels/rmsnorm.pyflashinfer/norm/kernels/fused_add_rmsnorm.py 파일의 수정입니다. 각 변경 사항을 자세히 살펴보겠습니다.

1. flashinfer/norm/kernels/rmsnorm.py 수정

이 파일에서는 QKRMSNormKernel의 스레드 구성을 최신 GPU 아키텍처에 맞게 조정합니다.

qk_rmsnorm: 작은 head_dim에 대한 스레드 수 조정

기존 QKRMSNormKernel은 RMSNorm의 스레드당 행(threads per row) 테이블을 재사용했습니다. 이로 인해 head_dim이 256 이하일 경우, 각 스레드는 단일 16-byte 벡터만 처리하게 됩니다. 이는 SM(Streaming Multiprocessor)당 최대 2048개의 활성 스레드를 고려할 때, SM당 32KB의 데이터만 처리하게 되어 대규모 M (시퀀스 길이)에서 지연 시간 병목 현상을 유발했습니다.

이번 PR에서는 _LATENCY_BOUND_SMS에 포함된 아키텍처(SM100, SM103, SM107)에 대해 _compute_threads_per_row 함수를 수정하여, 각 스레드가 더 많은 데이터를 처리하도록 조정했습니다. 구체적으로, head_dim이 64, 128, 256일 때 스레드당 각각 4, 4, 8개의 요소를 처리하도록 변경되었습니다 (기존에는 8, 16, 32개였습니다).

Before:

        self.threads_per_row = RMSNormKernel._compute_threads_per_row(self.head_dim)

After:

        self.threads_per_row = self._compute_threads_per_row(head_dim, self.sm_version)

    @staticmethod
    def _compute_threads_per_row(head_dim: int, sm_version: int) -> int:
        """Threads cooperating on one (batch, head) row.

        The shared RMSNorm table leaves each thread with a single 16-byte vector
        for small head_dim, which is too little in flight per SM on Blackwell and
        Rubin. There, use fewer threads per row so each thread owns ~32 elements.
        """
        default = RMSNormKernel._compute_threads_per_row(head_dim)
        if sm_version not in _LATENCY_BOUND_SMS:
            return default
        target = max(head_dim // 32, 1)
        target = 1 << (target.bit_length() - 1)  # power of two for warp shuffles
        return min(default, max(4, target))

또한, QKRMSNormKernel 생성자에 sm_version 인자가 추가되었으며, 이는 컴파일 캐시 키의 일부로 사용되어 아키텍처별 최적화가 올바르게 적용되도록 합니다.

Before:

def __init__(
        self,
        dtype: cutlass.Numeric,
        head_dim: int,
        weight_bias: float = 0.0,
    ):

After:

def __init__(
        self,
        dtype: cutlass.Numeric,
        head_dim: int,
        weight_bias: float = 0.0,
        sm_version: int | None = None,
    ):
        # ...
        self.sm_version = sm_version if sm_version is not None else get_sm_version()

그리고 _get_compiled_qk_rmsnorm_kernel 함수도 sm_version을 인자로 받도록 수정되었습니다.

Before:

def _get_compiled_qk_rmsnorm_kernel(
    dtype_str: str, head_dim: int, weight_bias: float, enable_pdl: bool
):
    # ...
    kernel_obj = QKRMSNormKernel(dtype, head_dim, weight_bias)

After:

def _get_compiled_qk_rmsnorm_kernel(
    dtype_str: str,
    head_dim: int,
    weight_bias: float,
    enable_pdl: bool,
    sm_version: int,
):
    # ...
    kernel_obj = QKRMSNormKernel(dtype, head_dim, weight_bias, sm_version=sm_version)

마지막으로, qk_rmsnorm_cute 함수에서 _get_compiled_qk_rmsnorm_kernel을 호출할 때 현재 GPU의 sm_version을 전달하도록 변경되었습니다.

Before:

    kernel = _get_compiled_qk_rmsnorm_kernel(
        dtype_str, head_dim, weight_bias, enable_pdl
    )

After:

    kernel = _get_compiled_qk_rmsnorm_kernel(
        dtype_str, head_dim, weight_bias, enable_pdl, get_sm_version(input.device)
    )

2. flashinfer/norm/kernels/fused_add_rmsnorm.py 수정

이 파일에서는 FusedAddRMSNormQuantKernel에서 특정 조건 하에 스레드 수를 조정하는 로직을 수정합니다.

fused_add_rmsnorm_quant: 256 스레드 증가 조건 완화

기존 코드에서는 H_per_cta (CTA당 처리하는 헤드 차원)가 8192보다 크고 num_threads가 256 미만일 경우, num_threads를 256으로 증가시켰습니다. 이는 H > 8192일 때 두 개의 행(row)을 CTA당 처리하게 되어 공유 메모리 사용량이 두 배가 되고 활성 CTA 수가 절반으로 줄어드는 효과가 있었습니다. 하지만 Blackwell 및 Rubin 아키텍처에서는 이로 인해 대역폭 병목이 발생할 수 있습니다.

이번 PR에서는 이 256 스레드 증가 로직이 self.sm_version not in _LATENCY_BOUND_SMS 조건 하에서만 적용되도록 변경되었습니다. 즉, 지연 시간 병목이 발생하는 아키텍처에서는 이 조건이 발동하지 않아 기존의 128 스레드 구성을 유지하게 됩니다.

Before:

        if self.H_per_cta > 8192 and self.num_threads < 256:
            self.num_threads = 256

After:

        if (
            self.H_per_cta > 8192
            and self.num_threads < 256
            and self.sm_version not in _LATENCY_BOUND_SMS
        ):
            self.num_threads = 256

왜 이게 좋은가?

성능 향상

이 PR은 최신 GPU 아키텍처에 대한 세밀한 튜닝을 통해 상당한 성능 향상을 가져왔습니다.

  • qk_rmsnorm 성능 (SM107, PDL on, bf16):

    • head_dim 64: M=512 (1.00×), M=8192 (1.41×), M=32768 (1.70×) 향상
    • head_dim 128: M=512 (1.09×), M=8192 (1.75×), M=32768 (2.13×) 향상
    • head_dim 256: M=512 (1.18×), M=8192 (1.93×), M=32768 (2.10×) 향상
  • qk_rmsnorm 성능 (B200/SM100, PDL off, bf16):

    • head_dim 64: M=2048 (1.00×), M=8192 (1.13×), M=32768 (1.16×) 향상
    • head_dim 128: M=2048 (1.00×), M=8192 (1.14×), M=32768 (1.20×) 향상
    • head_dim 256: M=2048 (1.13×), M=8192 (1.32×), M=32768 (1.36×) 향상
    • 전체 18개 셀에 대한 기하 평균 1.07×, M ≥ 2048의 경우 1.15× 향상
  • fused_add_rmsnorm_quant 성능 (M=32768, bf16):

    • SM107: H=12288 (1.07×), H=14336 (1.10×), H=16384 (0.98×) 향상

이러한 성능 향상은 주로 다음과 같은 이유로 달성되었습니다:

  1. 아키텍처별 최적화: 최신 GPU 아키텍처(SM100, 103, 107)는 이전 세대와 다른 성능 특성을 가집니다. 이 PR은 해당 아키텍처에서 발생하는 지연 시간 병목 현상을 정확히 파악하고, 스레드 구성을 조정하여 GPU 코어의 활용률을 높였습니다. 특히 qk_rmsnorm에서 작은 head_dim에 대해 스레드당 처리량을 늘린 것이 효과적이었습니다.
  2. 메모리 대역폭 및 지연 시간 균형: fused_add_rmsnorm_quant의 경우, H > 8192일 때 스레드 수를 256으로 늘리는 것이 항상 유리하지 않다는 점을 파악했습니다. 최신 아키텍처에서는 이로 인해 공유 메모리 사용량이 늘고 대역폭 병목이 발생할 수 있으므로, 해당 아키텍처에서는 이 최적화를 비활성화하여 성능을 유지하거나 개선했습니다.
  3. 일관된 성능: 특정 아키텍처에 대한 최적화는 다른 아키텍처의 성능을 저하시키지 않도록 주의 깊게 설계되었습니다. _LATENCY_BOUND_SMS 리스트를 사용하여 변경 사항을 명시적으로 제한함으로써, 기존 아키텍처에서의 성능은 그대로 유지됩니다.

일반적인 교훈

  • 하드웨어 특성 이해의 중요성: GPU 아키텍처는 계속 발전하며, 각 세대마다 성능 특성이 달라집니다. 라이브러리는 이러한 변화에 맞춰 커널을 지속적으로 튜닝해야 합니다. 단순히 일반적인 최적화를 적용하는 것을 넘어, 특정 아키텍처의 강점과 약점을 이해하고 이를 활용하는 것이 중요합니다.
  • 측정 기반 최적화: 성능 개선은 추측이 아닌 실제 측정 데이터를 기반으로 이루어져야 합니다. 이 PR은 다양한 head_dim, M, H 값 및 GPU 아키텍처에 대한 상세한 성능 벤치마크 결과를 제시하며, 이를 통해 최적화의 효과를 명확히 입증했습니다.
  • 조건부 최적화: 모든 최적화가 모든 상황에 적용되는 것은 아닙니다. 특정 조건(예: 하드웨어 아키텍처, 입력 데이터의 크기)에 따라 최적의 전략이 달라질 수 있습니다. 이를 고려하여 조건부 로직을 구현하는 것이 중요합니다.

References

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글