[flashinfer] FlashInfer, 최신 GPU 아키텍처를 위한 커널 튜닝으로 성능 극대화
PR 링크: flashinfer-ai/flashinfer#5305 상태: Merged | 변경: +38 / -6
들어가며
딥러닝 모델의 성능은 하드웨어의 잠재력을 얼마나 잘 끌어내는지에 달려있습니다. 특히 대규모 언어 모델(LLM)과 같이 연산 집약적인 작업에서는 GPU 커널 수준의 최적화가 전체 성능에 지대한 영향을 미칩니다. 이번 PR은 FlashInfer 라이브러리에서 최신 NVIDIA GPU 아키텍처(Blackwell 및 Rubin 시리즈, SM100, SM103, SM107)에 특화된 커널 튜닝을 통해 qk_rmsnorm 및 fused_add_rmsnorm_quant 연산의 성능을 향상시키는 것을 목표로 합니다.
기존에는 모든 아키텍처에 대해 동일한 커널 튜닝 전략이 적용되었지만, 이 PR에서는 특정 아키텍처 리스트(_LATENCY_BOUND_SMS)에 해당하는 GPU에서 발생하는 성능 병목 현상을 해결하기 위한 세밀한 조정을 수행합니다. 이는 커널의 스레드 구성 및 메모리 접근 패턴을 최적화하여, 특히 작은 head_dim이나 큰 시퀀스 길이를 처리할 때 발생하는 지연 시간을 줄이는 데 중점을 둡니다.
코드 분석
이번 PR의 핵심은 flashinfer/norm/kernels/rmsnorm.py와 flashinfer/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_dim64: M=512 (1.00×), M=8192 (1.41×), M=32768 (1.70×) 향상head_dim128: M=512 (1.09×), M=8192 (1.75×), M=32768 (2.13×) 향상head_dim256: M=512 (1.18×), M=8192 (1.93×), M=32768 (2.10×) 향상
-
qk_rmsnorm성능 (B200/SM100, PDL off, bf16):head_dim64: M=2048 (1.00×), M=8192 (1.13×), M=32768 (1.16×) 향상head_dim128: M=2048 (1.00×), M=8192 (1.14×), M=32768 (1.20×) 향상head_dim256: 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×) 향상
이러한 성능 향상은 주로 다음과 같은 이유로 달성되었습니다:
- 아키텍처별 최적화: 최신 GPU 아키텍처(SM100, 103, 107)는 이전 세대와 다른 성능 특성을 가집니다. 이 PR은 해당 아키텍처에서 발생하는 지연 시간 병목 현상을 정확히 파악하고, 스레드 구성을 조정하여 GPU 코어의 활용률을 높였습니다. 특히
qk_rmsnorm에서 작은head_dim에 대해 스레드당 처리량을 늘린 것이 효과적이었습니다. - 메모리 대역폭 및 지연 시간 균형:
fused_add_rmsnorm_quant의 경우,H > 8192일 때 스레드 수를 256으로 늘리는 것이 항상 유리하지 않다는 점을 파악했습니다. 최신 아키텍처에서는 이로 인해 공유 메모리 사용량이 늘고 대역폭 병목이 발생할 수 있으므로, 해당 아키텍처에서는 이 최적화를 비활성화하여 성능을 유지하거나 개선했습니다. - 일관된 성능: 특정 아키텍처에 대한 최적화는 다른 아키텍처의 성능을 저하시키지 않도록 주의 깊게 설계되었습니다.
_LATENCY_BOUND_SMS리스트를 사용하여 변경 사항을 명시적으로 제한함으로써, 기존 아키텍처에서의 성능은 그대로 유지됩니다.
일반적인 교훈
- 하드웨어 특성 이해의 중요성: GPU 아키텍처는 계속 발전하며, 각 세대마다 성능 특성이 달라집니다. 라이브러리는 이러한 변화에 맞춰 커널을 지속적으로 튜닝해야 합니다. 단순히 일반적인 최적화를 적용하는 것을 넘어, 특정 아키텍처의 강점과 약점을 이해하고 이를 활용하는 것이 중요합니다.
- 측정 기반 최적화: 성능 개선은 추측이 아닌 실제 측정 데이터를 기반으로 이루어져야 합니다. 이 PR은 다양한
head_dim,M,H값 및 GPU 아키텍처에 대한 상세한 성능 벤치마크 결과를 제시하며, 이를 통해 최적화의 효과를 명확히 입증했습니다. - 조건부 최적화: 모든 최적화가 모든 상황에 적용되는 것은 아닙니다. 특정 조건(예: 하드웨어 아키텍처, 입력 데이터의 크기)에 따라 최적의 전략이 달라질 수 있습니다. 이를 고려하여 조건부 로직을 구현하는 것이 중요합니다.
References
- NVIDIA CUDA Compute Capability 10.x (SM100, SM103, SM107 포함)
- [CuTe-DSL Documentation](https://github.com/ விளைinfer-ai/flashinfer/blob/main/docs/cutedsl.md) (커널 구현에 사용된 DSL)
참고 자료
- https://docs.nvidia.com/cuda/cuda-toolkit-release-notes/index.html#cuda-major-component-versions__sm-10
- https://github.com/flashinfer-ai/flashinfer/blob/main/docs/cutedsl.md
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer SM12x MoE 최적화: 정적 MoE 경로 통합 및 성능 향상
- [flashinfer] FlashInfer, MoE 모델의 성능을 극적으로 향상시키는 융합 커널과 최적화된 스케줄러 도입
- [flashinfer] FlashInfer의 plan() 함수 최적화: Python max()에서 Tensor.max()로의 전환
- [flashinfer] FlashInfer, BF16 활성화 및 MXFP8 가중치에 대한 Cake MegaMoE EP16 백엔드 최적화
- [vllm] vLLM의 PLE 메타데이터 전송 최적화: 비동기 전송으로 성능 향상
PR Analysis 의 다른글
- 이전글 [onnxruntime] ONNX Runtime CUDA 데이터 로딩 최적화: Pinned Buffer와 병렬 I/O를 통한 성능 개선
- 현재글 : [flashinfer] FlashInfer, 최신 GPU 아키텍처를 위한 커널 튜닝으로 성능 극대화
- 다음글 [sglang] [DeepSeek-V4.1] Blackwell(SM100) 성능의 한계를 끌어올리는 Fused WO-A 커널 최적화 분석
댓글