[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를 직접 계산하지 않고도 출력을 계산할 수 있습니다.
- Single-pass (
-
패널(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_k및d_v제약 조건: 커널이 정상적으로 작동하기 위한 몇 가지 제약 조건이 있습니다:d_k(Key/Query 헤드 차원)는 4의 배수여야 합니다 (d_k % 4 == 0). 이는LinearAttentionDotNeon함수의 4-wide tail 루프 때문입니다.d_k는MlasLinearAttentionNeonMaxKHeadSize(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은 여러 측면에서 훌륭한 최적화 및 개선 사례를 보여줍니다.
-
타겟 아키텍처 특화 최적화: ARM NEON 명령어셋의 특성을 깊이 이해하고 이를 활용하여 커널을 재설계했습니다. 단순히 기존 코드를 이식하는 것이 아니라, NEON의 128비트 레지스터 폭,
vfmaq_n_f32와 같은 명령어의 효율성, 그리고 컴파일러의 잠재적 성능 저하 패턴까지 고려하여 최적의 코드를 작성했습니다. 이는 특정 하드웨어 아키텍처에서 최고의 성능을 끌어내기 위한 엔지니어링의 정수를 보여줍니다. -
알고리즘적 최적화 (Single-pass vs Two-pass): LinearAttention 연산의 수학적 구조를 분석하여, 가능한 경우(linear/gated 규칙) 단일 패스로 연산을 완료하도록 구현했습니다. 이는 메모리 접근 횟수를 줄여 성능을 직접적으로 향상시킵니다. 리뷰어의 지적처럼, 단일 패스는 두 번의 패스보다 하나의 메모리 로드를 줄이는 효과가 있습니다. 이는 성능 향상의 핵심 동력입니다.
-
캐시 효율성 극대화:
MlasLinearAttentionNeonPanelWidth = 32라는 패널 크기 설정은 L1 캐시 크기를 고려하여 설계되었습니다. 데이터를 패널 단위로 처리하고 L1 캐시에 유지함으로써, 두 번째 패스에서 데이터를 다시 메인 메모리에서 읽어오는 비용을 크게 줄일 수 있습니다. 이는 특히 긴 시퀀스에서 반복적으로 발생하는 연산의 성능을 크게 향상시킵니다. -
정확성 보장 및 폴백 메커니즘: 새로운 NEON 커널은 특정 조건(d_k, d_v의 배수 관계, HeadsPerGroup=1 등)에서만 작동합니다. 이러한 조건에 맞지 않는 경우, 기존의 안정적이고 검증된 일반 커널로 안전하게 폴백(fallback)하도록 설계되었습니다. 이는 새로운 최적화된 경로의 버그가 전체 시스템의 안정성을 해치지 않도록 보장하는 중요한 안전장치입니다.
-
테스트 커버리지 강화: 새로운 최적화 경로의 다양한 엣지 케이스(예:
d_k % 4제약 조건,d_k최대치 경계)를 커버하는 테스트 케이스를 추가하여 코드의 신뢰성을 높였습니다. 이는 향후 유지보수 및 기능 추가 시 회귀(regression)를 방지하는 데 필수적입니다. -
코드 품질 및 유지보수성: 리뷰어의 피드백을 반영하여
switch문의default처리 방식을 개선하고, 구조체 초기화 시 명시적 초기화자(designated initializer)를 사용하는 등 코드의 안정성과 가독성을 높였습니다. 특히,switch문의default를 제거하고 각case끝에return을 사용하여 컴파일 타임에 새로운 규칙 추가 시 오류를 명확히 인지하도록 한 점은 매우 좋은 설계입니다.
성능 수치:
PR 설명에 따르면, 이 최적화는
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [onnxruntime] ONNX Runtime MLAS, AVX-512 최적화로 MobileClip-S0 모델 추론 속도 향상
- [onnxruntime] ONNX Runtime: AVX2 및 AVX-VNNI를 위한 2-bit 가중치 CPU 커널 최적화
- [onnxruntime] ONNX Runtime 스레드 풀의 지능형 대기: Exponential Backoff 도입으로 성능 및 전력 효율성 향상
- [onnxruntime] ONNX Runtime CUDA: int64 CumSum 연산 9배 가속화 최적화 분석
- [onnxruntime] ONNX Runtime의 CPU int4 가중치 프리패킹 최적화: 병렬 처리 효율성 개선
PR Analysis 의 다른글
- 이전글 [onnxruntime] ONNX Runtime의 ARM SVE i8mm QGEMM 최적화: 휴대용 머신 코드 전략
- 현재글 : [onnxruntime] ARM NEON 최적화: LinearAttention 커널 융합으로 3배 성능 향상
- 다음글 [flashinfer] SM120(Blackwell)을 위한 초고속 KDA Prefill 커널: FlashInfer의 CuTe DSL 백엔드 분석
댓글