[flashinfer] FlashInfer SM103 FP4 GEMM 최적화: Store256 및 Fused Epilogue 도입
PR 링크: flashinfer-ai/flashinfer#4063 상태: Merged | 변경: +602 / -50
들어가며
최신 GPU 아키텍처인 SM103(Blackwell) 환경에서 FP4 GEMM 연산은 추론 성능의 핵심입니다. 기존 FlashInfer의 구현은 범용적인 TMA(Tensor Memory Accelerator) 기반 스토어 경로를 사용하고 있었으나, 이는 특정 상황에서 메모리 대역폭을 완전히 활용하지 못하는 제약이 있었습니다. 본 PR은 SM103 아키텍처를 타겟으로 Store256 최적화와 mul.f32x2 기반의 Fused Epilogue를 도입하여 연산 효율을 극대화했습니다.
코드 분석
1. Store256 최적화 (include/flashinfer/gemm/fp4_gemm_cutlass_template_sm103.h)
기존에는 128-bit 정렬에 의존하던 스토어 방식을 256-bit로 확장했습니다. 이를 위해 isStore256OutputAligned 함수를 도입하여 런타임에 출력 텐서의 정렬 상태를 확인하고, 정렬이 보장될 경우 최적화된 run_store256 경로를 선택하도록 디스패처를 수정했습니다.
template <typename T>
bool isStore256OutputAligned(T const* D, int n) {
constexpr uintptr_t kStoreAlignmentBytes = 32;
return D != nullptr && reinterpret_cast<uintptr_t>(D) % kStoreAlignmentBytes == 0 &&
static_cast<uintptr_t>(n) * sizeof(T) % kStoreAlignmentBytes == 0;
}
2. JIT 컴파일 파이프라인 확장 (flashinfer/jit/gemm/core.py)
새로운 Store256 커널을 JIT 컴파일 과정에 포함하기 위해 generic_templates 리스트를 확장했습니다. 이를 통해 기존 커널과 최적화된 커널을 선택적으로 생성할 수 있는 구조를 갖추었습니다.
generic_templates = [
("fp4_gemm_cutlass.jinja", ""),
("fp4_gemm_cutlass_sm103_generic_store256.jinja", "_store256"),
]
왜 이게 좋은가
이번 최적화의 핵심은 메모리 대역폭 활용률의 극대화입니다.
- Store256: 128-bit에서 256-bit로 스토어 단위를 키움으로써, 글로벌 메모리로의 쓰기 트랜잭션 횟수를 줄여 대역폭 병목을 완화했습니다. 결과적으로 Native K768 전술에서 약 1.16배의 성능 향상을 달성했습니다.
- Fused F32x2:
mul.f32x2PTX 명령어를 사용하여alpha * accumulator연산을 쌍으로 처리함으로써 연산 파이프라인의 효율을 높였습니다. 특히.rnd옵션을 생략하여 기본 반올림 모드(Round-to-Nearest-Even)를 사용함으로써 컴파일러가 더 공격적인 최적화를 수행할 수 있도록 했습니다.
일반적 교훈
- 정렬(Alignment)의 중요성: GPU 메모리 연산에서 데이터 정렬은 하드웨어의 최대 대역폭을 끌어내기 위한 필수 조건입니다. 런타임 체크를 통해 정렬을 보장하고 최적화된 경로를 타는 패턴은 고성능 커널 설계의 정석입니다.
- JIT 템플릿화: 하드웨어 특화 커널을 JIT로 관리하면 유지보수 비용을 줄이면서도 아키텍처별 최적화 경로를 유연하게 추가할 수 있습니다.
리뷰어 피드백 반영
리뷰 과정에서 tiffany940107은 PTX 명령어 최적화와 코드 중복 제거에 대해 중요한 피드백을 주었습니다. 특히 LinearCombinationF32x2Impl을 별도로 추출하여 코드 재사용성을 높인 점은 유지보수 측면에서 매우 훌륭한 개선입니다. 또한, 테스트 코드를 내부 모듈 직접 호출 방식에서 flashinfer.mm_fp4 공개 API 사용 방식으로 변경하여 API 일관성을 확보했습니다.
참고 자료
- https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#floating-point-instructions-mul
- https://github.com/NVIDIA/cutlass
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] [FlashInfer] CUTLASS MoE 커널 최적화: 벡터화와 동적 스레드 할당으로 성능 한계 돌파하기
- [flashinfer] FlashInfer의 FP4 GEMM 최적화: 휴리스틱 개선과 Autotuning 효율화
- [cutlass] NVIDIA CUTLASS CuTeDSL: SM103 Grouped Block-Scaled GEMM 최적화 분석
- [flashinfer] FlashInfer의 GDN 커널 런칭 오버헤드 80% 절감하기: 호스트 측 최적화 전략
- [flashinfer] FlashInfer: SM120/SM121 아키텍처를 위한 네이티브 MXFP4 W4A4 Fused MoE 지원
PR Analysis 의 다른글
- 이전글 [sglang] SGLang Mamba 캐시의 상태 손상 및 슬롯 누수 버그 수정
- 현재글 : [flashinfer] FlashInfer SM103 FP4 GEMM 최적화: Store256 및 Fused Epilogue 도입
- 다음글 [onnxruntime] [ONNX Runtime] CPU GQA 성능의 한계를 넘다: FP16 입력과 양자화된 KV 캐시 최적화 분석
댓글