[flashinfer] FlashInfer, CuTe DSL을 활용한 저지연 GEMM 커널 도입으로 성능 극대화
PR 링크: flashinfer-ai/flashinfer#4685 상태: Merged | 변경: +3704 / -95
들어가며
최근 GPU 기반 딥러닝 모델의 규모가 기하급수적으로 증가하면서, 모델 추론 속도 향상을 위한 연산 최적화는 그 어느 때보다 중요해졌습니다. 특히 행렬 곱셈(GEMM, General Matrix Multiply)은 딥러닝 연산의 핵심으로, 이 연산의 효율성을 높이는 것은 전체 모델 성능에 직접적인 영향을 미칩니다. NVIDIA의 최신 Blackwell 아키텍처(SM100, SM103)는 FP4, FP8과 같은 저정밀도 데이터 타입을 활용하여 GEMM 연산의 처리량을 극대화할 수 있는 잠재력을 가지고 있습니다.
이번 글에서는 FlashInfer 라이브러리의 PR(#3141)에서 이루어진 중요한 코드 변경 사항을 분석하여, CuTe DSL(Domain Specific Language)을 활용한 저지연(low-latency) GEMM 커널 도입이 어떻게 FP4 및 FP8 연산의 성능을 혁신적으로 개선하는지 살펴보겠습니다. 이 PR은 기존 FlashInfer GEMM API에 Blackwell 아키텍처에 최적화된 새로운 커널을 통합하여, 특히 저정밀도 데이터 타입에서의 추론 속도를 크게 향상시키는 것을 목표로 합니다.
코드 분석
이번 PR의 핵심은 CuTe DSL을 활용한 저지연 블록 스케일드(block-scaled) GEMM 커널을 도입하고, 이를 기존 FlashInfer GEMM API에 통합하는 것입니다. 변경 사항은 주로 flashinfer/gemm/__init__.py와 flashinfer/gemm/gemm_base.py 파일에 집중되어 있습니다.
1. flashinfer/gemm/__init__.py
이 파일에서는 새로운 CuTe DSL 기반 커널을 FlashInfer 라이브러리 내부적으로 인식하고 사용할 수 있도록 등록하는 역할을 합니다.
Before:
from .kernels.cute_dsl.sm100_block_scaled_persistent_dense_gemm import (
Sm100BlockScaledPersistentDenseGemmKernel as Sm100BlockScaledPersistentDenseGemmKernel,
create_scale_factor_tensor as create_scale_factor_tensor,
)
_cute_dsl_kernels = [
"grouped_gemm_nt_masked",
"Sm100BlockScaledPersistentDenseGemmKernel",
"create_scale_factor_tensor",
]
After:
from .kernels.cute_dsl.sm100_block_scaled_persistent_dense_gemm import (
Sm100BlockScaledPersistentDenseGemmKernel as Sm100BlockScaledPersistentDenseGemmKernel,
create_scale_factor_tensor as create_scale_factor_tensor,
)
from .kernels.cute_dsl.low_latency_blockscaled_gemm import (
LowLatencyBlockscaledGemmKernel as LowLatencyBlockscaledGemmKernel,
)
_cute_dsl_kernels = [
"grouped_gemm_nt_masked",
"Sm100BlockScaledPersistentDenseGemmKernel",
"create_scale_factor_tensor",
"LowLatencyBlockscaledGemmKernel",
]
설명:
_cute_dsl_kernels 리스트에 새로 추가된 "LowLatencyBlockscaledGemmKernel"은 CuTe DSL을 사용하여 구현된 저지연 블록 스케일드 GEMM 커널을 나타냅니다. 이 커널은 FP4, FP8 및 혼합 FP8xFP4 연산을 지원하며, 이를 통해 기존 커널보다 더 낮은 지연 시간으로 높은 처리량을 달성할 수 있습니다.
2. flashinfer/gemm/gemm_base.py
이 파일은 GEMM 연산의 핵심 로직과 다양한 백엔드(backend)를 관리합니다. 이번 PR에서는 새로운 저지연 커널을 위한 요구사항 검증 함수와 tgv_gemm_sm100 함수의 수정이 이루어졌습니다.
2.1. _cutedsl_low_latency_blockscaled_tgv_requirement 함수 추가
이 함수는 새로운 저지연 블록 스케일드 GEMM 커널이 요구하는 입력 텐서의 속성(데이터 타입, 스케일 팩터, 차원, 정렬 등)을 검증합니다. FP4, FP8 연산 및 스케일 팩터의 존재 여부, 데이터 타입 일치 여부 등을 엄격하게 체크하여 커널이 올바르게 동작하도록 보장합니다.
코드 인용 (일부):
@supported_compute_capability([100, 103])
def _cutedsl_low_latency_blockscaled_tgv_requirement(
a: torch.Tensor,
b: torch.Tensor,
bias: torch.Tensor,
a_descale: Optional[torch.Tensor],
b_descale: Optional[torch.Tensor],
out: Optional[torch.Tensor] = None,
):
if a_descale is None or b_descale is None:
raise ValueError("Block-scaled TGV inputs require a_descale and b_descale")
fp4_dtype = get_native_fp4_dtype()
quantized_dtypes = (fp4_dtype, torch.float8_e4m3fn, torch.float8_e5m2)
if (
a.dtype not in quantized_dtypes
or b.dtype not in quantized_dtypes
or a_descale.dtype != b_descale.dtype
):
raise ValueError(
"Block-scaled TGV requires FP4/FP8 operands and matching scale dtypes"
)
# ... (이하 생략)
return True
설명:
이 함수는 저지연 커널의 정확성을 보장하기 위한 필수적인 전처리 단계입니다. FP4/FP8 연산은 스케일 팩터(a_descale, b_descale) 없이는 올바르게 수행될 수 없으므로, 이들이 반드시 제공되어야 함을 명시합니다. 또한, 입력 텐서의 데이터 타입과 스케일 팩터의 데이터 타입이 호환되는지 확인합니다. 이는 다양한 저정밀도 데이터 타입 조합에 대한 안정적인 동작을 보장합니다.
2.2. tgv_gemm_sm100 함수 수정
tgv_gemm_sm100 함수는 SM100 아키텍처를 위한 TGV GEMM 연산을 수행합니다. 이번 PR에서는 이 함수가 새로운 저지연 백엔드를 지원하도록 확장되었습니다.
Before (핵심 로직 일부):
# Verify SM100 architecture support
if not _match_sm_version(a.device, ["100", "103"]):
raise ValueError("TGV GEMM requires SM100, SM103 architecture")
# Verify dtype support
if a.dtype not in [torch.bfloat16, torch.float16]:
raise ValueError(
f"Unsupported dtype {a.dtype}. Only bfloat16 and float16 are supported."
)
if a.dtype != b.dtype:
raise ValueError(
f"Input tensors must have the same dtype. Got {a.dtype} and {b.dtype}."
)
if out is None:
out = torch.empty(
(a.shape[0], b.shape[1]),
device=a.device,
dtype=a.dtype,
)
else:
# ... (out tensor checks)
runners = []
use_sm_100f = is_sm100f_supported(a.device)
runners.append(get_tgv_gemm_sm10x_module(a.dtype, use_sm_100f).tgv_gemm_runner())
tuner = AutoTuner.get()
# ...
After (핵심 로직 일부):
fp4_dtype = get_native_fp4_dtype()
quantized_dtypes = (fp4_dtype, torch.float8_e4m3fn, torch.float8_e5m2)
is_blockscaled = a.dtype in quantized_dtypes
if is_blockscaled:
_cutedsl_low_latency_blockscaled_tgv_requirement(
a, b, bias, a_descale, b_descale, out
)
# cast: mypy doesn't know that a_descale and b_descale are not None
a_descale = cast(torch.Tensor, a_descale)
b_descale = cast(torch.Tensor, b_descale)
out_dtype = bias.dtype
else:
if a_descale is not None or b_descale is not None:
raise ValueError("Scale factors require block-scaled FP4/FP8 inputs")
if a.dtype not in [torch.bfloat16, torch.float16]:
raise ValueError(
f"Unsupported dtype {a.dtype}. Only bfloat16 and float16 are supported."
)
if a.dtype != b.dtype:
raise ValueError(
f"Input tensors must have the same dtype. Got {a.dtype} and {b.dtype}."
)
out_dtype = a.dtype
if out is None:
out = torch.empty(
(a.shape[0], b.shape[1]),
device=a.device,
dtype=out_dtype,
)
else:
# ... (out tensor checks)
if is_blockscaled:
runner = _cutedsl_low_latency_blockscaled_gemm_runner(
get_compute_capability(a.device)[0] * 10
+ get_compute_capability(a.device)[1],
pdl,
)
logical_k = a.shape[1] * (2 if a.dtype == fp4_dtype else 1)
inputs = [
b.T, # Note: b is column-major
a,
b_descale,
a_descale,
out.T, # Note: out is column-major
_get_cache_buf(
"tgv_gemm_sm100_blockscaled_workspace",
DEFAULT_WORKSPACE_SIZE,
a.device,
),
(b.shape[1], a.shape[0], logical_k, 1), # problem_mnkl
None, # alpha
bias,
]
runners = [runner]
tuning_config = TuningConfig()
dtype_str = f"{a.dtype}_{b.dtype}_{a_descale.dtype}"
else:
runners = [
get_tgv_gemm_sm10x_module(
a.dtype, is_sm100f_supported(a.device)
).tgv_gemm_runner()
]
inputs = [a, b, bias, pdl, out]
tuning_config = TuningConfig(
dynamic_tensor_specs=(
DynamicTensorSpec(
(0,),
(-2,),
get_hybrid_num_tokens_buckets,
map_to_hybrid_bucket_uncapped,
),
),
constraint_specs=(ConstraintSpec(4, -2, lambda shapes: shapes[0][-2]),),
)
dtype_str = "bf16" if a.dtype == torch.bfloat16 else "fp16"
tuner = AutoTuner.get()
runner, tactic = tuner.choose_one(
f"{dtype_str}_tgv_gemm",
runners,
inputs,
tuning_config,
)
# ... (rest of the function)
설명:
- 저정밀도 데이터 타입 감지:
is_blockscaled변수를 통해 입력 텐서의 데이터 타입이 FP4 또는 FP8인지 확인합니다. 이를 기반으로_cutedsl_low_latency_blockscaled_tgv_requirement함수를 호출하여 입력 유효성을 검사합니다. - 출력 텐서 타입 결정: 블록 스케일드 연산의 경우, 출력 텐서의 데이터 타입은
bias의 데이터 타입(bias.dtype)을 따릅니다. 이는 연산의 정확성을 유지하기 위함입니다. 일반적인 BF16/FP16 연산의 경우 기존과 같이 입력 텐서의 데이터 타입을 따릅니다. - 입력 텐서 구성: 블록 스케일드 연산 시
inputs리스트가 재구성됩니다. 기존의a,b,bias외에 스케일 팩터(a_descale,b_descale), 임시 작업 공간(workspace), 문제 차원(problem_mnkl), 그리고alpha값이 추가됩니다. 이는 새로운 저지연 커널이 요구하는 인자들입니다. 특히b와out텐서는 내부적으로 전치(transpose)되어 전달되는데, 이는 커널이 특정 레이아웃을 기대하기 때문입니다. - 자동 튜닝:
AutoTuner는 새로운 백엔드(_cutedsl_low_latency_blockscaled_gemm_runner)를 포함하여 최적의 연산 전술(tactic)을 선택합니다. 블록 스케일드 연산의 경우,TuningConfig가 단순화되어 새로운 커널에 최적화된 튜닝 전략을 사용합니다.
3. CuTe DSL 및 관련 라이브러리
이 PR은 nvidia-cutlass-dsl 라이브러리의 기능을 활용합니다. 특히 CuTe DSL은 GPU 하드웨어의 특성을 최대한 활용하여 고성능 커널을 생성하는 데 사용됩니다. 리뷰 코멘트에서 cutlass.cute.dsmem API의 부재 및 버전 호환성 문제가 지적되었으며, 이는 nvidia-cutlass-dsl 버전을 4.8.0a0 이상으로 업데이트하거나 해당 기능을 사용하는 코드를 수정하여 해결되었습니다. 또한, assume 매크로의 잘못된 사용으로 인한 잠재적 미스컴파일(miscompile) 문제도 발견되어 수정되었습니다.
왜 이게 좋은가?
이번 PR은 다음과 같은 이유로 매우 긍정적인 성능 개선을 가져옵니다:
- 저지연 및 고처리량: CuTe DSL 기반의 저지연 블록 스케일드 GEMM 커널은 NVIDIA Blackwell 아키텍처(SM100/SM103)에 특화되어 설계되었습니다. 이는 FP4, FP8과 같은 저정밀도 데이터 타입을 활용하여 기존 FP16/BF16 연산 대비 훨씬 낮은 지연 시간과 높은 처리량을 제공합니다. 특히, M 차원이 8 이하이고 K 차원이 특정 배수(FP4: 64, FP8: 128)인 경우에 최적화된 성능을 발휘합니다.
- 다양한 정밀도 지원: FP4 (NVFP4), FP8 (E4M3FN, E5M2), 혼합 FP8xFP4 연산을 모두 지원하며, 스케일 팩터와 옵션으로 Bias, FP32 출력 스케일을 적용할 수 있습니다. 이는 다양한 양자화(quantization) 기법 및 모델 요구사항에 유연하게 대응할 수 있게 합니다.
- 자동 최적화:
tgv_gemm_sm100함수는 입력 데이터의 특성에 따라 최적의 GEMM 커널(기존의 밀집(dense) 커널 또는 새로운 저지연 블록 스케일드 커널)을 자동으로 선택합니다. 또한,AutoTuner를 통해 다양한 전술(tactic)을 탐색하여 주어진 하드웨어 및 입력 조건에 가장 적합한 연산 방식을 찾아냅니다. - 코드 재사용 및 확장성: 기존의 TGV GEMM 프레임워크를 활용하면서 새로운 CuTe DSL 커널을 통합함으로써, 코드의 재사용성을 높이고 향후 새로운 아키텍처나 연산 방식에 대한 확장을 용이하게 합니다.
- 정확성 보장: PR 설명에 따르면, 새로운 커널과 관련된 모든 생성된 전술에 대해 광범위한 정확성 테스트(GPT-OSS-120B, DeepSeek-V3 등 다양한 모델의 형태 포함)가 수행되었습니다. 이는 성능 향상과 더불어 연산의 정확성을 보장하는 데 중점을 두었음을 보여줍니다.
성능 수치:
PR 설명에는 구체적인 성능 수치가 명시적으로 포함되어 있지 않지만, "low-latency"라는 이름 자체가 지연 시간 감소를 목표로 함을 시사합니다. CuTe DSL과 저정밀도 데이터 타입의 활용은 일반적으로 기존 방식 대비 수 배의 처리량 향상을 가져올 수 있습니다. 특히 Blackwell 아키텍처의 Tensor Core는 이러한 저정밀도 연산에 최적화되어 있어 상당한 성능 개선이 기대됩니다.
리뷰 피드백 및 교훈
리뷰 과정에서 몇 가지 중요한 기술적 논의와 개선 사항이 있었습니다:
- API 명명 규칙:
low_latency백엔드 이름이cutedsl_low_latency로 변경되었습니다. 이는 해당 백엔드가 CuTe DSL을 기반으로 함을 명확히 하여 API의 가독성과 이해도를 높입니다. - 라이브러리 버전 호환성:
nvidia-cutlass-dsl라이브러리의 특정 버전(4.8.0a0이상)이cute.dsmem과 같은 기능을 사용하기 위해 필요하다는 점이 지적되었습니다. 이는 외부 라이브러리 의존성 관리의 중요성을 보여줍니다. PR에서는pyproject.toml및requirements.txt파일의 버전 핀(pin)을 업데이트하여 이 문제를 해결했습니다. assume매크로의 올바른 사용: CuTe DSL의assume매크로는 컴파일러에게 특정 제약 조건을 알려주어 최적화를 유도하지만, 잘못 사용될 경우 미스컴파일(miscompile)을 유발할 수 있습니다. 리뷰어는assume(m, 32)와 같이 잘못된 가정으로 인해 실제 입력 크기와 충돌하는 경우를 발견하고 수정했습니다. 이는 DSL 사용 시 해당 제약 조건이 실제 데이터 흐름과 일치하는지 철저히 검증해야 함을 시사합니다.- 테스트 커버리지 및 버그 수정: 여러 테스트 파일에서 백엔드 이름이 변경된 것을 반영하지 않아 발생하는 버그가 발견되었습니다 (
test_mm_fp8.py,test_mm_fp4.py,test_mm_mxfp8.py). 또한, FP8 테스트에서atol=1e-1이 너무 커서 실제 연산 결과와 무관하게 통과하는 문제가 발견되어, 테스트의 정확성을 높이기 위해 FP8 연산의 결과 범위를 고려한 검증 로직이 추가되었습니다. 이는 새로운 기능 추가 시 기존 테스트 스위트와의 호환성 및 정확성 검증의 중요성을 강조합니다. - PDL(Parallel Data Layout) 관련 동기화 문제: PDL 사용 시 데이터 로딩 및 연산 간의 동기화 문제가 지적되었습니다. 특히 A 행렬과 스케일 팩터 로딩이 적절히 동기화되지 않아 발생할 수 있는 잠재적 레이스 컨디션(race condition)에 대한 논의가 있었습니다. PR에서는 이 부분에 대한 수정이 이루어졌으나, 리뷰어는 A 행렬의 사전 로딩 가능성 및 SFA(Scale Factor A) 로딩/알파 값 읽기 순서에 대한 추가적인 논의를 제기했습니다. 이는 병렬 처리 환경에서 데이터 종속성 관리가 얼마나 복잡하고 중요한지를 보여줍니다.
일반적인 교훈:
- 저정밀도 연산의 중요성: 최신 GPU 아키텍처는 저정밀도 연산(FP8, FP4 등)을 통해 성능을 극대화하도록 설계되었습니다. 이러한 기능을 라이브러리에 통합하는 것은 성능 향상의 핵심입니다.
- DSL 활용의 장단점: CuTe DSL과 같은 DSL은 하드웨어 특화된 고성능 커널을 생성하는 강력한 도구이지만, 해당 DSL의 버전 관리, 제약 조건 검증, 그리고 테스트 커버리지 확보가 매우 중요합니다.
- 철저한 테스트 및 검증: 새로운 기능을 추가할 때는 기존 기능과의 호환성, 다양한 엣지 케이스, 그리고 성능뿐만 아니라 정확성까지 보장하기 위한 포괄적인 테스트가 필수적입니다. 특히, 테스트의 허용 오차(tolerance) 설정은 실제 연산 결과의 범위를 고려해야 합니다.
- 리뷰 프로세스의 가치: 코드 리뷰는 단순한 버그 수정을 넘어, 잠재적인 성능 문제, 라이브러리 호환성, API 디자인 개선 등 심도 있는 기술적 논의를 이끌어내고 코드 품질을 향상시키는 데 결정적인 역할을 합니다.
결론
이번 FlashInfer PR은 CuTe DSL을 활용하여 Blackwell 아키텍처에 최적화된 저지연 GEMM 커널을 성공적으로 도입했습니다. FP4 및 FP8 연산 지원 확대, 자동 튜닝 기능 강화, 그리고 철저한 테스트를 통해 FlashInfer는 고성능 대규모 언어 모델 추론을 위한 핵심 라이브러리로서의 입지를 더욱 공고히 했습니다. 이 PR은 최신 하드웨어 기능을 활용하여 AI 연산 성능을 극한까지 끌어올리는 엔지니어링의 좋은 사례를 보여줍니다.
참고 자료
- https://github.com/NVIDIA/cutlass/blob/main/README.md
- https://github.com/flashinfer-ai/flashinfer/blob/main/docs/source/api/gemm.rst
- https://github.com/flashinfer-ai/flashinfer/blob/main/flashinfer/gemm/gemm_base.py
- https://github.com/flashinfer-ai/flashinfer/blob/main/flashinfer/gemm/kernels/cute_dsl/low_latency_blockscaled_gemm.py
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer, SM100 아키텍처를 위한 BF16 x FP4 GEMM 최적화로 성능 극대화
- [flashinfer] FlashInfer SM120 MoE GEMM 최적화: 웨이브+잔여물 비용 모델 도입
- [flashinfer] FlashInfer, MoE 및 FP8 GEMM 성능 향상을 위한 커널 업데이트
- [sglang] SGLang: MiniMax-M2.5 MoE 모델을 위한 FP8 FlashInfer TRT-LLM 라우팅 최적화
- [flashinfer] FlashInfer SM12x MoE 최적화: 정적 MoE 경로 통합 및 성능 향상
PR Analysis 의 다른글
- 이전글 [onnxruntime] ONNX Runtime CUDA: int64 CumSum 연산 9배 가속화 최적화 분석
- 현재글 : [flashinfer] FlashInfer, CuTe DSL을 활용한 저지연 GEMM 커널 도입으로 성능 극대화
- 다음글 [vllm] vLLM의 작은 배치 사이즈를 위한 Triton 기반 Split-row Top-p 샘플링 최적화
댓글