본문으로 건너뛰기

[flashinfer] NVIDIA Blackwell(SM120)을 위한 초고속 커널 최적화: MiniMax-H3 Fused FC1 + SwiGLU 분석

PR 링크: flashinfer-ai/flashinfer#5521 상태: Merged | 변경: +9225 / -0

들어가며

최근 생성형 AI 모델의 크기가 급격히 커짐에 따라, 추론 성능을 확보하기 위한 하드웨어 가속과 소프트웨어 최적화의 중요성이 그 어느 때보다 높아졌습니다. 특히 NVIDIA의 최신 아키텍처인 Blackwell(SM120)은 FP8뿐만 아니라 새로운 NVFP4(4-bit Floating Point) 데이터 타입을 지원하며 연산 성능의 새로운 지평을 열었습니다.

이번 글에서는 flashinfer-ai/flashinfer 레포지토리에 반영된 최신 PR을 통해, MiniMax-H3 확산 모델(Diffusion Model)의 핵심 레이어인 FC1 + SwiGLU를 Blackwell 아키텍처에 최적화하여 융합(Fusion)한 사례를 분석합니다. 이 PR은 기존의 분절된 연산 체인을 하나의 커널로 통합하여 메모리 대역폭 병목을 해결하고, RTX 5090 및 RTX PRO 6000 환경에서 압도적인 성능 향상을 이끌어냈습니다.


기존 방식의 문제점: Segmented Chain의 한계

일반적으로 PyTorch나 기존 라이브러리에서 RMSNorm, AdaLN, Quantization, GEMM, SwiGLU 연산을 수행할 때는 각 단계마다 별도의 커널을 실행합니다. 이를 Segmented Chain 방식이라고 합니다.

Before: Segmented Chain (Baseline)

벤치마크 코드(benchmarks/bench_minimax_h3_sm120_quant_fc1_swiglu.py)에서 볼 수 있듯이, 기존 방식은 다음과 같이 여러 단계의 함수 호출로 구성됩니다.

# Baseline: 여러 커널이 순차적으로 호출됨
def baseline_nvfp4(inputs, model, weights, act_global_scale, alpha):
    w_q, w_sf = weights
    # 1. RMSNorm + AdaLN (CPU/GPU 오버헤드 발생)
    a = torch_pre_norm(inputs["x"], model, inputs["adaln_index"])
    
    # 2. Activation Quantization (메모리 쓰기 발생)
    a_q, a_sf = fp4_quantize(
        a, act_global_scale, sf_vec_size=MINIMAX_H3_SF_BLOCK, ...
    )
    
    # 3. GEMM (FP4 Matrix Multiplication)
    h = mm_fp4(a_q, w_q.t(), a_sf, w_sf, alpha, torch.bfloat16, backend="cutlass")
    
    # 4. SwiGLU (SiLU and Multiply, 또 다른 메모리 읽기/쓰기)
    return silu_and_mul(h)

이 방식의 가장 큰 문제는 Intermediate Tensors(중간 텐서)입니다. 각 단계의 결과값이 GPU의 HBM(High Bandwidth Memory)에 기록되었다가 다음 커널에서 다시 읽혀야 하므로, 연산 속도보다 메모리 대역폭에 의해 전체 성능이 제한되는 Memory-bound 상황이 발생합니다.


핵심 최적화: Fused Kernel on SM120

이번 PR의 핵심은 이 모든 과정을 단 두 개의 커널 런칭으로 통합한 것입니다. 특히 Blackwell 아키텍처의 하드웨어 기능인 TMA(Tensor Memory Accelerator)와 MMA(Matrix-Multiply-Accumulate)를 적극 활용했습니다.

After: Fused MiniMax-H3 FC1 + SwiGLU

최적화된 코드는 단 한 줄의 API 호출로 모든 연산을 수행합니다.

# Fused: 모든 연산이 단일 커널 내에서 파이프라이닝됨
if args.variant == "fp8":
    fn = lambda: minimax_h3_fc1_swiglu_fp8(
        inputs["x"],               # Raw Input
        model["x_norm_weight"],    # RMSNorm Weight
        model["adaln_scale"],      # AdaLN Scale
        model["adaln_shift"],      # AdaLN Shift
        inputs["adaln_index"],     # Index for AdaLN
        fused_weights[0],          # Pre-processed FC1 Weight
        fused_weights[1]           # Weight Scale
    )

기술적 상세 구현 (Design)

  1. Persistent Tiles & TMA Ring: 128x128 출력 타일을 유지하면서 64바이트 로우 TMA 링 버퍼를 사용해 데이터를 효율적으로 로드합니다.
  2. Register-level SwiGLU: GEMM의 결과가 레지스터(Accumulator fragments)에 남아있는 상태에서 즉시 SwiGLU 연산을 수행합니다. HBM으로의 중간 쓰기 과정을 완전히 제거했습니다.
  3. Warp Specialization: 8개의 MMA 워프를 사용하며, Warp 0이 TMA Producer 역할을 겸임하도록 설계하여 레지스터 오버헤드를 최소화했습니다 (9번째 워프 사용 시 레지스터 부족으로 인한 Spill 발생 방지).
  4. NVFP4 지원: Blackwell의 핵심 기능인 mxf4nvf4.block_scale (m16n8k64) 명령어를 사용하여 4비트 양자화 연산을 가속화했습니다.

왜 이게 좋은가? (성능 및 효율성)

1. 압도적인 속도 향상 (Speedup)

PR 설명에 포함된 RTX PRO 6000 Blackwell 벤치마크 결과는 놀랍습니다.

  • FP8 Variant: 기존 Chain 대비 2.68x ~ 2.71x 속도 향상.
  • NVFP4 Variant: 기존 Chain 대비 1.52x ~ 1.53x 속도 향상.

특히 FP8에서 2.7배의 성능 향상이 나타난 이유는, 기존 torch._scaled_mm과 분절된 커널들이 Blackwell의 새로운 아키텍처 특성을 100% 활용하지 못했기 때문입니다. 전용 커널은 하드웨어의 최대 TFLOPS(FP8 기준 약 740 TFLOPS)에 근접하는 성능을 보여줍니다.

2. 메모리 사용량 절감

Variant M (Batch) Fused Peak GiB Chain Peak GiB 절감률
NVFP4 109,952 4.8 GiB 11.8 GiB ~60%
FP8 109,952 5.2 GiB 12.1 GiB ~57%

중간 텐서를 생성하지 않으므로 피크 메모리 사용량이 절반 이하로 줄어듭니다. 이는 더 큰 배치 사이즈를 처리하거나 더 긴 컨텍스트를 다룰 수 있게 해주는 결정적인 이점입니다.

3. 하드웨어 특화 최적화의 교훈

이 PR은 단순히 코드를 합친 것이 아니라, SM120 아키텍처의 레지스터 구조와 워프 스케줄링을 정밀하게 계산하여 작성되었습니다. 예를 들어, 워프 개수를 9개가 아닌 8개로 제한하여 레지스터 스필(Spill)을 막은 점은 시니어 엔지니어의 깊은 통찰력을 보여줍니다.


결론

이번 FlashInfer의 업데이트는 Blackwell 아키텍처를 사용하는 유저들에게 필수적인 최적화를 제공합니다. 특히 MiniMax-H3와 같은 최신 확산 모델에서 FP8/NVFP4 양자화를 실전 수준으로 끌어올렸다는 점에서 의의가 큽니다.

소프트웨어 엔지니어로서 우리는 단순히 라이브러리를 사용하는 것을 넘어, 하드웨어의 특성(TMA, MMA, Register file size)을 이해하고 이를 활용해 병목을 제거하는 것이 얼마나 큰 성능 차이를 만드는지 이 사례를 통해 배울 수 있습니다.


참고 문헌

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글