본문으로 건너뛰기

[onnxruntime] ONNX Runtime: AVX2 및 AVX-VNNI를 위한 2-bit 가중치 CPU 커널 최적화

PR 링크: microsoft/onnxruntime#29619 상태: Merged | 변경: +1983 / -12

들어가며

최근 LLM 및 경량화 모델의 추론 성능을 극대화하기 위해 2-bit 양자화(W2) 기법이 널리 사용되고 있습니다. 하지만 기존 microsoft/onnxruntimeMLAS(Microsoft Linear Algebra Subroutine) 라이브러리에서는 AVX-512를 지원하지 않는 일반적인 x86 CPU(Alder Lake, Zen 1-3 등) 환경에서 2-bit MatMul 연산 시, fp32 dequantizationSGEMM을 호출하는 비효율적인 폴백(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

교훈:

  1. 하드웨어 특화 커널의 중요성: 범용적인 dequant + SGEMM보다 데이터 레이아웃에 최적화된 네이티브 커널이 메모리 대역폭과 연산 효율 면에서 압도적입니다.
  2. 안전한 컴파일 플래그 관리: 공통 코드를 여러 ISA 타겟에서 공유할 때는 가장 낮은 공통 분모(Least Common Denominator)의 ISA 플래그를 적용해야 런타임 오류를 방지할 수 있습니다.

리뷰어 피드백 반영

리뷰 과정에서 BiasPtr에 대한 포인터 연산이 nullptr일 경우 발생할 수 있는 UB(Undefined Behavior) 문제가 지적되었습니다. 이를 if (BiasPtr != nullptr) 형태의 가드된 증가문으로 수정하여 코드 안정성을 확보했습니다.

References

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글