본문으로 건너뛰기

[onnxruntime] ARM NEON 최적화: LinearAttention 커널 융합으로 3배 성능 향상

PR 링크: microsoft/onnxruntime#32178 상태: Merged | 변경: +479 / -5

들어가며

최근 Microsoft ONNX Runtime 레포지토리에서는 [MLAS] Add a NEON fused kernel for LinearAttention이라는 제목의 PR이 올라왔습니다. 이 PR은 ARM 아키텍처에서 LinearAttention 연산의 성능을 획기적으로 개선하는 것을 목표로 합니다. 기존의 일반적인 디스패치 방식 대비 최대 3배의 속도 향상을 보고하며, 이는 모바일 및 엣지 디바이스에서의 AI 모델 추론 성능에 상당한 영향을 미칠 수 있습니다. 본 글에서는 이 PR이 어떻게 이러한 성능 향상을 달성했는지, 코드 변경 사항을 중심으로 심층적으로 분석하고 그 의미를 짚어보겠습니다.

LinearAttention은 Transformer 모델의 핵심 구성 요소 중 하나로, 특히 긴 시퀀스를 처리할 때 기존의 Scaled Dot-Product Attention의 계산 복잡도(시퀀스 길이의 제곱에 비례)를 극복하기 위해 제안되었습니다. 이 PR은 MLAS(Microsoft Linear Algebra Subprograms) 라이브러리 내에서 ARM NEON 명령어셋을 활용하여 LinearAttention 연산을 위한 새로운 융합(fused) 커널을 구현함으로써 성능을 최적화합니다.

코드 분석

이번 PR의 핵심은 onnxruntime/core/mlas/lib/linear_attention_kernel_neon.cpp 파일에 새로운 NEON 커널을 추가하고, 이를 빌드 시스템 및 기존 MLAS 구조에 통합하는 것입니다.

1. MLAS 빌드 시스템 통합 (cmake/onnxruntime_mlas.cmake)

먼저, 새로운 NEON 커널 소스 파일을 MLAS 빌드 시스템에 포함시키는 변경 사항입니다. Windows 및 비-Windows 환경 모두에서 ARM64 아키텍처를 위한 소스 목록에 linear_attention_kernel_neon.cpp가 추가되었습니다.

Before:

--- a/cmake/onnxruntime_mlas.cmake
+++ b/cmake/onnxruntime_mlas.cmake
@@ -143,6 +143,7 @@ function(setup_mlas_source_for_windows)
         ${MLAS_SRC_DIR}/eltwise_kernel_neon_fp16.cpp
         ${MLAS_SRC_DIR}/sqnbitgemm_kernel_neon_int8_i8mm.cpp
         ${MLAS_SRC_DIR}/sconv_nchw_depthwise_multiplier_1.cpp
+        ${MLAS_SRC_DIR}/linear_attention_kernel_neon.cpp
       )
 
       set(mlas_platform_preprocess_srcs
@@ -570,6 +571,7 @@ else() 
           ${MLAS_SRC_DIR}/eltwise_kernel_neon.cpp
           ${MLAS_SRC_DIR}/sqnbitgemm_kernel_neon_int8_i8mm.cpp
           ${MLAS_SRC_DIR}/sconv_nchw_depthwise_multiplier_1.cpp
+          ${MLAS_SRC_DIR}/linear_attention_kernel_neon.cpp
         )
 
         # Conditionally add the SVE implementation if compiler supports it

After: 위 diff에서 볼 수 있듯이, linear_attention_kernel_neon.cpp 파일이 ARM64 빌드 시 포함되도록 소스 목록에 추가되었습니다. 이는 MLAS가 ARM64 타겟을 위한 최적화된 커널을 컴파일하고 사용할 수 있도록 하는 기본적인 설정입니다.

2. NEON 융합 커널 구현 (onnxruntime/core/mlas/lib/linear_attention_kernel_neon.cpp)

이 PR의 핵심은 linear_attention_kernel_neon.cpp 파일에 구현된 새로운 NEON 커널입니다. 이 파일은 ARMv8-A ASIMD(Advanced SIMD) 명령어셋을 활용하여 LinearAttention 연산을 효율적으로 처리합니다. 기존의 AVX-512 커널과는 달리, NEON의 특성을 고려하여 재설계되었습니다.

주요 설계 원칙 및 구현:

  • 두 가지 연산 형태: LinearAttention은 크게 두 가지 연산 형태로 나뉩니다.

    • Single-pass (FusedTokenSinglePassNeon): linear/gated 규칙에 사용됩니다. 이 경우 upd (update) 값이 미리 알려져 있어, 상태 행렬 S를 한 번만 읽고 쓰고, 출력 o를 계산할 수 있습니다. 이는 이론적으로 최소한의 연산입니다.
    • Two-pass (FusedTokenTwoPassNeon): delta/gated_delta 규칙에 사용됩니다. 이 경우 upd 값이 S 행렬 전체에 대한 연산 결과를 필요로 하므로, S를 두 번 읽어야 합니다. 첫 번째 패스에서 필요한 중간 결과를 계산하고, 두 번째 패스에서 S를 다시 읽어 업데이트합니다. AVX-512 커널의 항등식을 활용하여 S_new를 직접 계산하지 않고도 출력을 계산할 수 있습니다.
  • 패널(Panel) 기반 처리: NEON은 128비트 레지스터를 사용하므로, AVX-512의 512비트 레지스터보다 처리할 수 있는 데이터 폭이 좁습니다. 이 커널은 MlasLinearAttentionNeonPanelWidth = 32 (float 기준) 크기의 패널을 사용하여 데이터를 처리합니다. 이는 8개의 float32x4_t 벡터에 해당합니다. 이 패널 크기는 L1 캐시에 적재되어 재사용될 수 있도록 설계되었습니다.

  • vfmaq_n_f32 활용: NEON에는 AVX-512의 임베디드 브로드캐스트 기능이 없으므로, 스칼라 값을 벡터 연산에 적용하기 위해 vfmaq_n_f32 (Fused Multiply-Accumulate with scalar)를 사용합니다. 이는 8개의 NEON 레인에서 각 스칼라 연산이 8개의 FMA(Fused Multiply-Add) 연산으로 상쇄되어 효율적입니다.

  • UnrolledLoop 헬퍼: GCC 컴파일러에서 루프 변수를 인덱스로 사용할 때 발생하는 성능 저하(고정 크기 배열을 스택에 유지하고 매번 로드/복원)를 피하기 위해 std::index_sequence를 이용한 컴파일 타임 루프 언롤링(UnrolledLoop<N>)을 사용합니다. 이는 FMA 파이프라인을 최대한 활용하는 데 중요합니다.

  • LinearAttentionDotNeon 함수: 두 개의 벡터(q0, kt)의 내적(dot product)을 계산하는 함수입니다. 4개의 float32x4_t 벡터를 사용하여 누적하고, 최종적으로 수평적 합산(vaddvq_f32)을 수행합니다. 이 함수는 delta/gated_delta 규칙에서 q.k 항을 계산하는 데 사용됩니다.

  • d_kd_v 제약 조건: 커널이 정상적으로 작동하기 위한 몇 가지 제약 조건이 있습니다:

    • d_k (Key/Query 헤드 차원)는 4의 배수여야 합니다 (d_k % 4 == 0). 이는 LinearAttentionDotNeon 함수의 4-wide tail 루프 때문입니다.
    • d_kMlasLinearAttentionNeonMaxKHeadSize (256) 이하이어야 합니다. 이는 두 번의 패스를 위한 스태이징 버퍼 크기 때문입니다.
    • d_v (Value/Output 헤드 차원)는 MlasLinearAttentionNeonPanelWidth (32)의 배수여야 합니다 (d_v % MlasLinearAttentionNeonPanelWidth == 0). 이는 패널 기반 처리의 효율성을 위함입니다.
    • Work->HeadsPerGroup은 1이어야 합니다. 여러 헤드를 그룹화하는 경우(GQA)에는 더 많은 누산기 레지스터가 필요하여 이 커널의 레지스터 압박을 초과할 수 있으므로, 이 경우 일반 커널로 폴백합니다.

이러한 제약 조건에 맞지 않거나 HeadsPerGroup != 1인 경우, 코드는 기존의 일반적인 MlasLinearAttentionProcessHead 함수로 폴백합니다.

// Fallback logic example from the code
const bool shape_ok = (d_k % 4 == 0) &&
                      (d_k <= MlasLinearAttentionNeonMaxKHeadSize) &&
                      (d_v % MlasLinearAttentionNeonPanelWidth == 0);

if (!shape_ok || Work->HeadsPerGroup != 1) {
    MlasLinearAttentionProcessHead(Work);
    return;
}

// NEON kernel execution...

3. 테스트 확장 (onnxruntime/test/mlas/unittest/test_linear_attention.cpp)

새로운 NEON 커널의 정확성과 성능을 검증하기 위해 테스트 케이스가 확장되었습니다. 특히, NEON 커널의 d_k % 4 == 0 제약 조건과 d_k <= 256 경계를 테스트하는 새로운 모양(shape)들이 추가되었습니다.

주요 테스트 케이스:

  • {12, 32}: d_k=12는 4의 배수이지만 16의 배수는 아닙니다. 이는 LinearAttentionDotNeon 함수의 4-wide tail 루프가 메인 루프와 함께 실행되는 경우를 테스트합니다.
  • {20, 32}: d_k=20 역시 4의 배수이지만 16의 배수는 아닙니다. 메인 루프(16개)와 테일 루프(4개)가 모두 실행되는 경우를 테스트합니다.
  • {256, 32}: d_k=256은 스태이징 버퍼의 최대 크기와 정확히 일치하는 경우로, 경계 조건을 테스트합니다.

이러한 테스트 케이스들은 새로운 NEON 커널이 다양한 d_k 값에 대해 올바르게 작동하는지, 특히 기존 AVX-512 커널이 지원하지 않는 d_k % 4 제약 조건을 어떻게 처리하는지 검증하는 데 중요합니다.

왜 이게 좋은가?

이 PR은 여러 측면에서 훌륭한 최적화 및 개선 사례를 보여줍니다.

  1. 타겟 아키텍처 특화 최적화: ARM NEON 명령어셋의 특성을 깊이 이해하고 이를 활용하여 커널을 재설계했습니다. 단순히 기존 코드를 이식하는 것이 아니라, NEON의 128비트 레지스터 폭, vfmaq_n_f32와 같은 명령어의 효율성, 그리고 컴파일러의 잠재적 성능 저하 패턴까지 고려하여 최적의 코드를 작성했습니다. 이는 특정 하드웨어 아키텍처에서 최고의 성능을 끌어내기 위한 엔지니어링의 정수를 보여줍니다.

  2. 알고리즘적 최적화 (Single-pass vs Two-pass): LinearAttention 연산의 수학적 구조를 분석하여, 가능한 경우(linear/gated 규칙) 단일 패스로 연산을 완료하도록 구현했습니다. 이는 메모리 접근 횟수를 줄여 성능을 직접적으로 향상시킵니다. 리뷰어의 지적처럼, 단일 패스는 두 번의 패스보다 하나의 메모리 로드를 줄이는 효과가 있습니다. 이는 성능 향상의 핵심 동력입니다.

  3. 캐시 효율성 극대화: MlasLinearAttentionNeonPanelWidth = 32라는 패널 크기 설정은 L1 캐시 크기를 고려하여 설계되었습니다. 데이터를 패널 단위로 처리하고 L1 캐시에 유지함으로써, 두 번째 패스에서 데이터를 다시 메인 메모리에서 읽어오는 비용을 크게 줄일 수 있습니다. 이는 특히 긴 시퀀스에서 반복적으로 발생하는 연산의 성능을 크게 향상시킵니다.

  4. 정확성 보장 및 폴백 메커니즘: 새로운 NEON 커널은 특정 조건(d_k, d_v의 배수 관계, HeadsPerGroup=1 등)에서만 작동합니다. 이러한 조건에 맞지 않는 경우, 기존의 안정적이고 검증된 일반 커널로 안전하게 폴백(fallback)하도록 설계되었습니다. 이는 새로운 최적화된 경로의 버그가 전체 시스템의 안정성을 해치지 않도록 보장하는 중요한 안전장치입니다.

  5. 테스트 커버리지 강화: 새로운 최적화 경로의 다양한 엣지 케이스(예: d_k % 4 제약 조건, d_k 최대치 경계)를 커버하는 테스트 케이스를 추가하여 코드의 신뢰성을 높였습니다. 이는 향후 유지보수 및 기능 추가 시 회귀(regression)를 방지하는 데 필수적입니다.

  6. 코드 품질 및 유지보수성: 리뷰어의 피드백을 반영하여 switch 문의 default 처리 방식을 개선하고, 구조체 초기화 시 명시적 초기화자(designated initializer)를 사용하는 등 코드의 안정성과 가독성을 높였습니다. 특히, switch 문의 default를 제거하고 각 case 끝에 return을 사용하여 컴파일 타임에 새로운 규칙 추가 시 오류를 명확히 인지하도록 한 점은 매우 좋은 설계입니다.

성능 수치:

PR 설명에 따르면, 이 최적화는

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

댓글

관련 포스트

PR Analysis 의 다른글