본문으로 건너뛰기

[flashinfer] FlashInfer, SM100 아키텍처를 위한 BF16 x FP4 GEMM 최적화로 성능 극대화

PR 링크: flashinfer-ai/flashinfer#4466 상태: Merged | 변경: +3193 / -27

들어가며

최근 대규모 언어 모델(LLM)의 발전과 함께, 모델의 추론 속도는 서비스의 품질을 결정하는 핵심 요소가 되었습니다. 특히 행렬 곱셈(GEMM) 연산은 딥러닝 모델의 연산량 대부분을 차지하기 때문에, 이 부분의 최적화는 전체 추론 성능 향상에 지대한 영향을 미칩니다. NVIDIA의 최신 GPU 아키텍처들은 저정밀도 연산(예: FP4, INT4)을 지원하며, 이를 활용하여 메모리 대역폭과 연산 효율성을 높이는 것이 중요해졌습니다.

이번 PR은 FlashInfer 라이브러리에 NVIDIA의 SM100 (Ampere 아키텍처) GPU를 위한 BF16 (BFloat16) 행렬과 FP4 (4-bit Floating Point) 행렬 간의 곱셈(GEMM) 연산을 최적화하는 새로운 커널을 추가합니다. 이는 특히 FP4 양자화된 가중치를 사용하는 모델의 추론 성능을 크게 향상시킬 것으로 기대됩니다.

코드 변경사항 분석

이번 PR의 핵심은 flashinfer/gemm/gemm_bf16_fp4_cute_dsl.py 파일에 새로운 아키텍처별 커널 로직을 추가하고, 기존 로직을 개선하는 것입니다. 특히 SM100 아키텍처에 특화된 Sm100DenseGemmBf16Fp4Kernel이 도입되었습니다.

1. flashinfer/gemm/gemm_bf16_fp4.py - 요구사항 검증 로직 개선

이 파일에서는 CuTe-DSL 백엔드를 사용하기 위한 입력 텐서의 데이터 타입 요구사항을 검증하는 로직이 수정되었습니다. 이전에는 무조건 torch.int32를 기대했지만, 이제는 GPU의 컴퓨팅 캐피빌리티(Compute Capability)에 따라 기대하는 데이터 타입이 달라지도록 개선되었습니다.

Before:

    if b.dtype != torch.int32:
        raise ValueError(
            f"cute-dsl bf16 x fp4 expects the int32 tile-packed weight from "
            f"prepare_bf16_fp4_weights(..., backend='cute-dsl'); got {b.dtype}."
        )

After:

    major, minor = get_compute_capability(a.device)
    cc = major * 10 + minor
    expected_dtype = torch.uint8 if cc in (100, 103) else torch.int32
    if b.dtype != expected_dtype:
        raise ValueError(
            f"cute-dsl bf16 x fp4 on SM{cc} expects the {expected_dtype} weight "
            "returned by prepare_bf16_fp4_weights(..., backend='cute-dsl'); "
            f"got {b.dtype}."
        )

왜 좋은가?

  • 아키텍처 특화 최적화: SM100 (CC 10.0) 및 SM103 (CC 10.3) 아키텍처에서는 FP4 가중치를 torch.uint8 형태로 저장하고 처리하는 것이 더 효율적입니다. 이 변경은 해당 아키텍처에 맞는 최적의 데이터 타입을 사용하도록 하여 성능을 향상시킵니다. 이전에는 모든 아키텍처에 대해 torch.int32를 강제하여 비효율이 발생할 수 있었습니다.
  • 명확한 오류 메시지: GPU 아키텍처에 따라 기대하는 데이터 타입이 달라지므로, 오류 메시지도 해당 아키텍처 정보를 포함하여 더 명확해졌습니다. 이는 개발자가 문제를 더 쉽게 진단하고 해결하는 데 도움을 줍니다.
  • get_compute_capability 활용: flashinfer.utils.get_compute_capability 함수를 사용하여 현재 실행 중인 GPU의 컴퓨팅 캐피빌리티를 동적으로 가져옵니다. 이는 라이브러리가 다양한 NVIDIA GPU에서 올바르게 작동하도록 보장합니다.

2. flashinfer/gemm/gemm_bf16_fp4_cute_dsl.py - SM100 커널 도입 및 준비 로직 분리

이 파일은 CuTe-DSL 백엔드를 사용하여 BF16 x FP4 GEMM을 구현하는 핵심 로직을 담고 있습니다. 이번 PR에서는 다음과 같은 주요 변경이 이루어졌습니다.

  • 커널 이름 변경 및 SM100 커널 도입: 기존의 BlackwellDenseGemmBf16Fp4KernelSm12xDenseGemmBf16Fp4Kernel로 이름이 변경되었고, 새로운 Sm100DenseGemmBf16Fp4Kernel이 추가되었습니다. 이는 각 아키텍처에 맞는 최적화된 커널을 사용하기 위함입니다.
  • 아키텍처별 준비 함수 분리: 가중치(b), 스케일 팩터(b_descale), 알파(alpha) 등의 입력을 커널이 요구하는 형태로 준비하는 로직이 아키텍처별로 분리되었습니다. _prepare_cute_dsl_sm100 함수가 새로 추가되었고, 기존의 _prepare_cute_dsl 함수는 이를 호출하도록 변경되었습니다.

Before (일부, _get_cute_dsl_bf16_fp4_gemm 함수 내):

    from .kernels.cute_dsl.dense_gemm_bf16_fp4_blackwell import \
        BlackwellDenseGemmBf16Fp4Kernel

    # ... (중략) ...

    gemm = BlackwellDenseGemmBf16Fp4Kernel(
        acc_dtype=cutlass.Float32,
        tile_shape_mnk=tile_shape_mnk,
        atom_layout=atom_layout,
        # ... (중략) ...
    )

After (일부, _get_cute_dsl_bf16_fp4_gemm 함수 내):

    from .kernels.cute_dsl.dense_gemm_bf16_fp4_sm12x import \
        Sm12xDenseGemmBf16Fp4Kernel

    # ... (중략) ...

    gemm = Sm12xDenseGemmBf16Fp4Kernel(
        acc_dtype=cutlass.Float32,
        tile_shape_mnk=tile_shape_mnk,
        atom_layout=atom_layout,
        # ... (중략) ...
    )

After (새로운 준비 함수 및 디스패치 로직):

def _prepare_cute_dsl_sm100(
    b: torch.Tensor,
    b_descale: torch.Tensor,
    alpha: Optional[torch.Tensor],
    block_size: int,
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
    """cute-DSL-backend prep: keep SF in swizzled layout.

    ``convert_sf_to_mma_layout`` only creates the six-dimensional strided view
    consumed by the SM100 TMA descriptor; it does not copy or linearize the
    canonical 128x4 scale-factor buffer.
    """
    from ..cute_dsl.utils import convert_sf_to_mma_layout

    n = int(b.shape[0])
    k = int(b.shape[1]) * 2
    b_descale = b_descale.contiguous()
    weight_sf = convert_sf_to_mma_layout(
        b_descale,
        m=n,
        k=k,
        num_groups=1,
        sf_vec_size=block_size,
    )
    return b.contiguous(), weight_sf, alpha

def _prepare_cute_dsl(
    b: torch.Tensor,
    b_descale: torch.Tensor,
    alpha: Optional[torch.Tensor],
    block_size: int,
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
    """Dispatch weight preparation to the architecture-specific DSL kernel."""
    major, minor = get_compute_capability(b.device)
    if (major, minor) in ((10, 0), (10, 3)):
        return _prepare_cute_dsl_sm100(b, b_descale, alpha, block_size)
    elif major == 12:
        return _prepare_cute_dsl_sm12x(b, b_descale, alpha, block_size)
    else:
        raise NotImplementedError(
            f"cute-dsl w4a16 GEMM only supports SM100/103 and SM12x; got {major}.{minor}"
        )

왜 좋은가?

  • 아키텍처별 최적화: SM100 아키텍처는 FP4 가중치와 스케일 팩터를 처리하는 방식이 SM12x (Ampere 이전)와 다릅니다. SM100은 TMA(Tensor Memory Accelerator)를 활용하여 스케일 팩터(SF)를 특정 레이아웃으로 직접 로드할 수 있습니다. _prepare_cute_dsl_sm100 함수는 convert_sf_to_mma_layout를 사용하여 이 아키텍처에 최적화된 스케일 팩터 레이아웃을 생성합니다. 이는 데이터 이동을 줄이고 GPU 코어에 데이터를 더 효율적으로 공급하여 성능을 향상시킵니다.
  • 코드 구조 개선: 아키텍처별 로직을 명확하게 분리함으로써 코드의 가독성과 유지보수성이 향상되었습니다. 각 함수는 특정 아키텍처에 대한 최적화에 집중할 수 있습니다.
  • Sm100DenseGemmBf16Fp4Kernel 도입: 이 새로운 커널은 SM100의 하드웨어 기능을 최대한 활용하도록 설계되었습니다. 예를 들어, FP4 가중치 로딩, 스케일 팩터 적용, BF16 행렬 곱셈 및 FP32 누적 등 각 단계가 최적화되었을 가능성이 높습니다.

3. SM100 전용 커널 및 튜닝 로직 추가

PR에는 SM100 아키텍처를 위한 GEMM 커널 컴파일 및 실행을 위한 세부 로직이 추가되었습니다. 이는 _get_sm100_bf16_fp4_kernel_launch_cute_dsl_sm100 함수에 구현되어 있습니다.

주요 코드 예시 (_get_sm100_bf16_fp4_kernel):

    from .kernels.cute_dsl.dense_gemm_bf16_fp4_sm100 import \
        Sm100DenseGemmBf16Fp4Kernel

    # ... (중략) ...

    kernel = Sm100DenseGemmBf16Fp4Kernel(
        acc_dtype=cutlass.Float32,
        use_2cta_instrs=use_2cta_instrs,
        mma_tiler_mnk=mma_tiler_mnk,
        cluster_shape_mn=cluster_shape_mn,
        enable_pdl=enable_pdl,
        raster_along_m=raster_along_m,
        transform_fragment_size=transform_fragment_size,
    )
    compiled = cute.compile(
        kernel.wrapper,
        weight_ptr,
        weight_sf_ptr,
        activation_ptr,
        alpha_ptr,
        output_ptr,
        n,
        m,
        k,
        max_active_clusters=max_active_clusters,
        stream=stream,
        options="--opt-level 2 --enable-tvm-ffi",
    )
    _SM100_BF16_FP4_KERNEL_CACHE[cache_key] = compiled
    return compiled

왜 좋은가?

  • CuTe-DSL 및 Cutlass 활용: NVIDIA의 CuTe-DSL과 Cutlass 라이브러리를 활용하여 GPU 하드웨어에 최적화된 커널을 생성합니다. cute.compile 함수는 런타임에 커널을 컴파일하여 특정 하드웨어 및 입력 크기에 최적화된 코드를 생성합니다. 이는 높은 성능을 달성하는 데 필수적입니다.
  • 커널 캐싱: _SM100_BF16_FP4_KERNEL_CACHE를 사용하여 컴파일된 커널을 캐싱합니다. 동일한 튜닝 파라미터(tactic)로 커널이 다시 요청될 경우, 재컴파일 없이 캐시된 커널을 재사용하여 컴파일 시간을 절약하고 런타임 오버헤드를 줄입니다.
  • 튜닝 파라미터 탐색: _SM100_BF16_FP4_TACTICS 튜플은 다양한 타일 크기(mma_tiler_mnk), 클러스터 모양(cluster_shape_mn), 래스터화 방향(raster_along_m) 등 잠재적으로 성능이 좋은 튜닝 파라미터 조합을 정의합니다. 런타임 시 이러한 파라미터들을 탐색하여 최적의 성능을 내는 커널을 선택하거나 컴파일합니다.
  • 메모리 정렬 요구사항: SM100 커널은 특정 데이터 포인터 정렬을 요구합니다 (weight, b_descale, a, alpha, out 텐서). 이는 GPU 하드웨어가 데이터를 효율적으로 로드하기 위한 필수 조건입니다. _launch_cute_dsl_sm100 함수는 이러한 정렬 요구사항을 검증하고, 필요시 사용자에게 경고를 제공합니다.

왜 이게 좋은가?

이번 PR은 다음과 같은 이유로 매우 훌륭한 최적화 및 개선이라고 할 수 있습니다.

  1. 획기적인 성능 향상: PR 설명에 제시된 성능 벤치마크 결과는 매우 인상적입니다. 예를 들어, N=6656, K=19968 조건에서 M=128일 때 FlashInfer cute-dsl은 Marlin NVFP4 대비 2.41배, torch-bf16 대비 1.77배 빠른 성능을 보여줍니다. 다른 조건에서도 유사하거나 더 큰 성능 향상이 관찰됩니다. 이는 FP4 양자화된 가중치를 사용하는 모델의 추론 속도를 크게 단축시킬 수 있음을 의미합니다. 특히 Marlin NVFP4 대비 상당한 우위를 보이는 경우가 많습니다. 이는 FlashInfer가 자체 개발한 CuTe-DSL 커널이 해당 아키텍처에 얼마나 잘 최적화되었는지를 보여줍니다.

    주요 성능 지표 (N=6656, K=19968, M=128):

    • torch-bf16: 59.76 µs (0.97×)
    • Marlin MXFP8: 120.26 µs (0.88×)
    • Marlin NVFP4: 106.01 µs (1.00×)
    • FlashInfer cute-dsl: 43.90 µs (2.41×)
  2. 아키텍처 특화 최적화: 최신 NVIDIA GPU 아키텍처(특히 SM100/103)의 하드웨어 기능을 최대한 활용하도록 커널을 설계했습니다. FP4 데이터 타입 처리, TMA(Tensor Memory Accelerator) 활용, 최적의 메모리 레이아웃 등을 고려하여 성능을 극대화했습니다. 이는 범용적인 커널보다 훨씬 높은 효율성을 제공합니다.

  3. 저정밀도 연산 지원 강화: FP4와 같은 초저정밀도 데이터 타입을 효과적으로 지원함으로써, 모델의 메모리 사용량을 줄이고 추론 속도를 높일 수 있습니다. 이는 더 큰 모델을 더 적은 리소스로 실행하거나, 동일한 하드웨어에서 더 빠른 응답 시간을 제공하는 데 기여합니다.

  4. 코드 품질 및 유지보수성 향상: 아키텍처별 로직 분리, 명확한 요구사항 검증, 커널 캐싱 메커니즘 도입 등을 통해 코드의 구조가 개선되었습니다. 이는 향후 새로운 아키텍처 지원 추가나 기존 커널 수정 시 유지보수성을 높여줍니다.

  5. 일반적인 교훈:

    • 하드웨어 이해의 중요성: GPU 아키텍처의 특성(예: 데이터 타입 지원, 메모리 접근 방식, TMA 등)을 깊이 이해하고 이를 커널 설계에 반영하는 것이 고성능 컴퓨팅의 핵심입니다.
    • DSL 및 컴파일러 활용: CuTe-DSL과 같은 도메인 특화 언어(DSL)와 런타임 컴파일러는 하드웨어에 최적화된 코드를 생성하는 강력한 도구입니다. 이를 잘 활용하면 복잡한 하드웨어 최적화를 추상화할 수 있습니다.
    • 튜닝 및 캐싱: 다양한 튜닝 파라미터를 탐색하고, 컴파일된 커널을 캐싱하는 전략은 성능과 런타임 효율성을 동시에 잡는 데 중요합니다.
    • 점진적 개선: 기존 커널의 요구사항 검증 로직을 개선하고, 새로운 아키텍처를 위한 커널을 점진적으로 추가하는 방식은 라이브러리의 안정성과 성능을 동시에 향상시킵니다.

리뷰어 노트 분석

리뷰 댓글을 보면, 주로 CI 파이프라인 실행 및 테스트 결과에 대한 내용이었습니다. 특히 tests.gemm.test_bmm_fp8.py에서 발생하는 실패는 이 PR의 직접적인 변경 사항과는 관련이 없는 것으로 보이며, 기존 테스트 환경이나 다른 변경 사항과의 충돌로 인한 문제일 가능성이 높습니다. IwakuraRein님의 댓글에서 해당 테스트가 PR 범위와 무관함을 명확히 하여, 핵심 코드 변경에 대한 기술적인 논의는 다른 부분에서 이루어졌음을 시사합니다. 이 PR의 핵심인 SM100 아키텍처용 BF16 x FP4 GEMM 커널 추가 및 최적화 자체에 대한 부정적인 피드백은 없었습니다.

References

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글