본문으로 건너뛰기

[flashinfer] FlashInfer: SM120/SM121 아키텍처를 위한 네이티브 MXFP4 W4A4 Fused MoE 지원

PR 링크: flashinfer-ai/flashinfer#4290 상태: Merged | 변경: +1160 / -184

들어가며

최신 대규모 언어 모델(LLM)의 핵심인 Mixture-of-Experts(MoE) 구조는 연산 효율성이 매우 중요합니다. 이번 FlashInfer 업데이트에서는 NVIDIA Blackwell(SM120/SM121) 아키텍처를 타겟으로, 기존 NVFP4를 넘어선 네이티브 MXFP4 W4A4 Fused MoE 지원을 추가했습니다. 이 변경사항은 데이터 양자화 효율을 극대화하여 연산 처리량을 높이고, 기존의 복잡한 파이프라인을 간소화하는 데 목적이 있습니다.

코드 분석

1. flashinfer/cute_dsl/fp4_common.py: MXFP4 양자화 로직 구현

MXFP4는 블록 단위 스케일링(Block-32)을 사용하여 정밀도를 유지하면서도 메모리 대역폭을 절약합니다. quantize_block_mxfp4 함수는 32개의 float32 값을 입력받아 UE8M0 스케일로 변환하고 패킹합니다.

Before (기존 방식):

# 기존 NVFP4 등 일반적인 FP4 양자화 로직
def quantize_block_fp4_fast(values, ...):
    # ... (기존 로직)

After (MXFP4 도입):

@cute.jit
def quantize_block_mxfp4(values: cute.Tensor, max_abs: Float32) -> Tuple[Uint64, Uint64, Uint8]:
    scale_u32 = cvt_f32_to_ue8m0(max_abs * rcp_approx_ftz(Float32(FLOAT4_E2M1_MAX)))
    scale_byte = Uint8(scale_u32 & Uint32(0xFF))
    # ... (32개 값에 대한 스케일링 및 패킹)
    return packed_lo, packed_hi, scale_byte

2. flashinfer/fused_moe/cute_dsl/b12x_moe.py: API 확장

B12xMoEWrapper에서 quant_mode="mxfp4"를 명시적으로 지원하도록 수정되었습니다. MXFP4는 자체적인 블록 스케일을 사용하므로, 기존의 input_global_scale 파라미터를 무시하도록 로직이 변경되었습니다.

# 변경된 파라미터 처리 로직
fc2_input_scale : Optional[torch.Tensor]
    # ...
    # ``quant_mode=

## 참고 자료
- nvfp4

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

댓글

관련 포스트

PR Analysis 의 다른글