[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)
- Persistent Tiles & TMA Ring: 128x128 출력 타일을 유지하면서 64바이트 로우 TMA 링 버퍼를 사용해 데이터를 효율적으로 로드합니다.
- Register-level SwiGLU: GEMM의 결과가 레지스터(Accumulator fragments)에 남아있는 상태에서 즉시 SwiGLU 연산을 수행합니다. HBM으로의 중간 쓰기 과정을 완전히 제거했습니다.
- Warp Specialization: 8개의 MMA 워프를 사용하며, Warp 0이 TMA Producer 역할을 겸임하도록 설계하여 레지스터 오버헤드를 최소화했습니다 (9번째 워프 사용 시 레지스터 부족으로 인한 Spill 발생 방지).
- 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)을 이해하고 이를 활용해 병목을 제거하는 것이 얼마나 큰 성능 차이를 만드는지 이 사례를 통해 배울 수 있습니다.
참고 문헌
- torch.nn.functional.rms_norm — 커널에서 융합된 RMSNorm의 기준 함수
- torch.cuda.Event — 벤치마크에서 정밀한 시간 측정을 위해 사용된 API
- NVIDIA Blackwell Architecture — SM120 아키텍처 및 NVFP4 지원에 대한 공식 정보
참고 자료
- https://pytorch.org/docs/stable/generated/torch.nn.functional.rms_norm.html
- https://pytorch.org/docs/stable/generated/torch.cuda.Event.html
- https://www.nvidia.com/en-us/data-center/blackwell-architecture/
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] Blackwell 아키텍처를 위한 MoE All-Reduce Fusion 최적화: FlashInfer의 'Cake' 백엔드 분석
- [flashinfer] FlashInfer, Blackwell 아키텍처를 위한 Recurrent KDA Prefill 최적화: Small-BH 커널 도입
- [flashinfer] Blackwell NVFP4 양자화 최적화: TMA OOB Zero-fill을 이용한 메모리 복사 오버헤드 제거
- [flashinfer] FlashInfer의 Fused SwiGLU 및 NVFP4 양자화 최적화 분석
- [flashinfer] FlashInfer의 Per-token NVFP4 Quantization 커널 최적화 분석
PR Analysis 의 다른글
- 이전글 [onnxruntime] ONNX Runtime, x86 CPU에서 FP16 LayerNorm 및 RMSNorm 성능 최적화: AVX2 활용
- 현재글 : [flashinfer] NVIDIA Blackwell(SM120)을 위한 초고속 커널 최적화: MiniMax-H3 Fused FC1 + SwiGLU 분석
- 다음글 [flashinfer] FlashInfer KDA: BF16 준비된 Prefill 계획 캐싱 및 FP32 중간 상태 최적화 분석
댓글