[flashinfer] NVIDIA SM110 GPU를 위한 실험적인 FP16 GQA 디코드 커널 추가
PR 링크: flashinfer-ai/flashinfer#5052 상태: Merged | 변경: +5108 / -2
들어가며
최근 대규모 언어 모델(LLM)의 발전은 GPU 하드웨어의 성능을 최대한 활용하는 효율적인 추론 커널의 중요성을 더욱 부각시키고 있습니다. 특히, Grouped-Query Attention (GQA)은 기존의 Multi-Query Attention (MQA)보다 더 나은 성능과 확장성을 제공하지만, 이를 위한 최적화된 하드웨어별 커널 구현은 여전히 중요한 연구 과제입니다.
이번 PR은 NVIDIA의 최신 GPU 아키텍처인 SM110 (Jetson AGX Thor)에 특화된 실험적인 FP16 GQA 디코드 커널을 FlashInfer 라이브러리에 추가합니다. 이 커널은 특정 GQA 설정 (32개의 쿼리 헤드, 8개의 키/밸류 헤드, 헤드 차원 128)에 대해 최적화되어 있으며, 기존의 일반적인 GQA 구현 대비 상당한 성능 향상을 목표로 합니다.
이 글에서는 해당 PR의 코드 변경 사항을 분석하고, 왜 이러한 최적화가 성능 향상에 기여하는지, 그리고 이를 통해 얻을 수 있는 일반적인 교훈은 무엇인지 살펴보겠습니다.
코드 분석
이번 PR은 주로 flashinfer/decode.py 파일의 수정과 함께 새로운 실험적인 커널 및 관련 테스트, 벤치마크 코드를 추가하는 데 중점을 둡니다.
1. flashinfer/__init__.py 및 flashinfer/decode.py - 새로운 API 노출
새로운 sm110_gqa_decode 함수가 FlashInfer 라이브러리의 최상위 레벨과 디코드 모듈에 추가되어 외부에서 쉽게 접근할 수 있도록 합니다.
Before:
# flashinfer/__init__.py
# ... (기존 import 문)
from .decode import single_decode_with_kv_cache as single_decode_with_kv_cache
# flashinfer/decode.py
# ... (기존 함수 정의)
After:
# flashinfer/__init__.py
# ... (기존 import 문)
from .decode import single_decode_with_kv_cache as single_decode_with_kv_cache
from .decode import sm110_gqa_decode as sm110_gqa_decode
# flashinfer/decode.py
@flashinfer_experimental_api(feature="SM110 GQA decode")
def sm110_gqa_decode(
q: torch.Tensor,
kv: torch.Tensor,
sequence_lengths: torch.Tensor,
*,
out: Optional[torch.Tensor] = None,
q_scale: float = 1.0,
) -> torch.Tensor:
r"""Decode one token with the exact-SM110 FP16 GQA specialization.
Parameters
----------
q : torch.Tensor
Contiguous FP16 query tensor with shape ``[batch, 32, 128]``.
kv : torch.Tensor
Contiguous FP16 stacked KV tensor with shape
``[batch, 2, 8, capacity, 128]``.
sequence_lengths : torch.Tensor
Contiguous CUDA int32 tensor with shape ``[batch]``.
out : Optional[torch.Tensor]
Optional caller-owned contiguous FP16 output with the same shape and
device as ``q``.
q_scale : float
Additional query scale applied before the standard ``1 / sqrt(128)``
attention scale.
Returns
-------
torch.Tensor
The decoded output in ``out`` when supplied, otherwise a newly
allocated tensor.
"""
from .experimental.sm110_gqa_decode import decode
return decode(
q,
kv,
sequence_lengths,
out=out,
q_scale=q_scale,
)
flashinfer/decode.py 내에서 @flashinfer_experimental_api 데코레이터를 사용하여 이 함수가 실험적인 기능임을 명시하고 있습니다. 이는 사용자가 의도적으로 이 기능을 사용하도록 유도하며, 향후 변경될 수 있음을 알립니다. 또한, q_scale 파라미터가 추가되어 쿼리 스케일링을 조절할 수 있게 되었습니다. 이는 특정 모델 아키텍처나 학습 설정에 맞춰 미세 조정을 가능하게 합니다.
2. flashinfer/experimental/sm110_gqa_decode/ - 핵심 커널 구현
이 디렉토리는 SM110 아키텍처에 최적화된 GQA 디코드 커널의 실제 구현을 포함합니다. PR 설명에 따르면, 이 커널은 다음과 같은 특징을 가집니다:
- 정확한 SM110 JIT 로딩: 특정 하드웨어에 대한 JIT 컴파일 및 SHA-256 검증을 통해 정확성을 보장합니다.
- 분리된 CUDA 커널: 256 스레드 단축 커널 (short-prefix)과 384 스레드 파이프라인 커널 (long-prefix)로 나뉘어, 시퀀스 길이에 따라 최적의 성능을 발휘하도록 설계되었습니다.
- FP16 정밀도: FP16 연산을 사용하여 메모리 대역폭과 연산 효율성을 높입니다.
- 고정된 GQA 설정: 32개의 쿼리 헤드, 8개의 키/밸류 헤드, 헤드 차원 128 (즉, 4:1 쿼리-대-KV 헤드 비율)에 특화되어 있습니다.
이 부분의 구체적인 C++/CUDA 코드는 diff에 직접 포함되지 않았지만, flashinfer/experimental/sm110_gqa_decode/README.md 파일에서 입력 텐서의 형태와 제약 조건에 대한 명확한 설명을 제공합니다.
3. benchmarks/bench_sm110_gqa_decode.py - 성능 벤치마킹
이 스크립트는 새로운 SM110 GQA 디코드 커널의 성능을 측정하고, PyTorch의 표준 Scaled Dot-Product Attention (SDPA) 구현과 비교합니다. CUPTI를 사용하여 cold-L2 캐시 상태에서의 순수 커널 실행 시간을 측정합니다.
주요 측정 로직:
def _measure(fn) -> float:
samples = bench_gpu_time_with_cupti(
fn,
dry_run_time_ms=100,
repeat_time_ms=1000,
cold_l2_cache=True,
)
return float(statistics.median(samples))
def _measure_shape(
label: str,
batch: int,
capacity: int,
lengths: list[int],
seed: int,
) -> dict[str, object]:
# ... (텐서 생성 및 커널 호출 준비)
def candidate() -> torch.Tensor:
return sm110_gqa_decode(q, kv, sequence_lengths)
def torch_sdpa() -> torch.Tensor:
return F.scaled_dot_product_attention(
sdpa_q,
sdpa_k,
sdpa_v,
attn_mask=mask,
scale=1.0 / math.sqrt(128),
enable_gqa=True,
)
candidate()
torch_sdpa()
torch.cuda.synchronize()
candidate_ms = _measure(candidate)
torch_ms = _measure(torch_sdpa)
return {
"label": label,
# ... (결과 포함)
"sm110_gqa_decode_ms": candidate_ms,
"torch_sdpa_ms": torch_ms,
"speedup_vs_torch_sdpa": torch_ms / candidate_ms,
}
이 스크립트는 다양한 배치 크기 및 시퀀스 길이 조합에 대해 sm110_gqa_decode와 torch.nn.functional.scaled_dot_product_attention의 실행 시간을 측정하고, 그 비율을 계산하여 성능 향상 정도를 보여줍니다.
4. examples/experimental/sm110_gqa_decode.py - 사용 예제
이 스크립트는 새로운 sm110_gqa_decode API를 어떻게 사용하는지 보여주는 간단한 예제입니다. 배치 크기, 최대 시퀀스 길이 등의 파라미터를 설정하고 커널을 실행한 후 결과 텐서의 정보를 출력합니다.
5. .pre-commit-config.yaml - Pre-commit Hook 설정
새로운 실험적 코드 경로를 pre-commit hook의 exclude 패턴에 추가하여, 해당 파일들이 자동으로 검사되지 않도록 설정합니다. 이는 실험적인 코드가 아직 안정화되지 않았거나, 특정 빌드/테스트 환경을 요구할 때 유용합니다.
Before:
- exclude: ^(?:csrc/(?:kda/flashkda_generated_.*|blackwell_msa/|concat_mla/|cake_all_gather_matmul/|cake_selective_state_update/(?:cuda|host)/|cake_mamba_ssd_combined/generated/|fused_moe/warp_decode/generated/|kda/(?:cake_flashkda_blackwell_evolution_.*|training_(?:fallback|grouped_row_wg8)_pointer_sm_(?:100a|103a))\.cu)|flashinfer/(?:experimental/cake_mxfp8_megamoe_ep16/csrc/|moe_ep/kernel_src/(?:cutedsl_megamoe|sm90/pull_style_cutedsl_megakernel|sm120/swapab_cutedsl_megakernel)/src/))
+exclude: ^(?:csrc/(?:kda/flashkda_generated_.*|blackwell_msa/|concat_mla/|cake_all_gather_matmul/|cake_selective_state_update/(?:cuda|host)/|cake_mamba_ssd_combined/generated/|fused_moe/warp_decode/generated/|kda/(?:cake_flashkda_blackwell_evolution_.*|training_(?:fallback|grouped_row_wg8)_pointer_sm_(?:100a|103a))\.cu)|flashinfer/(?:experimental/(?:cake_mxfp8_megamoe_ep16|sm110_gqa_decode)/csrc/|moe_ep/kernel_src/(?:cutedsl_megamoe|sm90/pull_style_cutedsl_megakernel|sm120/swapab_cutedsl_megakernel)/src/))
flashinfer/experimental/sm110_gqa_decode/csrc/ 경로가 추가된 것을 볼 수 있습니다.
왜 이게 좋은가?
1. 하드웨어 특화 최적화
이 PR의 핵심 가치는 NVIDIA SM110 아키텍처의 특성을 최대한 활용하는 커널을 제공한다는 점입니다. SM110은 특정 연산, 특히 FP16 정밀도의 행렬 곱셈 및 어텐션 메커니즘에서 높은 성능을 낼 수 있도록 설계되었습니다. 이 커널은:
- 정확한 하드웨어 타겟팅: SM110의 Tensor Core 활용, TMA (Tensor Memory Accelerator), Tensor Memory 등의 기능을 최적으로 사용하도록 설계되었습니다. 이는 일반적인 GPU 커널로는 달성하기 어려운 수준의 성능을 가능하게 합니다.
- 메모리 계층 구조 활용: short-prefix와 long-prefix 커널 분리를 통해 레지스터 사용, 공유 메모리 활용, 스레드 블록 동기화 등을 시퀀스 길이에 따라 최적화합니다. 이는 메모리 대역폭 병목 현상을 줄이고 연산 효율성을 극대화합니다.
2. 성능 향상
PR 설명에 제시된 벤치마크 결과는 이 최적화의 효과를 명확히 보여줍니다.
커널-전용 성능 비교 (vs. XQA):
| Shape | Valid lengths | Exported SM110 kernel | XQA | Speedup vs. XQA |
|---|---|---|---|---|
| B1, capacity 64 | [1] |
0.004512 ms | 0.005792 ms | 1.284x |
| B4, capacity 256 | [64, 127, 191, 256] |
0.026048 ms | 0.032768 ms | 1.258x |
| B1, capacity 1024 | [1024] |
0.035745 ms | 0.041984 ms | 1.175x |
| B1, capacity 4096 | [3968] |
0.117633 ms | 0.124513 ms | 1.058x |
모든 측정 항목에서 SM110 GQA 커널이 기존 XQA 구현보다 빠른 성능을 보였습니다. 특히 짧은 시퀀스 길이에서 1.28배 이상의 속도 향상이 있었습니다.
공개 API 성능 비교 (vs. PyTorch SDPA):
| Shape | Valid lengths | SM110 GQA decode | PyTorch SDPA | Speedup |
|---|---|---|---|---|
| B1, capacity 64 | [1] |
0.161282 ms | 0.541734 ms | 3.36x |
| B4, capacity 256 | [64, 127, 191, 256] |
0.184097 ms | 0.632294 ms | 3.43x |
| B1, capacity 1024 | [1024] |
0.193154 ms | 0.629159 ms | 3.26x |
| B1, capacity 4096 | [3968] |
0.277633 ms | 2.361688 ms | 8.51x |
PyTorch SDPA와의 비교에서는 훨씬 더 큰 성능 향상이 나타났습니다. 이는 PyTorch SDPA가 더 일반적인 목적의 구현인 반면, sm110_gqa_decode는 특정 하드웨어와 GQA 설정에 극도로 최적화되었기 때문입니다. 특히 긴 시퀀스 길이 (4096)에서 8.5배 이상의 압도적인 성능 향상을 보였습니다.
3. 실험적 API의 장점
@flashinfer_experimental_api 데코레이터와 flashinfer/experimental/ 디렉토리 구조는 다음과 같은 이점을 제공합니다:
- 안정성: 아직 개발 중이거나 특정 하드웨어에 종속적인 코드를 메인 코드베이스와 분리하여, 라이브러리 전체의 안정성을 유지합니다.
- 명확한 사용 의도: 사용자가 실험적인 기능을 사용하려면 명시적으로
sm110_gqa_decode와 같은 API를 호출해야 하므로, 의도치 않은 사용이나 호환성 문제를 방지할 수 있습니다. - 점진적 발전: 실험적인 기능을 통해 실제 하드웨어에서의 성능을 검증하고, 피드백을 반영하여 점진적으로 개선한 후 정식 기능으로 편입시킬 수 있습니다.
4. 일반적인 교훈
- 하드웨어 특화는 성능의 핵심: 최신 하드웨어 아키텍처의 기능을 최대한 활용하는 것은 LLM 추론 성능을 극대화하는 데 필수적입니다.
- GQA 최적화의 중요성: GQA는 LLM의 확장성을 높이는 중요한 기법이며, 이를 위한 효율적인 커널 구현은 필수적입니다.
- 실험적 기능 관리: 새로운 기능, 특히 하드웨어 종속적인 기능은 실험적 API로 관리하여 안정성과 명확성을 확보하는 것이 좋습니다.
- 벤치마킹의 중요성: 실제 하드웨어에서 다양한 조건 (cold-L2, 다양한 시퀀스 길이 등)으로 성능을 측정하고 비교하는 것은 최적화의 효과를 검증하는 데 매우 중요합니다.
결론
이번 PR은 NVIDIA SM110 GPU를 위한 실험적인 FP16 GQA 디코드 커널을 성공적으로 추가했습니다. 이는 특정 하드웨어에 대한 깊은 이해를 바탕으로 한 최적화의 좋은 예시이며, LLM 추론 성능 향상에 크게 기여할 것으로 기대됩니다. 특히 PyTorch SDPA 대비 압도적인 성능 향상은 하드웨어 특화 커널의 중요성을 다시 한번 강조합니다. 향후 이 실험적인 기능이 어떻게 발전하고 정식 기능으로 편입될지 주목할 만합니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html
- https://github.com/flashinfer-ai/flashinfer/blob/main/flashinfer/decode.py
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [vllm] vLLM의 멀티모달 추론 성능 극대화: Triton/FlashInfer 복합 어텐션 도입
- [flashinfer] NVIDIA Blackwell(SM103a)을 위한 극한의 커널 퓨전: MiniMax-H3 BF16 Pre-attention 최적화 분석
- [flashinfer] NVFP4 MoE All-to-All 성능 최적화: Phased Dispatch 기법 분석
- [flashinfer] FlashInfer의 NVFP4 KV 캐시 성능 최적화: FP4 연산의 병목 현상 해소
- [flashinfer] FlashInfer SM120 NVFP4 어텐션 최적화: N64 스코어-슬롯 재사용을 통한 성능 향상
PR Analysis 의 다른글
- 이전글 [sglang] SGLang에서 NPU를 위한 LTX-2/2.3 추론 성능 최적화 및 호환성 개선
- 현재글 : [flashinfer] NVIDIA SM110 GPU를 위한 실험적인 FP16 GQA 디코드 커널 추가
- 다음글 [LlamaFactory] LLaMA Factory v1: 멀티모달 및 메모리 효율적인 SFT를 위한 Ulysses CP와 Chunk Loss 지원
댓글