본문으로 건너뛰기

[onnxruntime] ONNX Runtime: Arm64 KleidiAI 기반 FP16 GEMM 및 Convolution 최적화

PR 링크: microsoft/onnxruntime#28786 상태: Merged | 변경: +3263 / -32

들어가며

최근 AI 모델의 경량화와 추론 속도 향상을 위해 FP16(Half-precision) 연산의 중요성이 커지고 있습니다. 특히 Arm64 아키텍처에서 KleidiAI 라이브러리를 활용하면 CPU 기반의 FP16 연산 성능을 극대화할 수 있습니다. 이번 PR은 ONNX Runtime의 MLAS(Microsoft Linear Algebra Subprograms) 레이어에 Arm64 KleidiAI를 위한 FP16 GEMM 및 Convolution 지원을 추가하여, 기존 CPU 연산 대비 향상된 성능을 제공하고자 합니다.

코드 분석

1. MLAS API 확장 (onnxruntime/core/mlas/inc/mlas.h)

이번 변경의 핵심은 MLAS_HALF_GEMM_DATA_PARAMS 구조체에 BackendKernelSelectorConfig를 추가하고, BIsBackendNativePacked 플래그를 도입하여 백엔드 전용 패킹 레이아웃을 지원하는 것입니다.

// Before
struct MLAS_HALF_GEMM_DATA_PARAMS {
    const MLAS_HALF_GEMM_POSTPROCESSOR* OutputProcessor = nullptr;
    bool AIsfp32 = false;
    bool BIsfp32 = false;
};

// After
struct MLAS_HALF_GEMM_DATA_PARAMS {
    // ... 기존 필드
    bool BIsPacked = false;
    const MLAS_BACKEND_KERNEL_SELECTOR_CONFIG* BackendKernelSelectorConfig = nullptr;
    bool BIsBackendNativePacked = false;
};

이 변경을 통해 KleidiAI와 같은 특정 백엔드가 직접 소비할 수 있는 최적화된 데이터 레이아웃을 전달할 수 있게 되었습니다.

2. HalfConv 디스패치 구현 (onnxruntime/core/mlas/lib/halfconv.cpp)

새롭게 추가된 halfconv.cppMlasHalfConv 관련 API를 플랫폼 오버라이드 방식으로 디스패치합니다. 이는 KleidiAI가 지원되지 않는 환경에서는 안전하게 폴백(fallback)할 수 있는 구조를 제공합니다.

bool MLASCALL MlasHalfConv(...) {
    if (GetMlasPlatform().MlasHalfConvOverride == nullptr) {
        return false;
    }
    return GetMlasPlatform().MlasHalfConvOverride(...);
}

왜 이게 좋은가

  1. 성능 최적화: KleidiAI의 SME/SME2 기반 FP16 HGEMM 및 IMATMUL 커널을 직접 호출함으로써, 범용 NEON 코드 대비 높은 연산 효율을 달성합니다.
  2. 유연한 백엔드 선택: MLAS_BACKEND_KERNEL_SELECTOR_CONFIG를 통해 런타임에 KleidiAI 사용 여부를 제어할 수 있어, 특정 하드웨어에서 문제가 발생할 경우 즉시 비활성화할 수 있는 안전장치를 마련했습니다.
  3. 안전한 메모리 관리: thread_local 기반의 스크래치 버퍼를 사용하되, 8MB 제한을 두어 메모리 누수나 과도한 할당을 방지했습니다.

리뷰어 피드백 반영

리뷰 과정에서 MLAS_THROW_EX 사용에 대한 논의가 있었습니다. 초기에는 예외 발생이 MLAS의 관례와 다르다는 지적이 있었으나, native-packed B와 같은 메모리 안전성이 직결된 계약 위반 상황에서는 명시적인 예외 처리가 필요하다는 점에 합의했습니다. 또한, FP16 연산 시 발생하는 정밀도 변화를 고려하여 테스트 허용 오차(tolerance)를 조정하고, 그 이유를 코드 주석으로 명시하여 향후 유지보수성을 높였습니다.

결론

이번 PR은 단순히 기능을 추가하는 것을 넘어, MLAS의 API 계약을 더욱 견고하게 만들고 Arm64 환경에서의 FP16 추론 성능을 한 단계 끌어올리는 중요한 인프라 개선입니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글