[onnxruntime] ONNX Runtime: AVX2 및 AVX-VNNI를 위한 2-bit 가중치 CPU 커널 최적화
PR 링크: microsoft/onnxruntime#29619 상태: Merged | 변경: +1983 / -12
들어가며
최근 LLM 및 경량화 모델의 추론 성능을 극대화하기 위해 2-bit 양자화(W2) 기법이 널리 사용되고 있습니다. 하지만 기존 microsoft/onnxruntime의 MLAS(Microsoft Linear Algebra Subroutine) 라이브러리에서는 AVX-512를 지원하지 않는 일반적인 x86 CPU(Alder Lake, Zen 1-3 등) 환경에서 2-bit MatMul 연산 시, fp32 dequantization 후 SGEMM을 호출하는 비효율적인 폴백(fallback) 경로를 거쳐야 했습니다. 본 PR은 AVX2 및 AVX-VNNI를 지원하는 CPU를 위해 네이티브 2-bit 커널을 추가하여 이 병목 현상을 해결합니다.
코드 분석
1. CMake 빌드 설정 수정 (cmake/onnxruntime_mlas.cmake)
가장 중요한 변경 중 하나는 sqnbitgemm_kernel_avx512_2bit.cpp의 컴파일 플래그 제거입니다. 이 파일은 AVX2와 AVX-512 디스패치 테이블 모두에서 호출되는 공통 스칼라 헬퍼를 포함하고 있습니다.
Before:
set_source_files_properties(${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.cpp PROPERTIES
COMPILE_FLAGS "-mfma -mavx512bw -mavx512dq -mavx512vl")
After:
set_source_files_properties(${MLAS_SRC_DIR}/sqnbitgemm_kernel_avx512_2bit.cpp PROPERTIES
COMPILE_FLAGS "")
분석: AVX2 전용 호스트에서 AVX-512 관련 플래그가 활성화된 상태로 컴파일되면, 컴파일러가 스칼라 루프를 EVEX 명령어로 자동 벡터화(autovectorize)하여 SIGILL(Illegal Instruction) 오류가 발생합니다. 이를 방지하기 위해 플래그를 제거하여 베이스라인 코드를 생성하도록 변경했습니다.
2. AVX2 커널 디스패치 및 구현 (onnxruntime/core/mlas/lib/sqnbitgemm_kernel_avx2.cpp)
AVX2 및 AVX-VNNI 디스패치 테이블에 2-bit 네이티브 커널을 등록했습니다.
After (Dispatch):
d.SQ2BitGemmKernel_BlkSum_CompInt8 = onnxruntime::mlas::sq2bit_avx2::SQ2BitGemmKernel_BlkSum_CompInt8_Avx2_Dispatch;
분석: BlkLen에 따라 적절한 커널을 선택하는 라우팅 로직을 추가했습니다. BlkLen 32/64는 R2xC4 타일링을 사용하여 레지스터 효율을 높였고, 128은 레지스터 압박을 고려해 R1xC4 타일을 유지했습니다.
왜 이게 좋은가
이번 최적화는 단순히 기능을 추가한 것을 넘어, 실제 하드웨어 성능을 크게 향상시켰습니다. Arrow Lake 아키텍처 기준, 기존 fp32 dequant + SGEMM 대비 성능 수치는 다음과 같습니다.
| BlkLen | VNNI c/MAC | maddubs c/MAC (Fallback) | 성능 향상 (VNNI 기준) |
|---|---|---|---|
| 32 | ~0.051 | ~0.066 | 1.26x |
| 64 | ~0.030 | ~0.047 | 1.52x |
| 128 | ~0.027 | ~0.043 | 1.56x |
교훈:
- 하드웨어 특화 커널의 중요성: 범용적인
dequant + SGEMM보다 데이터 레이아웃에 최적화된 네이티브 커널이 메모리 대역폭과 연산 효율 면에서 압도적입니다. - 안전한 컴파일 플래그 관리: 공통 코드를 여러 ISA 타겟에서 공유할 때는 가장 낮은 공통 분모(Least Common Denominator)의 ISA 플래그를 적용해야 런타임 오류를 방지할 수 있습니다.
리뷰어 피드백 반영
리뷰 과정에서 BiasPtr에 대한 포인터 연산이 nullptr일 경우 발생할 수 있는 UB(Undefined Behavior) 문제가 지적되었습니다. 이를 if (BiasPtr != nullptr) 형태의 가드된 증가문으로 수정하여 코드 안정성을 확보했습니다.
References
참고 자료
- https://github.com/microsoft/onnxruntime/tree/main/onnxruntime/core/mlas
- https://www.intel.com/content/www/us/en/docs/intrinsics-guide/index.html#techs=AVX_VNNI
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [onnxruntime] ONNX Runtime: Arm64 KleidiAI 기반 FP16 GEMM 및 Convolution 최적화
- [onnxruntime] ONNX Runtime: MoE Router GEMV 최적화 및 Bias Fusion 구현
- [onnxruntime] RISC-V 벡터(RVV) 최적화: ONNX Runtime LLM 추론 성능 극대화
- [onnxruntime] ONNX Runtime의 RISC-V Vector(RVV) 최적화: SGEMM과 Softmax 성능을 3배로 끌어올리기
- [triton] Triton에서 Ragged Mode를 위한 X Scale Swizzling 최적화
PR Analysis 의 다른글
- 이전글 [onnxruntime] ONNX Runtime: fpA_intB GEMM 최적화 및 CUDA 그래프 호환성 강화
- 현재글 : [onnxruntime] ONNX Runtime: AVX2 및 AVX-VNNI를 위한 2-bit 가중치 CPU 커널 최적화
- 다음글 [onnxruntime] ONNX Runtime: Arm64 KleidiAI 기반 FP16 GEMM 및 Convolution 최적화
댓글