본문으로 건너뛰기

[flashinfer] Blackwell 아키텍처를 위한 극한의 최적화: MiniMax-H3 Fused 커널의 256-Column Tiling 전략

PR 링크: flashinfer-ai/flashinfer#5506 상태: Merged | 변경: +713 / -617

들어가며

최근 LLM 및 Diffusion 모델의 추론 속도를 높이기 위해 연산자 퓨전(Operator Fusion)과 양자화(Quantization)는 필수적인 기법이 되었습니다. 특히 NVIDIA의 최신 아키텍처인 Blackwell(SM100, SM103)은 TMEM(Tensor Memory)과 TMA(Tensor Memory Accelerator)라는 강력한 하드웨어 기능을 제공하여, 기존 Hopper 아키텍처보다 더 높은 연산 밀도를 가능하게 합니다.

이번 글에서는 flashinfer 라이브러리에 적용된 MiniMax-H3 fused RMSNorm + AdaLN + FC1 + SwiGLU 커널의 최적화 PR을 분석합니다. 이 PR은 기존의 Tiling 전략을 전면 수정하여 Blackwell 하드웨어의 잠재력을 극한으로 끌어올렸으며, 결과적으로 이전 구현 대비 최대 1.32배, 일반적인 segmented chain 대비 최대 3.05배의 성능 향상을 달성했습니다.


코드 분석: 핵심 변경 사항

1. Tiling 전략의 변화: 224에서 256으로

가장 큰 변화는 GEMM의 출력 타일(Output Tile) 크기를 확장한 것입니다. 기존 224-column 방식에서 256-column 방식으로 변경하여 하드웨어 정렬(Alignment)과 연산 효율을 높였습니다.

Before:

#define BLOCK_N 224
#define B_HALF_N 112
#define N_TILES 128

After:

#define BLOCK_N 256
#define B_HALF_N 128
#define N_TILES 112

BLOCK_N이 256으로 증가하면서, CTA(Cooperative Thread Array) 쌍당 처리하는 데이터 양이 늘어났습니다. 이는 Blackwell의 텐서 코어가 256 단위의 연산에 더 최적화되어 있음을 시사합니다.

2. 스레드 수 및 공유 메모리(Shared Memory) 최적화

더 큰 타일을 처리하기 위해 커널의 스레드 수와 공유 메모리 할당량이 조정되었습니다.

Before:

#define SM_TOTAL 194688
__global__ __launch_bounds__(192) __cluster_dims__(2,1,1) void
kernel_minimax_h3_fc1_swiglu_e4m3(...)

After:

#define SM_TOTAL 206976
__global__ __launch_bounds__(320) __cluster_dims__(2,1,1) void
kernel_minimax_h3_fc1_swiglu_e4m3(...)

스레드 수가 192개에서 320개로 대폭 증가했습니다. 이는 더 복잡해진 에필로그(Epilogue) 연산을 병렬화하고, 메모리 레이턴시를 숨기기(Latency Hiding) 위한 결정입니다. 또한 SMEM_TOTAL이 약 207KB로 증가하며 Blackwell의 넉넉한 공유 메모리 자원을 적극적으로 활용하고 있습니다.

3. TMEM 및 TMA 활용의 고도화

Blackwell의 핵심인 TMEM(Tensor Memory) 할당 방식이 변경되었습니다. TMEM_NCOLS가 줄어든 것처럼 보이지만, 이는 256-column accumulator를 하나로 통합하여 더 효율적으로 관리하기 위함입니다.

Before:

#define TMEM_NCOLS 472
#define TMEM_TMEM_SFA_OFFSET 448
#define TMEM_TMEM_SFB_OFFSET 456

After:

#define TMEM_NCOLS 280
#define TMEM_TMEM_SFA_OFFSET 256
#define TMEM_TMEM_SFB_OFFSET 264

또한, Scale 타일을 가져올 때 TMA(Tensor Memory Accelerator)의 로드 단위를 16-byte에서 128-byte로 변경했습니다. 이는 메모리 버스 대역폭을 더 효율적으로 사용하여 데이터 로딩 병목을 줄이는 결정적인 역할을 합니다.

4. 인라인 어셈블리를 통한 TMEM 로드 최적화

이번 PR에서는 TMEM에서 데이터를 효율적으로 읽어오기 위해 tcgen05 명령어를 사용하는 인라인 어셈블리 함수가 추가되었습니다.

__device__ __forceinline__ void tmem_ld_x32(float* dst, int tmem_addr) {
    asm volatile(
        "tcgen05.ld.sync.aligned.32x32b.x32.b32"
        " {%0, %1, %2, %3, %4, %5, %6, %7," // ... 중략 ...
        "  %24, %25, %26, %27, %28, %29, %30, %31}, [%32];"
        : "=f"(dst[0]), // ... 중략 ...
        : "r"(tmem_addr));
}

이 함수는 32x32비트 데이터를 한 번에 동기적으로 로드하여, FP4/MXFP8 양자화 연산 이후의 에필로그 처리를 가속화합니다.


왜 이게 좋은가?

1. 하드웨어 한계(Ceiling)에 근접한 성능

PR 설명에 따르면, 최적화된 NVFP4 GEMM 커널은 Operand 로드 시간을 제외한 순수 연산에서 B200 기준 피크 성능의 85-90%, B300 기준 92-94%에 도달했습니다. 이는 소프트웨어 수준에서 할 수 있는 최적화가 거의 정점에 도달했음을 의미합니다.

2. 메모리 효율성 극대화

Scale 타일을 128-byte TMA row로 읽어오는 방식은 메모리 트랜잭션 횟수를 줄여 L2 캐시와 SM 간의 데이터 전송 효율을 높였습니다. GEMM 성능의 병목이 주로 L2 -> SM 데이터 전달에 있다는 점을 파악하고 이를 정밀 타격한 결과입니다.

3. 실질적인 가속 수치

  • B200 (SM100a): 기존 대비 1.22~1.32배 향상
  • B300 (SM103a): 기존 대비 1.08~1.31배 향상

이러한 수치는 단순한 이론적 향상이 아니라, 실제 프로덕션 환경에서 사용되는 다양한 Shape(M=4184~109952)에서 검증된 결과입니다.


결론

이번 최적화는 단순히 코드를 정리하는 수준을 넘어, Blackwell 아키텍처의 하드웨어 특성(TMA, TMEM, 텐서 코어 레이아웃)을 깊이 이해하고 이를 소프트웨어 설계에 반영한 훌륭한 사례입니다.

특히 256-column Tiling으로의 전환과 TMA 로드 크기 최적화는 향후 Blackwell 기반의 다른 커널 최적화에서도 표준적인 접근 방식으로 참고될 가치가 충분합니다. 고성능 컴퓨팅(HPC) 분야의 엔지니어라면, 하드웨어의 물리적 한계(Ceiling)를 측정하고 그 간극을 좁혀나가는 이 PR의 분석 과정을 눈여겨볼 필요가 있습니다.


References

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글