From cc30a146370218897f0f0d057cd131661d9bc091 Mon Sep 17 00:00:00 2001 From: Rohan Date: Wed, 11 Jun 2025 15:25:05 +0000 Subject: [PATCH 01/26] Rewire ORT to support a NEON version of NCHWc Conv --- cmake/onnxruntime_mlas.cmake | 1 + onnxruntime/core/mlas/lib/mlasi.h | 11 +++++++++++ onnxruntime/core/mlas/lib/platform.cpp | 6 ++++++ onnxruntime/core/mlas/lib/snchwc.cpp | 16 ++++++++++------ 4 files changed, 28 insertions(+), 6 deletions(-) diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index 24cecf07e8e36..3ae159502c294 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -442,6 +442,7 @@ else() ${MLAS_SRC_DIR}/aarch64/QgemmS8S8KernelSmmla.S ${MLAS_SRC_DIR}/aarch64/QgemmU8X8KernelUmmla.S ${MLAS_SRC_DIR}/aarch64/SbgemmKernelNeon.S + ${MLAS_SRC_DIR}/aarch64/SConvKernelNeon.S ${MLAS_SRC_DIR}/activate_fp16.cpp ${MLAS_SRC_DIR}/dwconv.cpp ${MLAS_SRC_DIR}/halfgemm_kernel_neon.cpp diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h index a099bcf8438fe..033968da3e189 100644 --- a/onnxruntime/core/mlas/lib/mlasi.h +++ b/onnxruntime/core/mlas/lib/mlasi.h @@ -949,6 +949,10 @@ extern "C" { #if defined(__aarch64__) && defined(__linux__) MLAS_SBGEMM_FLOAT_KERNEL MlasSbgemmKernelZero; MLAS_SBGEMM_FLOAT_KERNEL MlasSbgemmKernelAdd; + MLAS_CONV_FLOAT_KERNEL MlasConvNchwFloatKernelNeon; + MLAS_CONV_FLOAT_KERNEL MlasConvNchwcFloatKernelNeon; + MLAS_CONV_DEPTHWISE_FLOAT_KERNEL MlasConvDepthwiseFloatKernelNeon; + MLAS_CONV_POINTWISE_FLOAT_KERNEL MlasConvPointwiseFloatKernelNeon; #endif MLAS_GEMM_DOUBLE_KERNEL MlasDgemmKernelZero; MLAS_GEMM_DOUBLE_KERNEL MlasDgemmKernelAdd; @@ -1332,6 +1336,13 @@ struct MLAS_PLATFORM { const MLAS_GEMM_QUANT_DISPATCH* GemmU8U8Dispatch; const MLAS_GEMM_QUANT_DISPATCH* GemmU8S8Dispatch; const MLAS_GEMM_QUANT_DISPATCH* GemmS8S8Dispatch; +#endif +#if defined(__aarch64__) && defined(__linux__) + MLAS_CONV_FLOAT_KERNEL* ConvNchwFloatKernel; + MLAS_CONV_FLOAT_KERNEL* ConvNchwcFloatKernel; + MLAS_CONV_DEPTHWISE_FLOAT_KERNEL* ConvDepthwiseFloatKernel; + MLAS_CONV_POINTWISE_FLOAT_KERNEL* ConvPointwiseFloatKernel; + uint32_t NchwcBlockSize; #endif const MLAS_SYMM_QGEMM_DISPATCH* SymmQgemmDispatch{nullptr}; diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp index 3256dadb856d3..9489042c28ccd 100644 --- a/onnxruntime/core/mlas/lib/platform.cpp +++ b/onnxruntime/core/mlas/lib/platform.cpp @@ -558,6 +558,12 @@ Return Value: this->SoftmaxDispatch = &MlasSoftmaxDispatchNeon; this->EltwiseDispatch = &MlasEltwiseDispatchNeon; + this->ConvNchwFloatKernel = MlasConvNchwFloatKernelNeon; + this->ConvNchwcFloatKernel = MlasConvNchwcFloatKernelNeon; + this->ConvDepthwiseFloatKernel = MlasConvDepthwiseFloatKernelNeon; + this->ConvPointwiseFloatKernel = MlasConvPointwiseFloatKernelNeon; + this->NchwcBlockSize = 16; // What is it supposed to be? + // // Check if the processor supports ASIMD dot product instructions. // diff --git a/onnxruntime/core/mlas/lib/snchwc.cpp b/onnxruntime/core/mlas/lib/snchwc.cpp index f9cf1605787aa..c36a4a8fac71a 100644 --- a/onnxruntime/core/mlas/lib/snchwc.cpp +++ b/onnxruntime/core/mlas/lib/snchwc.cpp @@ -101,7 +101,7 @@ Return Value: --*/ { -#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) +#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) || defined(MLAS_TARGET_ARM64) return GetMlasPlatform().NchwcBlockSize; #else return 1; @@ -674,7 +674,7 @@ struct MLAS_NCHWC_CONV_NCHWC_ALGORITHM : MLAS_NCHWC_GROUPED_CONV_ALGORITHM const size_t BlockedOutputWidth = BlockSize * OutputWidth; -#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) +#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) || defined(MLAS_TARGET_ARM64) MLAS_CONV_FLOAT_KERNEL* Kernel = GetMlasPlatform().ConvNchwcFloatKernel; #else MLAS_CONV_FLOAT_KERNEL* Kernel = MlasConvNchwcFloatKernel; @@ -784,7 +784,7 @@ struct MLAS_NCHWC_CONV_NCHW_ALGORITHM : MLAS_NCHWC_GROUPED_CONV_ALGORITHM const size_t BlockedOutputWidth = BlockSize * OutputWidth; -#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) +#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) || defined(MLAS_TARGET_ARM64) MLAS_CONV_FLOAT_KERNEL* Kernel = GetMlasPlatform().ConvNchwFloatKernel; #else MLAS_CONV_FLOAT_KERNEL* Kernel = MlasConvNchwFloatKernel; @@ -879,7 +879,7 @@ struct MLAS_NCHWC_CONV_POINTWISE_ALGORITHM : MLAS_NCHWC_GROUPED_CONV_ALGORITHM const size_t FilterStrideBytes = BlockSize * InputChannels * sizeof(float); const size_t OutputStrideBytes = BlockSize * OutputSize * sizeof(float); -#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) +#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) || defined(MLAS_TARGET_ARM64) MLAS_CONV_POINTWISE_FLOAT_KERNEL* Kernel = GetMlasPlatform().ConvPointwiseFloatKernel; #else MLAS_CONV_POINTWISE_FLOAT_KERNEL* Kernel = MlasConvPointwiseFloatKernel; @@ -1016,7 +1016,7 @@ struct MLAS_NCHWC_CONV_DEPTHWISE_ALGORITHM : MLAS_NCHWC_CONV_ALGORITHM const size_t BlockedOutputWidth = BlockSize * OutputWidth; -#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) +#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) || defined(MLAS_TARGET_ARM64) MLAS_CONV_DEPTHWISE_FLOAT_KERNEL* Kernel = GetMlasPlatform().ConvDepthwiseFloatKernel; #else MLAS_CONV_DEPTHWISE_FLOAT_KERNEL* Kernel = MlasConvDepthwiseFloatKernel; @@ -1621,7 +1621,7 @@ Return Value: } } -#if !defined(MLAS_TARGET_AMD64) && !defined(MLAS_TARGET_LARCH64) +#if !defined(MLAS_TARGET_AMD64) && !defined(MLAS_TARGET_LARCH64) && !defined(MLAS_TARGET_ARM64) // // Convolution and pooling kernel stubs for architectures that do not yet have @@ -1788,6 +1788,10 @@ MlasConvPointwiseFloatKernel( MLAS_UNREFERENCED_PARAMETER(Flags); } +#endif + +#if !defined(MLAS_TARGET_AMD64) && !defined(MLAS_TARGET_LARCH64) + void MLASCALL MlasPoolMaximumFloatKernel( From f190c0d706d01e3219b588d272e62de64e5bb27e Mon Sep 17 00:00:00 2001 From: Rohan Date: Tue, 17 Jun 2025 18:12:39 +0000 Subject: [PATCH 02/26] Remove reference to assembly file --- cmake/onnxruntime_mlas.cmake | 2 +- onnxruntime/core/mlas/lib/sconv.h | 0 onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp | 0 3 files changed, 1 insertion(+), 1 deletion(-) create mode 100644 onnxruntime/core/mlas/lib/sconv.h create mode 100644 onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index 3ae159502c294..80445dd7746be 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -442,7 +442,6 @@ else() ${MLAS_SRC_DIR}/aarch64/QgemmS8S8KernelSmmla.S ${MLAS_SRC_DIR}/aarch64/QgemmU8X8KernelUmmla.S ${MLAS_SRC_DIR}/aarch64/SbgemmKernelNeon.S - ${MLAS_SRC_DIR}/aarch64/SConvKernelNeon.S ${MLAS_SRC_DIR}/activate_fp16.cpp ${MLAS_SRC_DIR}/dwconv.cpp ${MLAS_SRC_DIR}/halfgemm_kernel_neon.cpp @@ -450,6 +449,7 @@ else() ${MLAS_SRC_DIR}/qgemm_kernel_smmla.cpp ${MLAS_SRC_DIR}/qgemm_kernel_ummla.cpp ${MLAS_SRC_DIR}/sbgemm_kernel_neon.cpp + ${MLAS_SRC_DIR}/sconv_kernel_neon.cpp ${MLAS_SRC_DIR}/cast_kernel_neon.cpp ${MLAS_SRC_DIR}/hqnbitgemm_kernel_neon_fp16.cpp ${MLAS_SRC_DIR}/rotary_embedding_kernel_neon_fp16.cpp diff --git a/onnxruntime/core/mlas/lib/sconv.h b/onnxruntime/core/mlas/lib/sconv.h new file mode 100644 index 0000000000000..e69de29bb2d1d diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp new file mode 100644 index 0000000000000..e69de29bb2d1d From 632870bbf0f1d8d2c194ca1c4511882c80667bb5 Mon Sep 17 00:00:00 2001 From: Rohan Date: Fri, 20 Jun 2025 22:30:33 +0000 Subject: [PATCH 03/26] Add a NEON kernel for Pointwise Convolution --- onnxruntime/core/mlas/lib/platform.cpp | 2 +- onnxruntime/core/mlas/lib/sconv.h | 139 +++++ .../core/mlas/lib/sconv_kernel_neon.cpp | 486 ++++++++++++++++++ 3 files changed, 626 insertions(+), 1 deletion(-) diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp index 9489042c28ccd..00dfc86b5eaf1 100644 --- a/onnxruntime/core/mlas/lib/platform.cpp +++ b/onnxruntime/core/mlas/lib/platform.cpp @@ -562,7 +562,7 @@ Return Value: this->ConvNchwcFloatKernel = MlasConvNchwcFloatKernelNeon; this->ConvDepthwiseFloatKernel = MlasConvDepthwiseFloatKernelNeon; this->ConvPointwiseFloatKernel = MlasConvPointwiseFloatKernelNeon; - this->NchwcBlockSize = 16; // What is it supposed to be? + this->NchwcBlockSize = 4; // What is it supposed to be? // // Check if the processor supports ASIMD dot product instructions. diff --git a/onnxruntime/core/mlas/lib/sconv.h b/onnxruntime/core/mlas/lib/sconv.h index e69de29bb2d1d..9cdd45abdc9ec 100644 --- a/onnxruntime/core/mlas/lib/sconv.h +++ b/onnxruntime/core/mlas/lib/sconv.h @@ -0,0 +1,139 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + sconv.h + +Abstract: + + This module contains the private data structures and procedure prototypes + for the single precision convolution operation. + +--*/ + +#pragma once + +#include + +// +// Define the calling convention for MLAS functions. +// + +#ifndef MLASCALL +#if defined(_WIN32) && !defined(_WIN64) +#define MLASCALL __stdcall +#else +#define MLASCALL +#endif +#endif + +// +// Define the convolution kernel flags. +// + +#define MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT 0x00000001 +#define MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION 0x00000002 +#define MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION 0x00000004 +#define MLAS_CONV_KERNEL_FLAG_OTHER_ACTIVATION 0x00000008 + +// +// Define the prototypes of the NEON convolution kernels. +// + +#if defined(__aarch64__) || defined(_M_ARM64) + +extern "C" { + +void +MLASCALL +MlasConvNchwFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t FilterCount, + size_t InputStride, + size_t FilterStride, + size_t OutputStride, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + const float* Bias, + unsigned KernelFlags + ); + +void +MLASCALL +MlasConvNchwcFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t FilterCount, + size_t InputStride, + size_t FilterStride, + size_t OutputStride, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + const float* Bias, + unsigned KernelFlags + ); + +void +MLASCALL +MlasConvDepthwiseFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t InputStride, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + const float* Bias, + unsigned KernelFlags + ); + +void +MLASCALL +MlasConvPointwiseFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t InputChannels, + size_t FilterCount, + size_t InputStride, + size_t FilterStride, + size_t OutputStride, + size_t OutputCount, + const float* Bias, + unsigned KernelFlags + ); + +} + +#endif diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index e69de29bb2d1d..a0ba4a937d120 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -0,0 +1,486 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + sconv_kernel_neon.cpp + +Abstract: + + This module implements the single precision convolution kernels for ARM NEON. + +--*/ + +#include "sconv.h" + +#if defined(__aarch64__) || defined(_M_ARM64) + +#include +#include + +#include "arm_neon.h" +#include "mlasi.h" + +void + MLASCALL + MlasConvNchwFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t FilterCount, + size_t InputStride, + size_t FilterStride, + size_t OutputStride, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + const float* Bias, + unsigned KernelFlags + ) +{ + // Mark unused parameters + (void)InputBase; + (void)InputWidth; + + const bool AccumulateOutput = (KernelFlags & MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT) != 0; + const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; + const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; + + const float32x4_t ZeroVector = vdupq_n_f32(0.0f); + + // Convert byte strides to element strides + const size_t StrideWidthElements = StrideWidth / sizeof(float); + const size_t DilationWidthElements = DilationWidth / sizeof(float); + const size_t InputStrideElements = InputStride / sizeof(float); + const size_t FilterStrideElements = FilterStride / sizeof(float); + const size_t OutputStrideElements = OutputStride / sizeof(float); + const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); + + for (size_t f = 0; f < FilterCount; f++) { + const float* filter = Filter + f * FilterStrideElements; + float* output = Output + f * OutputStrideElements; + const float bias = BiasAddition ? Bias[f] : 0.0f; + const float32x4_t BiasVector = vdupq_n_f32(bias); + + const float* input_ptr = Input; + size_t OutputIndex = 0; + + // Handle left padding + for (size_t i = 0; i < OutputCountLeftPad; i++) { + float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; + const float filter_value = filter[kh * KernelWidth + kw]; + const float32x4_t FilterVector = vdupq_n_f32(filter_value); + const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + + MlasStoreFloat32x4(&output[OutputIndex], Accumulator); + OutputIndex += 4; + input_ptr += StrideWidthElements; + } + + // Handle main output region + for (size_t i = 0; i < OutputCount; i++) { + float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; + const float filter_value = filter[kh * KernelWidth + kw]; + const float32x4_t FilterVector = vdupq_n_f32(filter_value); + const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + + MlasStoreFloat32x4(&output[OutputIndex], Accumulator); + OutputIndex += 4; + input_ptr += StrideWidthElements; + } + + // Handle right padding + for (size_t i = 0; i < OutputCountRightPad; i++) { + float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; + const float filter_value = filter[kh * KernelWidth + kw]; + const float32x4_t FilterVector = vdupq_n_f32(filter_value); + const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + + MlasStoreFloat32x4(&output[OutputIndex], Accumulator); + OutputIndex += 4; + input_ptr += StrideWidthElements; + } + + // Move to next row + Input += InputStrideElements; + } +} + +// +// Implementation of MlasConvNchwcFloatKernelNeon +// +// This kernel performs 2D convolution optimized for NCHWc format. +// It processes multiple filters in blocks to improve cache efficiency. +// + +void + MLASCALL + MlasConvNchwcFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t FilterCount, + size_t InputStride, + size_t FilterStride, + size_t OutputStride, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + const float* Bias, + unsigned KernelFlags + ) +{ + // Mark unused parameters + (void)InputBase; + (void)InputWidth; + + const bool AccumulateOutput = (KernelFlags & MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT) != 0; + const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; + const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; + + const float32x4_t ZeroVector = vdupq_n_f32(0.0f); + + // Convert byte strides to element strides + const size_t StrideWidthElements = StrideWidth / sizeof(float); + const size_t DilationWidthElements = DilationWidth / sizeof(float); + const size_t InputStrideElements = InputStride / sizeof(float); + const size_t FilterStrideElements = FilterStride / sizeof(float); + const size_t OutputStrideElements = OutputStride / sizeof(float); + const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); + + // Mark unused variables in NCHWc kernel + (void)InputStrideElements; + + // Process all filters + for (size_t f = 0; f < FilterCount; f++) { + const float* filter = Filter + f * FilterStrideElements; + float* output = Output + f * OutputStrideElements; + const float bias = BiasAddition ? Bias[f] : 0.0f; + const float32x4_t BiasVector = vdupq_n_f32(bias); + + const float* input_ptr = Input; + size_t OutputIndex = 0; + + // Handle left padding + for (size_t i = 0; i < OutputCountLeftPad; i++) { + float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; + const float filter_value = filter[kh * KernelWidth + kw]; + const float32x4_t FilterVector = vdupq_n_f32(filter_value); + const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + + MlasStoreFloat32x4(&output[OutputIndex], Accumulator); + OutputIndex += 4; + input_ptr += StrideWidthElements; + } + + // Handle main output region + for (size_t i = 0; i < OutputCount; i++) { + float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; + const float filter_value = filter[kh * KernelWidth + kw]; + const float32x4_t FilterVector = vdupq_n_f32(filter_value); + const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + + MlasStoreFloat32x4(&output[OutputIndex], Accumulator); + OutputIndex += 4; + input_ptr += StrideWidthElements; + } + + // Handle right padding + for (size_t i = 0; i < OutputCountRightPad; i++) { + float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; + const float filter_value = filter[kh * KernelWidth + kw]; + const float32x4_t FilterVector = vdupq_n_f32(filter_value); + const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + + MlasStoreFloat32x4(&output[OutputIndex], Accumulator); + OutputIndex += 4; + input_ptr += StrideWidthElements; + } + } +} + +// +// Implementation of MlasConvDepthwiseFloatKernelNeon +// +// This kernel performs depthwise separable convolution where each input channel +// is convolved with its own filter. This is more efficient than standard convolution +// for certain network architectures like MobileNets. +// + +void + MLASCALL + MlasConvDepthwiseFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t InputStride, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + const float* Bias, + unsigned KernelFlags + ) +{ + // Mark unused parameters + (void)InputBase; + (void)InputWidth; + + const bool AccumulateOutput = (KernelFlags & MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT) != 0; + const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; + const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; + + const float32x4_t ZeroVector = vdupq_n_f32(0.0f); + + // Convert byte strides to element strides + const size_t StrideWidthElements = StrideWidth / sizeof(float); + const size_t DilationWidthElements = DilationWidth / sizeof(float); + const size_t InputStrideElements = InputStride / sizeof(float); + const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); + + // Mark unused variables in Depthwise kernel + (void)InputStrideElements; + + // For depthwise convolution, process 4 channels at a time + const float bias = BiasAddition ? Bias[0] : 0.0f; + const float32x4_t BiasVector = vdupq_n_f32(bias); + + const float* input_ptr = Input; + size_t OutputIndex = 0; + + // Handle left padding + for (size_t i = 0; i < OutputCountLeftPad; i++) { + float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&Output[OutputIndex]) : BiasVector; + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; + const float filter_value = Filter[kh * KernelWidth + kw]; + const float32x4_t FilterVector = vdupq_n_f32(filter_value); + const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + + MlasStoreFloat32x4(&Output[OutputIndex], Accumulator); + OutputIndex += 4; + input_ptr += StrideWidthElements; + } + + // Handle main output region + for (size_t i = 0; i < OutputCount; i++) { + float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&Output[OutputIndex]) : BiasVector; + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; + const float filter_value = Filter[kh * KernelWidth + kw]; + const float32x4_t FilterVector = vdupq_n_f32(filter_value); + const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + + MlasStoreFloat32x4(&Output[OutputIndex], Accumulator); + OutputIndex += 4; + input_ptr += StrideWidthElements; + } + + // Handle right padding + for (size_t i = 0; i < OutputCountRightPad; i++) { + float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&Output[OutputIndex]) : BiasVector; + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; + const float filter_value = Filter[kh * KernelWidth + kw]; + const float32x4_t FilterVector = vdupq_n_f32(filter_value); + const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + + MlasStoreFloat32x4(&Output[OutputIndex], Accumulator); + OutputIndex += 4; + input_ptr += StrideWidthElements; + } +} + +// +// Implementation of MlasConvPointwiseFloatKernelNeon +// +// This kernel performs pointwise (1x1) convolution which is essentially +// a matrix multiplication across the channel dimension. It's optimized +// for cases where the kernel size is 1x1. +// + +void + MLASCALL + MlasConvPointwiseFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t InputChannels, + size_t FilterCount, + size_t InputStride, + size_t FilterStride, + size_t OutputStride, + size_t OutputCount, + const float* Bias, + unsigned KernelFlags + ) +{ + const bool AccumulateOutput = (KernelFlags & MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT) != 0; + const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; + const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; + + const size_t StrideWidthElements = StrideWidth / sizeof(float); + const size_t InputStrideElements = InputStride / sizeof(float); + const size_t FilterStrideElements = FilterStride / sizeof(float); + const size_t OutputStrideElements = OutputStride / sizeof(float); + + const size_t BlockSize = MlasNchwcGetBlockSize(); + const float32x4_t ZeroVector = vdupq_n_f32(0.0f); + + for (size_t i = 0; i < OutputCount; i++) { + for (size_t f = 0; f < FilterCount; f++) { + const float* filter = Filter + f * FilterStrideElements; + float* output = Output + f * OutputStrideElements; + float32x4_t Accumulator; + if (AccumulateOutput) { + Accumulator = MlasLoadFloat32x4(&output[i * BlockSize]); + } else if (BiasAddition) { + Accumulator = MlasLoadFloat32x4(&Bias[f * BlockSize]); + } else { + Accumulator = vdupq_n_f32(0.0f); + } + for (size_t c = 0; c < InputChannels; c++) { + const float* input_ptr = Input + c * InputStrideElements + i * StrideWidthElements; + for (size_t input_b = 0; input_b < BlockSize; input_b++) { + const float input_value = input_ptr[input_b]; + const float32x4_t InputVector = vdupq_n_f32(input_value); + const float* filter_ptr = filter + (c * BlockSize + input_b) * BlockSize; + const float32x4_t FilterVector = MlasLoadFloat32x4(filter_ptr); + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + MlasStoreFloat32x4(&output[i * BlockSize], Accumulator); + } + } +} + +#endif // __aarch64__ || _M_ARM64 \ No newline at end of file From 159570a9372344b02dd653d2cac90cc7cac72fa3 Mon Sep 17 00:00:00 2001 From: Rohan Date: Thu, 26 Jun 2025 20:32:54 +0000 Subject: [PATCH 04/26] Add a NEON kernel for Depthwise --- .../core/mlas/lib/sconv_kernel_neon.cpp | 103 +++++++----------- 1 file changed, 38 insertions(+), 65 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index a0ba4a937d120..ea44cca7d973e 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -317,90 +317,65 @@ void unsigned KernelFlags ) { - // Mark unused parameters - (void)InputBase; - (void)InputWidth; - const bool AccumulateOutput = (KernelFlags & MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT) != 0; const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; + const size_t BlockSize = MlasNchwcGetBlockSize(); const float32x4_t ZeroVector = vdupq_n_f32(0.0f); - // Convert byte strides to element strides const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); const size_t InputStrideElements = InputStride / sizeof(float); const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); - // Mark unused variables in Depthwise kernel (void)InputStrideElements; - // For depthwise convolution, process 4 channels at a time - const float bias = BiasAddition ? Bias[0] : 0.0f; - const float32x4_t BiasVector = vdupq_n_f32(bias); + const size_t InputWidthElements = InputWidth / sizeof(float); - const float* input_ptr = Input; - size_t OutputIndex = 0; + const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; - // Handle left padding - for (size_t i = 0; i < OutputCountLeftPad; i++) { - float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&Output[OutputIndex]) : BiasVector; - - for (size_t kh = 0; kh < KernelHeight; kh++) { - for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float filter_value = Filter[kh * KernelWidth + kw]; - const float32x4_t FilterVector = vdupq_n_f32(filter_value); - const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { + bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); - } - } + float32x4_t Accumulator; - if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + if (AccumulateOutput) { + Accumulator = MlasLoadFloat32x4(&Output[output_idx * BlockSize]); + } else if (BiasAddition) { + Accumulator = MlasLoadFloat32x4(Bias); + } else { + Accumulator = vdupq_n_f32(0.0f); } - MlasStoreFloat32x4(&Output[OutputIndex], Accumulator); - OutputIndex += 4; - input_ptr += StrideWidthElements; - } - - // Handle main output region - for (size_t i = 0; i < OutputCount; i++) { - float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&Output[OutputIndex]) : BiasVector; - for (size_t kh = 0; kh < KernelHeight; kh++) { for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float filter_value = Filter[kh * KernelWidth + kw]; - const float32x4_t FilterVector = vdupq_n_f32(filter_value); - const float32x4_t InputVector = MlasLoadFloat32x4(input_element); - - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); - } - } - - if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); - } - - MlasStoreFloat32x4(&Output[OutputIndex], Accumulator); - OutputIndex += 4; - input_ptr += StrideWidthElements; - } - - // Handle right padding - for (size_t i = 0; i < OutputCountRightPad; i++) { - float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&Output[OutputIndex]) : BiasVector; + size_t kernel_pos = kh * KernelWidth + kw; + + const float* input_base = Input + output_idx * StrideWidthElements + + kh * DilatedInputWidthElements + kw * DilationWidthElements; + + float32x4_t InputVector; + + if (is_main_region) { + InputVector = MlasLoadFloat32x4(input_base); + } else { + float input_values[4]; + for (size_t i = 0; i < BlockSize; i++) { + const float* input_element = input_base + i; + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; + + if (input_element >= input_row_start && input_element < input_row_end) { + input_values[i] = *input_element; + } else { + input_values[i] = 0.0f; + } + } + InputVector = MlasLoadFloat32x4(input_values); + } - for (size_t kh = 0; kh < KernelHeight; kh++) { - for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float filter_value = Filter[kh * KernelWidth + kw]; - const float32x4_t FilterVector = vdupq_n_f32(filter_value); - const float32x4_t InputVector = MlasLoadFloat32x4(input_element); + const float32x4_t FilterVector = MlasLoadFloat32x4(&Filter[kernel_pos * BlockSize]); Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); } @@ -410,9 +385,7 @@ void Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); } - MlasStoreFloat32x4(&Output[OutputIndex], Accumulator); - OutputIndex += 4; - input_ptr += StrideWidthElements; + MlasStoreFloat32x4(&Output[output_idx * BlockSize], Accumulator); } } From 52f09bf4685c4de1fa23b1e68766e058dd17cdff Mon Sep 17 00:00:00 2001 From: Rohan Date: Thu, 26 Jun 2025 22:33:44 +0000 Subject: [PATCH 05/26] Remove placeholder implementations --- .../core/mlas/lib/sconv_kernel_neon.cpp | 207 +----------------- 1 file changed, 2 insertions(+), 205 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index ea44cca7d973e..36e25ee51232d 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -48,108 +48,7 @@ void unsigned KernelFlags ) { - // Mark unused parameters - (void)InputBase; - (void)InputWidth; - - const bool AccumulateOutput = (KernelFlags & MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT) != 0; - const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; - const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; - - const float32x4_t ZeroVector = vdupq_n_f32(0.0f); - - // Convert byte strides to element strides - const size_t StrideWidthElements = StrideWidth / sizeof(float); - const size_t DilationWidthElements = DilationWidth / sizeof(float); - const size_t InputStrideElements = InputStride / sizeof(float); - const size_t FilterStrideElements = FilterStride / sizeof(float); - const size_t OutputStrideElements = OutputStride / sizeof(float); - const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); - - for (size_t f = 0; f < FilterCount; f++) { - const float* filter = Filter + f * FilterStrideElements; - float* output = Output + f * OutputStrideElements; - const float bias = BiasAddition ? Bias[f] : 0.0f; - const float32x4_t BiasVector = vdupq_n_f32(bias); - - const float* input_ptr = Input; - size_t OutputIndex = 0; - - // Handle left padding - for (size_t i = 0; i < OutputCountLeftPad; i++) { - float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; - - for (size_t kh = 0; kh < KernelHeight; kh++) { - for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float filter_value = filter[kh * KernelWidth + kw]; - const float32x4_t FilterVector = vdupq_n_f32(filter_value); - const float32x4_t InputVector = MlasLoadFloat32x4(input_element); - - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); - } - } - - if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); - } - - MlasStoreFloat32x4(&output[OutputIndex], Accumulator); - OutputIndex += 4; - input_ptr += StrideWidthElements; - } - - // Handle main output region - for (size_t i = 0; i < OutputCount; i++) { - float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; - - for (size_t kh = 0; kh < KernelHeight; kh++) { - for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float filter_value = filter[kh * KernelWidth + kw]; - const float32x4_t FilterVector = vdupq_n_f32(filter_value); - const float32x4_t InputVector = MlasLoadFloat32x4(input_element); - - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); - } - } - - if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); - } - - MlasStoreFloat32x4(&output[OutputIndex], Accumulator); - OutputIndex += 4; - input_ptr += StrideWidthElements; - } - - // Handle right padding - for (size_t i = 0; i < OutputCountRightPad; i++) { - float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; - - for (size_t kh = 0; kh < KernelHeight; kh++) { - for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float filter_value = filter[kh * KernelWidth + kw]; - const float32x4_t FilterVector = vdupq_n_f32(filter_value); - const float32x4_t InputVector = MlasLoadFloat32x4(input_element); - - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); - } - } - - if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); - } - - MlasStoreFloat32x4(&output[OutputIndex], Accumulator); - OutputIndex += 4; - input_ptr += StrideWidthElements; - } - - // Move to next row - Input += InputStrideElements; - } + } // @@ -183,109 +82,7 @@ void unsigned KernelFlags ) { - // Mark unused parameters - (void)InputBase; - (void)InputWidth; - - const bool AccumulateOutput = (KernelFlags & MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT) != 0; - const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; - const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; - - const float32x4_t ZeroVector = vdupq_n_f32(0.0f); - - // Convert byte strides to element strides - const size_t StrideWidthElements = StrideWidth / sizeof(float); - const size_t DilationWidthElements = DilationWidth / sizeof(float); - const size_t InputStrideElements = InputStride / sizeof(float); - const size_t FilterStrideElements = FilterStride / sizeof(float); - const size_t OutputStrideElements = OutputStride / sizeof(float); - const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); - - // Mark unused variables in NCHWc kernel - (void)InputStrideElements; - - // Process all filters - for (size_t f = 0; f < FilterCount; f++) { - const float* filter = Filter + f * FilterStrideElements; - float* output = Output + f * OutputStrideElements; - const float bias = BiasAddition ? Bias[f] : 0.0f; - const float32x4_t BiasVector = vdupq_n_f32(bias); - - const float* input_ptr = Input; - size_t OutputIndex = 0; - - // Handle left padding - for (size_t i = 0; i < OutputCountLeftPad; i++) { - float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; - - for (size_t kh = 0; kh < KernelHeight; kh++) { - for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float filter_value = filter[kh * KernelWidth + kw]; - const float32x4_t FilterVector = vdupq_n_f32(filter_value); - const float32x4_t InputVector = MlasLoadFloat32x4(input_element); - - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); - } - } - - if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); - } - - MlasStoreFloat32x4(&output[OutputIndex], Accumulator); - OutputIndex += 4; - input_ptr += StrideWidthElements; - } - - // Handle main output region - for (size_t i = 0; i < OutputCount; i++) { - float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; - - for (size_t kh = 0; kh < KernelHeight; kh++) { - for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float filter_value = filter[kh * KernelWidth + kw]; - const float32x4_t FilterVector = vdupq_n_f32(filter_value); - const float32x4_t InputVector = MlasLoadFloat32x4(input_element); - - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); - } - } - - if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); - } - - MlasStoreFloat32x4(&output[OutputIndex], Accumulator); - OutputIndex += 4; - input_ptr += StrideWidthElements; - } - - // Handle right padding - for (size_t i = 0; i < OutputCountRightPad; i++) { - float32x4_t Accumulator = AccumulateOutput ? MlasLoadFloat32x4(&output[OutputIndex]) : BiasVector; - - for (size_t kh = 0; kh < KernelHeight; kh++) { - for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_element = input_ptr + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float filter_value = filter[kh * KernelWidth + kw]; - const float32x4_t FilterVector = vdupq_n_f32(filter_value); - const float32x4_t InputVector = MlasLoadFloat32x4(input_element); - - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); - } - } - - if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); - } - - MlasStoreFloat32x4(&output[OutputIndex], Accumulator); - OutputIndex += 4; - input_ptr += StrideWidthElements; - } - } + } // From b505bd64df8039bed001778f47b32724f31e25b9 Mon Sep 17 00:00:00 2001 From: Rohan Date: Tue, 8 Jul 2025 19:32:07 +0000 Subject: [PATCH 06/26] Add placeholder kernel for MlasConvNchwcFloatKernelNeon --- .../core/mlas/lib/sconv_kernel_neon.cpp | 102 +++++++++++++++++- 1 file changed, 99 insertions(+), 3 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index 36e25ee51232d..c31f1e9b800bb 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -48,14 +48,33 @@ void unsigned KernelFlags ) { - + (void)Input; + (void)Filter; + (void)Output; + (void)StrideWidth; + (void)DilationWidth; + (void)FilterCount; + (void)InputStride; + (void)FilterStride; + (void)OutputStride; + (void)KernelHeight; + (void)KernelWidth; + (void)InputBase; + (void)InputWidth; + (void)DilatedInputWidth; + (void)OutputCountLeftPad; + (void)OutputCount; + (void)OutputCountRightPad; + (void)Bias; + (void)KernelFlags; } // // Implementation of MlasConvNchwcFloatKernelNeon // // This kernel performs 2D convolution optimized for NCHWc format. -// It processes multiple filters in blocks to improve cache efficiency. +// NCHWc format organizes data in blocks where both input and output channels +// are blocked. The filter is organized as [OutputChannelBlocks][InputChannelBlocks][KH][KW][BlockSize][BlockSize]. // void @@ -82,7 +101,84 @@ void unsigned KernelFlags ) { - + const bool AccumulateOutput = (KernelFlags & MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT) != 0; + const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; + const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; + + const size_t BlockSize = MlasNchwcGetBlockSize(); + + const size_t StrideWidthElements = StrideWidth / sizeof(float); + const size_t DilationWidthElements = DilationWidth / sizeof(float); + const size_t FilterStrideElements = FilterStride / sizeof(float); + const size_t OutputStrideElements = OutputStride / sizeof(float); + const size_t InputWidthElements = InputWidth / sizeof(float); + const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); + + (void)InputStride; + + const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; + + for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { + bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); + + for (size_t filterSetBlock = 0; filterSetBlock < FilterCount; filterSetBlock++) { + const float* filter = Filter + filterSetBlock * FilterStrideElements; + float* output = Output + filterSetBlock * OutputStrideElements; + + float accumulator[BlockSize]; + + for (size_t i = 0; i < BlockSize; i++) { + accumulator[i] = 0.0f; + } + if (AccumulateOutput) { + for (size_t i = 0; i < BlockSize; i++) { + accumulator[i] = output[output_idx * BlockSize + i]; + } + } + if (BiasAddition) { + for (size_t i = 0; i < BlockSize; i++) { + accumulator[i] += Bias[filterSetBlock * BlockSize + i]; + } + } + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_base = Input + output_idx * StrideWidthElements + + kh * DilatedInputWidthElements + kw * DilationWidthElements; + + for (size_t filterBlock = 0; filterBlock < BlockSize; filterBlock++) { + for (size_t ic = 0; ic < BlockSize; ic++) { + size_t kernel_pos = kh * (KernelWidth * BlockSize * BlockSize) + + kw * (BlockSize * BlockSize) + + filterBlock * (BlockSize) + + ic; + + const float* input_element = input_base + ic; + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; + + float input_value; + if (is_main_region && input_element >= input_row_start && input_element < input_row_end) { + input_value = *input_element; + } else { + input_value = 0.0f; + } + + float filter_value = filter[kernel_pos]; + accumulator[filterBlock] += input_value * filter_value; + } + } + } + } + + for (size_t i = 0; i < BlockSize; i++) { + if (ReluActivation && accumulator[i] < 0.0f) { + accumulator[i] = 0.0f; + } + output[output_idx * BlockSize + i] = accumulator[i]; + } + } + } } // From 790cc7eda8361244d3180090ab69ef600e3ca6ab Mon Sep 17 00:00:00 2001 From: Rohan Date: Tue, 8 Jul 2025 22:22:48 +0000 Subject: [PATCH 07/26] Fix MlasConvNchwcFloatKernelNeon --- onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index c31f1e9b800bb..ea380150034ee 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -153,21 +153,22 @@ void filterBlock * (BlockSize) + ic; - const float* input_element = input_base + ic; + const float* input_element = input_base + filterBlock; const float* input_row_start = InputBase + kh * DilatedInputWidthElements; const float* input_row_end = input_row_start + InputWidthElements; float input_value; - if (is_main_region && input_element >= input_row_start && input_element < input_row_end) { + if (is_main_region || (input_element >= input_row_start && input_element < input_row_end)) { input_value = *input_element; } else { input_value = 0.0f; } float filter_value = filter[kernel_pos]; - accumulator[filterBlock] += input_value * filter_value; + accumulator[ic] += input_value * filter_value; } } + } } From 906393af28e9b4b00fb5be2c991fbf35a423fa2a Mon Sep 17 00:00:00 2001 From: Rohan Date: Tue, 8 Jul 2025 22:31:43 +0000 Subject: [PATCH 08/26] Use MLAS intrinsics for MlasConvNchwcFloatKernelNeon --- .../core/mlas/lib/sconv_kernel_neon.cpp | 66 +++++++++---------- 1 file changed, 31 insertions(+), 35 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index ea380150034ee..ef16487becb5d 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -106,6 +106,7 @@ void const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; const size_t BlockSize = MlasNchwcGetBlockSize(); + const float32x4_t ZeroVector = vdupq_n_f32(0.0f); const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); @@ -125,20 +126,17 @@ void const float* filter = Filter + filterSetBlock * FilterStrideElements; float* output = Output + filterSetBlock * OutputStrideElements; - float accumulator[BlockSize]; + float32x4_t Accumulator; - for (size_t i = 0; i < BlockSize; i++) { - accumulator[i] = 0.0f; - } if (AccumulateOutput) { - for (size_t i = 0; i < BlockSize; i++) { - accumulator[i] = output[output_idx * BlockSize + i]; - } + Accumulator = MlasLoadFloat32x4(&output[output_idx * BlockSize]); + } else { + Accumulator = vdupq_n_f32(0.0f); } + if (BiasAddition) { - for (size_t i = 0; i < BlockSize; i++) { - accumulator[i] += Bias[filterSetBlock * BlockSize + i]; - } + const float32x4_t BiasVector = MlasLoadFloat32x4(&Bias[filterSetBlock * BlockSize]); + Accumulator = vaddq_f32(Accumulator, BiasVector); } for (size_t kh = 0; kh < KernelHeight; kh++) { @@ -147,37 +145,35 @@ void kh * DilatedInputWidthElements + kw * DilationWidthElements; for (size_t filterBlock = 0; filterBlock < BlockSize; filterBlock++) { - for (size_t ic = 0; ic < BlockSize; ic++) { - size_t kernel_pos = kh * (KernelWidth * BlockSize * BlockSize) + - kw * (BlockSize * BlockSize) + - filterBlock * (BlockSize) + - ic; - - const float* input_element = input_base + filterBlock; - const float* input_row_start = InputBase + kh * DilatedInputWidthElements; - const float* input_row_end = input_row_start + InputWidthElements; - - float input_value; - if (is_main_region || (input_element >= input_row_start && input_element < input_row_end)) { - input_value = *input_element; - } else { - input_value = 0.0f; - } - - float filter_value = filter[kernel_pos]; - accumulator[ic] += input_value * filter_value; + const float* input_element = input_base + filterBlock; + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; + + float input_value; + if (is_main_region || (input_element >= input_row_start && input_element < input_row_end)) { + input_value = *input_element; + } else { + input_value = 0.0f; } - } + const float32x4_t InputVector = vdupq_n_f32(input_value); + + size_t kernel_base_pos = kh * (KernelWidth * BlockSize * BlockSize) + + kw * (BlockSize * BlockSize) + + filterBlock * BlockSize; + + const float32x4_t FilterVector = MlasLoadFloat32x4(&filter[kernel_base_pos]); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } } } - for (size_t i = 0; i < BlockSize; i++) { - if (ReluActivation && accumulator[i] < 0.0f) { - accumulator[i] = 0.0f; - } - output[output_idx * BlockSize + i] = accumulator[i]; + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); } + + MlasStoreFloat32x4(&output[output_idx * BlockSize], Accumulator); } } } From 4d322e643e5ae64e1b6a9e0d1ccc7a0983245e69 Mon Sep 17 00:00:00 2001 From: Rohan Date: Wed, 9 Jul 2025 16:33:01 +0000 Subject: [PATCH 09/26] Add MlasConvNchwFloatKernelNeon --- .../core/mlas/lib/sconv_kernel_neon.cpp | 91 ++++++++++++++----- 1 file changed, 69 insertions(+), 22 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index ef16487becb5d..bc27e0b413ead 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -48,34 +48,81 @@ void unsigned KernelFlags ) { - (void)Input; - (void)Filter; - (void)Output; - (void)StrideWidth; - (void)DilationWidth; - (void)FilterCount; + const bool AccumulateOutput = (KernelFlags & MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT) != 0; + const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; + const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; + + const size_t BlockSize = MlasNchwcGetBlockSize(); + const float32x4_t ZeroVector = vdupq_n_f32(0.0f); + + const size_t StrideWidthElements = StrideWidth / sizeof(float); + const size_t DilationWidthElements = DilationWidth / sizeof(float); + const size_t FilterStrideElements = FilterStride / sizeof(float); + const size_t OutputStrideElements = OutputStride / sizeof(float); + const size_t InputWidthElements = InputWidth / sizeof(float); + const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); + (void)InputStride; - (void)FilterStride; - (void)OutputStride; - (void)KernelHeight; - (void)KernelWidth; - (void)InputBase; - (void)InputWidth; - (void)DilatedInputWidth; - (void)OutputCountLeftPad; - (void)OutputCount; - (void)OutputCountRightPad; - (void)Bias; - (void)KernelFlags; + + const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; + + for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { + bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); + + for (size_t filterSetBlock = 0; filterSetBlock < FilterCount; filterSetBlock++) { + const float* filter = Filter + filterSetBlock * FilterStrideElements; + float* output = Output + filterSetBlock * OutputStrideElements; + + float32x4_t Accumulator; + + if (AccumulateOutput) { + Accumulator = MlasLoadFloat32x4(&output[output_idx * BlockSize]); + } else { + Accumulator = vdupq_n_f32(0.0f); + } + + if (BiasAddition) { + const float32x4_t BiasVector = MlasLoadFloat32x4(&Bias[filterSetBlock * BlockSize]); + Accumulator = vaddq_f32(Accumulator, BiasVector); + } + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_base = Input + output_idx * StrideWidthElements + + kh * DilatedInputWidthElements + kw * DilationWidthElements; + + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; + + float input_value; + if (is_main_region || (input_base >= input_row_start && input_base < input_row_end)) { + input_value = *input_base; + } else { + input_value = 0.0f; + } + + const float32x4_t InputVector = vdupq_n_f32(input_value); + + size_t kernel_base_pos = kh * KernelWidth + kw; + + const float32x4_t FilterVector = MlasLoadFloat32x4(&filter[kernel_base_pos * BlockSize]); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } + } + + if (ReluActivation) { + Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + } + + MlasStoreFloat32x4(&output[output_idx * BlockSize], Accumulator); + } + } } // // Implementation of MlasConvNchwcFloatKernelNeon // -// This kernel performs 2D convolution optimized for NCHWc format. -// NCHWc format organizes data in blocks where both input and output channels -// are blocked. The filter is organized as [OutputChannelBlocks][InputChannelBlocks][KH][KW][BlockSize][BlockSize]. -// void MLASCALL From cb06a1a47e807b46b044317a3e8e29a24b04e141 Mon Sep 17 00:00:00 2001 From: Rohan Date: Thu, 10 Jul 2025 21:17:00 +0000 Subject: [PATCH 10/26] Add placeholder NCHWc Pool --- cmake/onnxruntime_mlas.cmake | 1 + onnxruntime/core/mlas/lib/mlasi.h | 6 +- onnxruntime/core/mlas/lib/platform.cpp | 3 + onnxruntime/core/mlas/lib/snchwc.cpp | 10 +- onnxruntime/core/mlas/lib/spool.h | 130 +++++++++++++++++ .../core/mlas/lib/spool_kernel_neon.cpp | 132 ++++++++++++++++++ 6 files changed, 274 insertions(+), 8 deletions(-) create mode 100644 onnxruntime/core/mlas/lib/spool.h create mode 100644 onnxruntime/core/mlas/lib/spool_kernel_neon.cpp diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index 80445dd7746be..c8cb173f635c0 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -450,6 +450,7 @@ else() ${MLAS_SRC_DIR}/qgemm_kernel_ummla.cpp ${MLAS_SRC_DIR}/sbgemm_kernel_neon.cpp ${MLAS_SRC_DIR}/sconv_kernel_neon.cpp + ${MLAS_SRC_DIR}/spool_kernel_neon.cpp ${MLAS_SRC_DIR}/cast_kernel_neon.cpp ${MLAS_SRC_DIR}/hqnbitgemm_kernel_neon_fp16.cpp ${MLAS_SRC_DIR}/rotary_embedding_kernel_neon_fp16.cpp diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h index 033968da3e189..4a9b2abca5a73 100644 --- a/onnxruntime/core/mlas/lib/mlasi.h +++ b/onnxruntime/core/mlas/lib/mlasi.h @@ -952,7 +952,10 @@ extern "C" { MLAS_CONV_FLOAT_KERNEL MlasConvNchwFloatKernelNeon; MLAS_CONV_FLOAT_KERNEL MlasConvNchwcFloatKernelNeon; MLAS_CONV_DEPTHWISE_FLOAT_KERNEL MlasConvDepthwiseFloatKernelNeon; - MLAS_CONV_POINTWISE_FLOAT_KERNEL MlasConvPointwiseFloatKernelNeon; + MLAS_CONV_POINTWISE_FLOAT_KERNEL MlasConvPointwiseFloatKernelNeon; + MLAS_POOL_FLOAT_KERNEL MlasPoolMaximumFloatKernelNeon; + MLAS_POOL_FLOAT_KERNEL MlasPoolAverageExcludePadFloatKernelNeon; + MLAS_POOL_FLOAT_KERNEL MlasPoolAverageIncludePadFloatKernelNeon; #endif MLAS_GEMM_DOUBLE_KERNEL MlasDgemmKernelZero; MLAS_GEMM_DOUBLE_KERNEL MlasDgemmKernelAdd; @@ -1342,6 +1345,7 @@ struct MLAS_PLATFORM { MLAS_CONV_FLOAT_KERNEL* ConvNchwcFloatKernel; MLAS_CONV_DEPTHWISE_FLOAT_KERNEL* ConvDepthwiseFloatKernel; MLAS_CONV_POINTWISE_FLOAT_KERNEL* ConvPointwiseFloatKernel; + MLAS_POOL_FLOAT_KERNEL* PoolFloatKernel[MlasPoolingKindCount]; uint32_t NchwcBlockSize; #endif const MLAS_SYMM_QGEMM_DISPATCH* SymmQgemmDispatch{nullptr}; diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp index 00dfc86b5eaf1..9e1b0906eeb93 100644 --- a/onnxruntime/core/mlas/lib/platform.cpp +++ b/onnxruntime/core/mlas/lib/platform.cpp @@ -562,6 +562,9 @@ Return Value: this->ConvNchwcFloatKernel = MlasConvNchwcFloatKernelNeon; this->ConvDepthwiseFloatKernel = MlasConvDepthwiseFloatKernelNeon; this->ConvPointwiseFloatKernel = MlasConvPointwiseFloatKernelNeon; + this->PoolFloatKernel[MlasMaximumPooling] = MlasPoolMaximumFloatKernelNeon; + this->PoolFloatKernel[MlasAveragePoolingExcludePad] = MlasPoolAverageExcludePadFloatKernelNeon; + this->PoolFloatKernel[MlasAveragePoolingIncludePad] = MlasPoolAverageIncludePadFloatKernelNeon; this->NchwcBlockSize = 4; // What is it supposed to be? // diff --git a/onnxruntime/core/mlas/lib/snchwc.cpp b/onnxruntime/core/mlas/lib/snchwc.cpp index c36a4a8fac71a..2fc27d6d4ad7f 100644 --- a/onnxruntime/core/mlas/lib/snchwc.cpp +++ b/onnxruntime/core/mlas/lib/snchwc.cpp @@ -1093,7 +1093,7 @@ struct MLAS_NCHWC_CONV_DEPTHWISE_ALGORITHM : MLAS_NCHWC_CONV_ALGORITHM struct MLAS_NCHWC_POOL_ALGORITHM : MLAS_NCHWC_NN_ALGORITHM { -#if !defined(MLAS_TARGET_AMD64) && !defined(MLAS_TARGET_LARCH64) +#if !defined(MLAS_TARGET_AMD64) && !defined(MLAS_TARGET_LARCH64) && !defined(MLAS_TARGET_ARM64) static MLAS_POOL_FLOAT_KERNEL* const PoolKernels[]; #endif @@ -1131,7 +1131,7 @@ struct MLAS_NCHWC_POOL_ALGORITHM : MLAS_NCHWC_NN_ALGORITHM const size_t DilatedInputWidthBytes = BlockSize * DilationHeight * InputWidth * sizeof(float); const size_t InputStrideBytes = DilatedInputWidthBytes - KernelWidth * DilationWidthBytes; -#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) +#if defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_LARCH64) || defined(MLAS_TARGET_ARM64) MLAS_POOL_FLOAT_KERNEL* Kernel = GetMlasPlatform().PoolFloatKernel[WorkBlock->PoolingKind]; #else MLAS_POOL_FLOAT_KERNEL* Kernel = PoolKernels[WorkBlock->PoolingKind]; @@ -1197,7 +1197,7 @@ struct MLAS_NCHWC_POOL_ALGORITHM : MLAS_NCHWC_NN_ALGORITHM } }; -#if !defined(MLAS_TARGET_AMD64) && !defined(MLAS_TARGET_LARCH64) +#if !defined(MLAS_TARGET_AMD64) && !defined(MLAS_TARGET_LARCH64) && !defined(MLAS_TARGET_ARM64) MLAS_POOL_FLOAT_KERNEL* const MLAS_NCHWC_POOL_ALGORITHM::PoolKernels[] = { @@ -1788,10 +1788,6 @@ MlasConvPointwiseFloatKernel( MLAS_UNREFERENCED_PARAMETER(Flags); } -#endif - -#if !defined(MLAS_TARGET_AMD64) && !defined(MLAS_TARGET_LARCH64) - void MLASCALL MlasPoolMaximumFloatKernel( diff --git a/onnxruntime/core/mlas/lib/spool.h b/onnxruntime/core/mlas/lib/spool.h new file mode 100644 index 0000000000000..99e2d0ec8be7e --- /dev/null +++ b/onnxruntime/core/mlas/lib/spool.h @@ -0,0 +1,130 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + spool.h + +Abstract: + + This module contains the private data structures and procedure prototypes + for the single precision pooling operation. + +--*/ + +#pragma once + +#include + +// +// Define the calling convention for MLAS functions. +// + +#ifndef MLASCALL +#if defined(_WIN32) && !defined(_WIN64) +#define MLASCALL __stdcall +#else +#define MLASCALL +#endif +#endif + +// +// Define the prototypes of the NEON convolution kernels. +// + +#if defined(__aarch64__) || defined(_M_ARM64) + +extern "C" { + +void +MLASCALL +MlasConvNchwFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t FilterCount, + size_t InputStride, + size_t FilterStride, + size_t OutputStride, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + const float* Bias, + unsigned KernelFlags + ); + +void +MLASCALL +MlasConvNchwcFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t FilterCount, + size_t InputStride, + size_t FilterStride, + size_t OutputStride, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + const float* Bias, + unsigned KernelFlags + ); + +void +MLASCALL +MlasConvDepthwiseFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t InputStride, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + const float* Bias, + unsigned KernelFlags + ); + +void +MLASCALL +MlasConvPointwiseFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t InputChannels, + size_t FilterCount, + size_t InputStride, + size_t FilterStride, + size_t OutputStride, + size_t OutputCount, + const float* Bias, + unsigned KernelFlags + ); + +} + +#endif diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp new file mode 100644 index 0000000000000..17e0c3b76a0d7 --- /dev/null +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -0,0 +1,132 @@ +/*++ + +Copyright (c) Microsoft Corporation. All rights reserved. + +Licensed under the MIT License. + +Module Name: + + spool_kernel_neon.cpp + +Abstract: + + This module implements the single precision pooling kernels for ARM NEON. + +--*/ + +// #include "spool.h" + +#if defined(__aarch64__) || defined(_M_ARM64) + +#include +#include + +#include "arm_neon.h" +#include "mlasi.h" + +void + MLASCALL + MlasPoolMaximumFloatKernelNeon( + const float* Input, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t InputStride, + size_t ActualKernelSize, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad + ) +{ + MLAS_UNREFERENCED_PARAMETER(Input); + MLAS_UNREFERENCED_PARAMETER(Output); + MLAS_UNREFERENCED_PARAMETER(StrideWidth); + MLAS_UNREFERENCED_PARAMETER(DilationWidth); + MLAS_UNREFERENCED_PARAMETER(InputStride); + MLAS_UNREFERENCED_PARAMETER(ActualKernelSize); + MLAS_UNREFERENCED_PARAMETER(KernelHeight); + MLAS_UNREFERENCED_PARAMETER(KernelWidth); + MLAS_UNREFERENCED_PARAMETER(InputBase); + MLAS_UNREFERENCED_PARAMETER(InputWidth); + MLAS_UNREFERENCED_PARAMETER(DilatedInputWidth); + MLAS_UNREFERENCED_PARAMETER(OutputCountLeftPad); + MLAS_UNREFERENCED_PARAMETER(OutputCount); + MLAS_UNREFERENCED_PARAMETER(OutputCountRightPad); +} + +void + MLASCALL + MlasPoolAverageExcludePadFloatKernelNeon( + const float* Input, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t InputStride, + size_t ActualKernelSize, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad + ) +{ + MLAS_UNREFERENCED_PARAMETER(Input); + MLAS_UNREFERENCED_PARAMETER(Output); + MLAS_UNREFERENCED_PARAMETER(StrideWidth); + MLAS_UNREFERENCED_PARAMETER(DilationWidth); + MLAS_UNREFERENCED_PARAMETER(InputStride); + MLAS_UNREFERENCED_PARAMETER(ActualKernelSize); + MLAS_UNREFERENCED_PARAMETER(KernelHeight); + MLAS_UNREFERENCED_PARAMETER(KernelWidth); + MLAS_UNREFERENCED_PARAMETER(InputBase); + MLAS_UNREFERENCED_PARAMETER(InputWidth); + MLAS_UNREFERENCED_PARAMETER(DilatedInputWidth); + MLAS_UNREFERENCED_PARAMETER(OutputCountLeftPad); + MLAS_UNREFERENCED_PARAMETER(OutputCount); + MLAS_UNREFERENCED_PARAMETER(OutputCountRightPad); +} + +void + MLASCALL + MlasPoolAverageIncludePadFloatKernelNeon( + const float* Input, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t InputStride, + size_t ActualKernelSize, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad + ) +{ + MLAS_UNREFERENCED_PARAMETER(Input); + MLAS_UNREFERENCED_PARAMETER(Output); + MLAS_UNREFERENCED_PARAMETER(StrideWidth); + MLAS_UNREFERENCED_PARAMETER(DilationWidth); + MLAS_UNREFERENCED_PARAMETER(InputStride); + MLAS_UNREFERENCED_PARAMETER(ActualKernelSize); + MLAS_UNREFERENCED_PARAMETER(KernelHeight); + MLAS_UNREFERENCED_PARAMETER(KernelWidth); + MLAS_UNREFERENCED_PARAMETER(InputBase); + MLAS_UNREFERENCED_PARAMETER(InputWidth); + MLAS_UNREFERENCED_PARAMETER(DilatedInputWidth); + MLAS_UNREFERENCED_PARAMETER(OutputCountLeftPad); + MLAS_UNREFERENCED_PARAMETER(OutputCount); + MLAS_UNREFERENCED_PARAMETER(OutputCountRightPad); +} + +#endif // __aarch64__ || _M_ARM64 \ No newline at end of file From 00caa4c1a2682d38de13d933813e589a6e2990e6 Mon Sep 17 00:00:00 2001 From: Rohan Date: Thu, 10 Jul 2025 21:48:16 +0000 Subject: [PATCH 11/26] Vanilla C++ implementation --- .../core/mlas/lib/spool_kernel_neon.cpp | 185 ++++++++++++++---- 1 file changed, 146 insertions(+), 39 deletions(-) diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp index 17e0c3b76a0d7..202a81cf99c00 100644 --- a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -20,6 +20,7 @@ Module Name: #include #include +#include #include "arm_neon.h" #include "mlasi.h" @@ -43,20 +44,52 @@ void size_t OutputCountRightPad ) { - MLAS_UNREFERENCED_PARAMETER(Input); - MLAS_UNREFERENCED_PARAMETER(Output); - MLAS_UNREFERENCED_PARAMETER(StrideWidth); - MLAS_UNREFERENCED_PARAMETER(DilationWidth); - MLAS_UNREFERENCED_PARAMETER(InputStride); MLAS_UNREFERENCED_PARAMETER(ActualKernelSize); - MLAS_UNREFERENCED_PARAMETER(KernelHeight); - MLAS_UNREFERENCED_PARAMETER(KernelWidth); - MLAS_UNREFERENCED_PARAMETER(InputBase); - MLAS_UNREFERENCED_PARAMETER(InputWidth); - MLAS_UNREFERENCED_PARAMETER(DilatedInputWidth); - MLAS_UNREFERENCED_PARAMETER(OutputCountLeftPad); - MLAS_UNREFERENCED_PARAMETER(OutputCount); - MLAS_UNREFERENCED_PARAMETER(OutputCountRightPad); + MLAS_UNREFERENCED_PARAMETER(InputStride); + + constexpr size_t BlockSize = 4; + const size_t StrideWidthElements = StrideWidth / sizeof(float); + const size_t DilationWidthElements = DilationWidth / sizeof(float); + const size_t InputWidthElements = InputWidth / sizeof(float); + const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); + + const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; + + for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { + // Initialize maximum values to negative infinity + float max_values[BlockSize]; + for (size_t i = 0; i < BlockSize; i++) { + max_values[i] = std::numeric_limits::lowest(); + } + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_base = Input + output_idx * StrideWidthElements + + kh * DilatedInputWidthElements + kw * DilationWidthElements; + + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; + + for (size_t block_idx = 0; block_idx < BlockSize; block_idx++) { + const float* input_element = input_base + block_idx; + + float input_value; + if (input_element >= input_row_start && input_element < input_row_end) { + input_value = *input_element; + } else { + input_value = std::numeric_limits::lowest(); + } + + max_values[block_idx] = std::max(max_values[block_idx], input_value); + } + } + } + + // Store the results + for (size_t i = 0; i < BlockSize; i++) { + Output[output_idx * BlockSize + i] = max_values[i]; + } + } } void @@ -78,20 +111,55 @@ void size_t OutputCountRightPad ) { - MLAS_UNREFERENCED_PARAMETER(Input); - MLAS_UNREFERENCED_PARAMETER(Output); - MLAS_UNREFERENCED_PARAMETER(StrideWidth); - MLAS_UNREFERENCED_PARAMETER(DilationWidth); - MLAS_UNREFERENCED_PARAMETER(InputStride); MLAS_UNREFERENCED_PARAMETER(ActualKernelSize); - MLAS_UNREFERENCED_PARAMETER(KernelHeight); - MLAS_UNREFERENCED_PARAMETER(KernelWidth); - MLAS_UNREFERENCED_PARAMETER(InputBase); - MLAS_UNREFERENCED_PARAMETER(InputWidth); - MLAS_UNREFERENCED_PARAMETER(DilatedInputWidth); - MLAS_UNREFERENCED_PARAMETER(OutputCountLeftPad); - MLAS_UNREFERENCED_PARAMETER(OutputCount); - MLAS_UNREFERENCED_PARAMETER(OutputCountRightPad); + MLAS_UNREFERENCED_PARAMETER(InputStride); + + constexpr size_t BlockSize = 4; + const size_t StrideWidthElements = StrideWidth / sizeof(float); + const size_t DilationWidthElements = DilationWidth / sizeof(float); + const size_t InputWidthElements = InputWidth / sizeof(float); + const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); + + const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; + + for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { + // Initialize sum values and count + float sum_values[BlockSize]; + size_t valid_count[BlockSize]; + for (size_t i = 0; i < BlockSize; i++) { + sum_values[i] = 0.0f; + valid_count[i] = 0; + } + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_base = Input + output_idx * StrideWidthElements + + kh * DilatedInputWidthElements + kw * DilationWidthElements; + + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; + + for (size_t block_idx = 0; block_idx < BlockSize; block_idx++) { + const float* input_element = input_base + block_idx; + + if (input_element >= input_row_start && input_element < input_row_end) { + float input_value = *input_element; + sum_values[block_idx] += input_value; + valid_count[block_idx]++; + } + } + } + } + + // Store the results (average excluding padding) + for (size_t i = 0; i < BlockSize; i++) { + if (valid_count[i] > 0) { + Output[output_idx * BlockSize + i] = sum_values[i] / static_cast(valid_count[i]); + } else { + Output[output_idx * BlockSize + i] = 0.0f; + } + } + } } void @@ -113,20 +181,59 @@ void size_t OutputCountRightPad ) { - MLAS_UNREFERENCED_PARAMETER(Input); - MLAS_UNREFERENCED_PARAMETER(Output); - MLAS_UNREFERENCED_PARAMETER(StrideWidth); - MLAS_UNREFERENCED_PARAMETER(DilationWidth); MLAS_UNREFERENCED_PARAMETER(InputStride); - MLAS_UNREFERENCED_PARAMETER(ActualKernelSize); - MLAS_UNREFERENCED_PARAMETER(KernelHeight); - MLAS_UNREFERENCED_PARAMETER(KernelWidth); - MLAS_UNREFERENCED_PARAMETER(InputBase); - MLAS_UNREFERENCED_PARAMETER(InputWidth); - MLAS_UNREFERENCED_PARAMETER(DilatedInputWidth); - MLAS_UNREFERENCED_PARAMETER(OutputCountLeftPad); - MLAS_UNREFERENCED_PARAMETER(OutputCount); - MLAS_UNREFERENCED_PARAMETER(OutputCountRightPad); + + constexpr size_t BlockSize = 4; + const size_t StrideWidthElements = StrideWidth / sizeof(float); + const size_t DilationWidthElements = DilationWidth / sizeof(float); + const size_t InputWidthElements = InputWidth / sizeof(float); + const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); + + const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; + + // Use ActualKernelSize if provided, otherwise compute from dimensions + const float KernelSize = (ActualKernelSize > 0) ? + static_cast(ActualKernelSize) : + static_cast(KernelHeight * KernelWidth); + + for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { + bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); + + // Initialize sum values + float sum_values[BlockSize]; + for (size_t i = 0; i < BlockSize; i++) { + sum_values[i] = 0.0f; + } + + for (size_t kh = 0; kh < KernelHeight; kh++) { + for (size_t kw = 0; kw < KernelWidth; kw++) { + const float* input_ptr = Input + output_idx * StrideWidthElements + + kh * DilatedInputWidthElements + kw * DilationWidthElements; + + // Check bounds for this kernel position + const float* row_start = InputBase + kh * DilatedInputWidthElements; + const float* row_end = row_start + InputWidthElements; + + for (size_t block_idx = 0; block_idx < BlockSize; block_idx++) { + const float* element_ptr = input_ptr + block_idx; + + float value; + if (is_main_region || (element_ptr >= row_start && element_ptr < row_end)) { + value = *element_ptr; + } else { + value = 0.0f; // Padding values are treated as 0 + } + + sum_values[block_idx] += value; + } + } + } + + // Store the results (divide by total kernel size for include pad) + for (size_t i = 0; i < BlockSize; i++) { + Output[output_idx * BlockSize + i] = sum_values[i] / KernelSize; + } + } } #endif // __aarch64__ || _M_ARM64 \ No newline at end of file From 4cead5ec263956556c56ff866a047cb381247161 Mon Sep 17 00:00:00 2001 From: Rohan Date: Thu, 10 Jul 2025 22:26:35 +0000 Subject: [PATCH 12/26] Intrinsics for Pooling --- .../core/mlas/lib/spool_kernel_neon.cpp | 202 +++++++++++------- 1 file changed, 119 insertions(+), 83 deletions(-) diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp index 202a81cf99c00..7c7ecd6e0755a 100644 --- a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -14,15 +14,12 @@ Module Name: --*/ -// #include "spool.h" - #if defined(__aarch64__) || defined(_M_ARM64) -#include #include #include +#include -#include "arm_neon.h" #include "mlasi.h" void @@ -46,8 +43,8 @@ void { MLAS_UNREFERENCED_PARAMETER(ActualKernelSize); MLAS_UNREFERENCED_PARAMETER(InputStride); - - constexpr size_t BlockSize = 4; + + const size_t BlockSize = MlasNchwcGetBlockSize(); const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); const size_t InputWidthElements = InputWidth / sizeof(float); @@ -55,40 +52,46 @@ void const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; + // Initialize negative infinity vector for out-of-bounds values using MLAS intrinsics + const MLAS_FLOAT32X4 NegInfVector = MlasBroadcastFloat32x4(-std::numeric_limits::infinity()); + for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { - // Initialize maximum values to negative infinity - float max_values[BlockSize]; - for (size_t i = 0; i < BlockSize; i++) { - max_values[i] = std::numeric_limits::lowest(); - } + // Initialize maximum values to negative infinity using MLAS intrinsics + MLAS_FLOAT32X4 MaxVector = NegInfVector; for (size_t kh = 0; kh < KernelHeight; kh++) { for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_base = Input + output_idx * StrideWidthElements + - kh * DilatedInputWidthElements + kw * DilationWidthElements; - - const float* input_row_start = InputBase + kh * DilatedInputWidthElements; - const float* input_row_end = input_row_start + InputWidthElements; - - for (size_t block_idx = 0; block_idx < BlockSize; block_idx++) { - const float* input_element = input_base + block_idx; - - float input_value; - if (input_element >= input_row_start && input_element < input_row_end) { - input_value = *input_element; - } else { - input_value = std::numeric_limits::lowest(); - } + const float* input_ptr = Input + output_idx * StrideWidthElements + + kh * DilatedInputWidthElements + kw * DilationWidthElements; + + const float* row_start = InputBase + kh * DilatedInputWidthElements; + const float* row_end = row_start + InputWidthElements; - max_values[block_idx] = std::max(max_values[block_idx], input_value); + MLAS_FLOAT32X4 InputVector; + if (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end) { + // All elements are within bounds, load directly using MLAS intrinsics + InputVector = MlasLoadFloat32x4(input_ptr); + } else { + // Some elements might be out of bounds, load individually + std::vector values(BlockSize); + for (size_t i = 0; i < BlockSize; i++) { + const float* element_ptr = input_ptr + i; + if (element_ptr >= row_start && element_ptr < row_end) { + values[i] = *element_ptr; + } else { + values[i] = -std::numeric_limits::infinity(); + } + } + InputVector = MlasLoadFloat32x4(values.data()); } + + // Update maximum using MLAS intrinsics + MaxVector = MlasMaximumFloat32x4(MaxVector, InputVector); } } - // Store the results - for (size_t i = 0; i < BlockSize; i++) { - Output[output_idx * BlockSize + i] = max_values[i]; - } + // Store the results using MLAS intrinsics + MlasStoreFloat32x4(&Output[output_idx * BlockSize], MaxVector); } } @@ -113,8 +116,8 @@ void { MLAS_UNREFERENCED_PARAMETER(ActualKernelSize); MLAS_UNREFERENCED_PARAMETER(InputStride); - - constexpr size_t BlockSize = 4; + + const size_t BlockSize = MlasNchwcGetBlockSize(); const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); const size_t InputWidthElements = InputWidth / sizeof(float); @@ -122,43 +125,68 @@ void const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; + // Initialize zero vector using MLAS intrinsics + const MLAS_FLOAT32X4 ZeroVector = MlasZeroFloat32x4(); + for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { - // Initialize sum values and count - float sum_values[BlockSize]; - size_t valid_count[BlockSize]; - for (size_t i = 0; i < BlockSize; i++) { - sum_values[i] = 0.0f; - valid_count[i] = 0; - } + // Initialize sum vector using MLAS intrinsics + MLAS_FLOAT32X4 SumVector = ZeroVector; + + // Track valid count for each element (allocate dynamically for variable block size) + std::vector valid_count(BlockSize, 0); for (size_t kh = 0; kh < KernelHeight; kh++) { for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_base = Input + output_idx * StrideWidthElements + - kh * DilatedInputWidthElements + kw * DilationWidthElements; - - const float* input_row_start = InputBase + kh * DilatedInputWidthElements; - const float* input_row_end = input_row_start + InputWidthElements; - - for (size_t block_idx = 0; block_idx < BlockSize; block_idx++) { - const float* input_element = input_base + block_idx; - - if (input_element >= input_row_start && input_element < input_row_end) { - float input_value = *input_element; - sum_values[block_idx] += input_value; - valid_count[block_idx]++; + const float* input_ptr = Input + output_idx * StrideWidthElements + + kh * DilatedInputWidthElements + kw * DilationWidthElements; + + const float* row_start = InputBase + kh * DilatedInputWidthElements; + const float* row_end = row_start + InputWidthElements; + + MLAS_FLOAT32X4 InputVector; + + if (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end) { + // All elements are within bounds, load directly using MLAS intrinsics + InputVector = MlasLoadFloat32x4(input_ptr); + // All elements are valid + for (size_t i = 0; i < BlockSize; i++) { + valid_count[i]++; + } + } else { + // Some elements might be out of bounds, handle individually + std::vector values(BlockSize); + for (size_t i = 0; i < BlockSize; i++) { + const float* element_ptr = input_ptr + i; + if (element_ptr >= row_start && element_ptr < row_end) { + values[i] = *element_ptr; + valid_count[i]++; + } else { + values[i] = 0.0f; + } } + InputVector = MlasLoadFloat32x4(values.data()); } + + // Add to sum using MLAS intrinsics + SumVector = MlasAddFloat32x4(SumVector, InputVector); } } - // Store the results (average excluding padding) + // Compute average by dividing by valid count for each element + std::vector results(BlockSize); + MlasStoreFloat32x4(results.data(), SumVector); + for (size_t i = 0; i < BlockSize; i++) { if (valid_count[i] > 0) { - Output[output_idx * BlockSize + i] = sum_values[i] / static_cast(valid_count[i]); + results[i] = results[i] / static_cast(valid_count[i]); } else { - Output[output_idx * BlockSize + i] = 0.0f; + results[i] = 0.0f; } } + + // Store the results using MLAS intrinsics + MLAS_FLOAT32X4 ResultVector = MlasLoadFloat32x4(results.data()); + MlasStoreFloat32x4(&Output[output_idx * BlockSize], ResultVector); } } @@ -182,57 +210,65 @@ void ) { MLAS_UNREFERENCED_PARAMETER(InputStride); - - constexpr size_t BlockSize = 4; + + const size_t BlockSize = MlasNchwcGetBlockSize(); const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); const size_t InputWidthElements = InputWidth / sizeof(float); const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; - - // Use ActualKernelSize if provided, otherwise compute from dimensions - const float KernelSize = (ActualKernelSize > 0) ? - static_cast(ActualKernelSize) : - static_cast(KernelHeight * KernelWidth); + + // Use ActualKernelSize as provided by the caller + const float KernelSize = static_cast(ActualKernelSize); + + // Initialize vectors using MLAS intrinsics + const MLAS_FLOAT32X4 ZeroVector = MlasZeroFloat32x4(); + const MLAS_FLOAT32X4 KernelSizeVector = MlasBroadcastFloat32x4(KernelSize); for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); - // Initialize sum values - float sum_values[BlockSize]; - for (size_t i = 0; i < BlockSize; i++) { - sum_values[i] = 0.0f; - } + // Initialize sum vector using MLAS intrinsics + MLAS_FLOAT32X4 SumVector = ZeroVector; for (size_t kh = 0; kh < KernelHeight; kh++) { for (size_t kw = 0; kw < KernelWidth; kw++) { const float* input_ptr = Input + output_idx * StrideWidthElements + kh * DilatedInputWidthElements + kw * DilationWidthElements; - // Check bounds for this kernel position const float* row_start = InputBase + kh * DilatedInputWidthElements; const float* row_end = row_start + InputWidthElements; - for (size_t block_idx = 0; block_idx < BlockSize; block_idx++) { - const float* element_ptr = input_ptr + block_idx; - - float value; - if (is_main_region || (element_ptr >= row_start && element_ptr < row_end)) { - value = *element_ptr; - } else { - value = 0.0f; // Padding values are treated as 0 + MLAS_FLOAT32X4 InputVector; + + if (is_main_region || (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end)) { + // All elements are within bounds or in main region, load directly using MLAS intrinsics + InputVector = MlasLoadFloat32x4(input_ptr); + } else { + // Some elements might be out of bounds, handle individually + std::vector values(BlockSize); + for (size_t i = 0; i < BlockSize; i++) { + const float* element_ptr = input_ptr + i; + if (is_main_region || (element_ptr >= row_start && element_ptr < row_end)) { + values[i] = *element_ptr; + } else { + values[i] = 0.0f; // Padding values are treated as 0 + } } - - sum_values[block_idx] += value; + InputVector = MlasLoadFloat32x4(values.data()); } + + // Add to sum using MLAS intrinsics + SumVector = MlasAddFloat32x4(SumVector, InputVector); } } - // Store the results (divide by total kernel size for include pad) - for (size_t i = 0; i < BlockSize; i++) { - Output[output_idx * BlockSize + i] = sum_values[i] / KernelSize; - } + // Compute average by dividing by kernel size using MLAS intrinsics + MLAS_FLOAT32X4 ResultVector = MlasDivideFloat32x4(SumVector, KernelSizeVector); + + // Store the results using MLAS intrinsics + MlasStoreFloat32x4(&Output[output_idx * BlockSize], ResultVector); } } From abd54916096dc54b12fc7fc453d0b4b49961a248 Mon Sep 17 00:00:00 2001 From: Rohan Date: Thu, 10 Jul 2025 22:35:21 +0000 Subject: [PATCH 13/26] Refactored to share code --- .../core/mlas/lib/spool_kernel_neon.cpp | 200 +++++++++--------- 1 file changed, 102 insertions(+), 98 deletions(-) diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp index 7c7ecd6e0755a..460846a77f800 100644 --- a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -95,28 +95,25 @@ void } } -void - MLASCALL - MlasPoolAverageExcludePadFloatKernelNeon( - const float* Input, - float* Output, - size_t StrideWidth, - size_t DilationWidth, - size_t InputStride, - size_t ActualKernelSize, - size_t KernelHeight, - size_t KernelWidth, - const float* InputBase, - size_t InputWidth, - size_t DilatedInputWidth, - size_t OutputCountLeftPad, - size_t OutputCount, - size_t OutputCountRightPad - ) +// Helper function for average pooling kernels +static void +MlasPoolAverageFloatKernelNeonImpl( + const float* Input, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t ActualKernelSize, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + bool ExcludePad +) { - MLAS_UNREFERENCED_PARAMETER(ActualKernelSize); - MLAS_UNREFERENCED_PARAMETER(InputStride); - const size_t BlockSize = MlasNchwcGetBlockSize(); const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); @@ -128,12 +125,24 @@ void // Initialize zero vector using MLAS intrinsics const MLAS_FLOAT32X4 ZeroVector = MlasZeroFloat32x4(); + // For include pad, prepare kernel size vector + MLAS_FLOAT32X4 KernelSizeVector; + if (!ExcludePad) { + const float KernelSize = static_cast(ActualKernelSize); + KernelSizeVector = MlasBroadcastFloat32x4(KernelSize); + } + for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { + bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); + // Initialize sum vector using MLAS intrinsics MLAS_FLOAT32X4 SumVector = ZeroVector; - // Track valid count for each element (allocate dynamically for variable block size) - std::vector valid_count(BlockSize, 0); + // Track valid count for each element (only needed for exclude pad) + std::vector valid_count; + if (ExcludePad) { + valid_count.resize(BlockSize, 0); + } for (size_t kh = 0; kh < KernelHeight; kh++) { for (size_t kw = 0; kw < KernelWidth; kw++) { @@ -145,23 +154,34 @@ void MLAS_FLOAT32X4 InputVector; - if (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end) { - // All elements are within bounds, load directly using MLAS intrinsics + // Determine if we can do a fast vector load + bool can_fast_load = ExcludePad ? (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end) : (is_main_region || (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end)); + + if (can_fast_load) { + // All elements are within bounds (or in main region for include pad) InputVector = MlasLoadFloat32x4(input_ptr); - // All elements are valid - for (size_t i = 0; i < BlockSize; i++) { - valid_count[i]++; + + if (ExcludePad) { + // All elements are valid for exclude pad + for (size_t i = 0; i < BlockSize; i++) { + valid_count[i]++; + } } } else { // Some elements might be out of bounds, handle individually std::vector values(BlockSize); for (size_t i = 0; i < BlockSize; i++) { const float* element_ptr = input_ptr + i; - if (element_ptr >= row_start && element_ptr < row_end) { + + bool is_valid = ExcludePad ? (element_ptr >= row_start && element_ptr < row_end) : (is_main_region || (element_ptr >= row_start && element_ptr < row_end)); + + if (is_valid) { values[i] = *element_ptr; - valid_count[i]++; + if (ExcludePad) { + valid_count[i]++; + } } else { - values[i] = 0.0f; + values[i] = 0.0f; // Padding values are treated as 0 } } InputVector = MlasLoadFloat32x4(values.data()); @@ -172,27 +192,35 @@ void } } - // Compute average by dividing by valid count for each element - std::vector results(BlockSize); - MlasStoreFloat32x4(results.data(), SumVector); + // Compute final result based on pooling type + MLAS_FLOAT32X4 ResultVector; + if (ExcludePad) { + // Exclude pad: divide by valid count for each element + std::vector results(BlockSize); + MlasStoreFloat32x4(results.data(), SumVector); - for (size_t i = 0; i < BlockSize; i++) { - if (valid_count[i] > 0) { - results[i] = results[i] / static_cast(valid_count[i]); - } else { - results[i] = 0.0f; + for (size_t i = 0; i < BlockSize; i++) { + if (valid_count[i] > 0) { + results[i] = results[i] / static_cast(valid_count[i]); + } else { + results[i] = 0.0f; + } } + + ResultVector = MlasLoadFloat32x4(results.data()); + } else { + // Include pad: divide by total kernel size + ResultVector = MlasDivideFloat32x4(SumVector, KernelSizeVector); } // Store the results using MLAS intrinsics - MLAS_FLOAT32X4 ResultVector = MlasLoadFloat32x4(results.data()); MlasStoreFloat32x4(&Output[output_idx * BlockSize], ResultVector); } } void MLASCALL - MlasPoolAverageIncludePadFloatKernelNeon( + MlasPoolAverageExcludePadFloatKernelNeon( const float* Input, float* Output, size_t StrideWidth, @@ -211,65 +239,41 @@ void { MLAS_UNREFERENCED_PARAMETER(InputStride); - const size_t BlockSize = MlasNchwcGetBlockSize(); - const size_t StrideWidthElements = StrideWidth / sizeof(float); - const size_t DilationWidthElements = DilationWidth / sizeof(float); - const size_t InputWidthElements = InputWidth / sizeof(float); - const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); - - const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; - - // Use ActualKernelSize as provided by the caller - const float KernelSize = static_cast(ActualKernelSize); - - // Initialize vectors using MLAS intrinsics - const MLAS_FLOAT32X4 ZeroVector = MlasZeroFloat32x4(); - const MLAS_FLOAT32X4 KernelSizeVector = MlasBroadcastFloat32x4(KernelSize); - - for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { - bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); - - // Initialize sum vector using MLAS intrinsics - MLAS_FLOAT32X4 SumVector = ZeroVector; - - for (size_t kh = 0; kh < KernelHeight; kh++) { - for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_ptr = Input + output_idx * StrideWidthElements + - kh * DilatedInputWidthElements + kw * DilationWidthElements; - - const float* row_start = InputBase + kh * DilatedInputWidthElements; - const float* row_end = row_start + InputWidthElements; - - MLAS_FLOAT32X4 InputVector; - - if (is_main_region || (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end)) { - // All elements are within bounds or in main region, load directly using MLAS intrinsics - InputVector = MlasLoadFloat32x4(input_ptr); - } else { - // Some elements might be out of bounds, handle individually - std::vector values(BlockSize); - for (size_t i = 0; i < BlockSize; i++) { - const float* element_ptr = input_ptr + i; - if (is_main_region || (element_ptr >= row_start && element_ptr < row_end)) { - values[i] = *element_ptr; - } else { - values[i] = 0.0f; // Padding values are treated as 0 - } - } - InputVector = MlasLoadFloat32x4(values.data()); - } - - // Add to sum using MLAS intrinsics - SumVector = MlasAddFloat32x4(SumVector, InputVector); - } - } + MlasPoolAverageFloatKernelNeonImpl( + Input, Output, StrideWidth, DilationWidth, ActualKernelSize, + KernelHeight, KernelWidth, InputBase, InputWidth, DilatedInputWidth, + OutputCountLeftPad, OutputCount, OutputCountRightPad, + true // ExcludePad = true + ); +} - // Compute average by dividing by kernel size using MLAS intrinsics - MLAS_FLOAT32X4 ResultVector = MlasDivideFloat32x4(SumVector, KernelSizeVector); +void + MLASCALL + MlasPoolAverageIncludePadFloatKernelNeon( + const float* Input, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t InputStride, + size_t ActualKernelSize, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad + ) +{ + MLAS_UNREFERENCED_PARAMETER(InputStride); - // Store the results using MLAS intrinsics - MlasStoreFloat32x4(&Output[output_idx * BlockSize], ResultVector); - } + MlasPoolAverageFloatKernelNeonImpl( + Input, Output, StrideWidth, DilationWidth, ActualKernelSize, + KernelHeight, KernelWidth, InputBase, InputWidth, DilatedInputWidth, + OutputCountLeftPad, OutputCount, OutputCountRightPad, + false // ExcludePad = false + ); } #endif // __aarch64__ || _M_ARM64 \ No newline at end of file From 74e0e3b0fcf9e6562f6c354d9fb1969b60eaee4b Mon Sep 17 00:00:00 2001 From: Rohan Date: Fri, 11 Jul 2025 15:01:12 +0000 Subject: [PATCH 14/26] Format file & delete unused header --- onnxruntime/core/mlas/lib/spool.h | 130 ------------------ .../core/mlas/lib/spool_kernel_neon.cpp | 22 +-- 2 files changed, 1 insertion(+), 151 deletions(-) delete mode 100644 onnxruntime/core/mlas/lib/spool.h diff --git a/onnxruntime/core/mlas/lib/spool.h b/onnxruntime/core/mlas/lib/spool.h deleted file mode 100644 index 99e2d0ec8be7e..0000000000000 --- a/onnxruntime/core/mlas/lib/spool.h +++ /dev/null @@ -1,130 +0,0 @@ -/*++ - -Copyright (c) Microsoft Corporation. All rights reserved. - -Licensed under the MIT License. - -Module Name: - - spool.h - -Abstract: - - This module contains the private data structures and procedure prototypes - for the single precision pooling operation. - ---*/ - -#pragma once - -#include - -// -// Define the calling convention for MLAS functions. -// - -#ifndef MLASCALL -#if defined(_WIN32) && !defined(_WIN64) -#define MLASCALL __stdcall -#else -#define MLASCALL -#endif -#endif - -// -// Define the prototypes of the NEON convolution kernels. -// - -#if defined(__aarch64__) || defined(_M_ARM64) - -extern "C" { - -void -MLASCALL -MlasConvNchwFloatKernelNeon( - const float* Input, - const float* Filter, - float* Output, - size_t StrideWidth, - size_t DilationWidth, - size_t FilterCount, - size_t InputStride, - size_t FilterStride, - size_t OutputStride, - size_t KernelHeight, - size_t KernelWidth, - const float* InputBase, - size_t InputWidth, - size_t DilatedInputWidth, - size_t OutputCountLeftPad, - size_t OutputCount, - size_t OutputCountRightPad, - const float* Bias, - unsigned KernelFlags - ); - -void -MLASCALL -MlasConvNchwcFloatKernelNeon( - const float* Input, - const float* Filter, - float* Output, - size_t StrideWidth, - size_t DilationWidth, - size_t FilterCount, - size_t InputStride, - size_t FilterStride, - size_t OutputStride, - size_t KernelHeight, - size_t KernelWidth, - const float* InputBase, - size_t InputWidth, - size_t DilatedInputWidth, - size_t OutputCountLeftPad, - size_t OutputCount, - size_t OutputCountRightPad, - const float* Bias, - unsigned KernelFlags - ); - -void -MLASCALL -MlasConvDepthwiseFloatKernelNeon( - const float* Input, - const float* Filter, - float* Output, - size_t StrideWidth, - size_t DilationWidth, - size_t InputStride, - size_t KernelHeight, - size_t KernelWidth, - const float* InputBase, - size_t InputWidth, - size_t DilatedInputWidth, - size_t OutputCountLeftPad, - size_t OutputCount, - size_t OutputCountRightPad, - const float* Bias, - unsigned KernelFlags - ); - -void -MLASCALL -MlasConvPointwiseFloatKernelNeon( - const float* Input, - const float* Filter, - float* Output, - size_t StrideWidth, - size_t InputChannels, - size_t FilterCount, - size_t InputStride, - size_t FilterStride, - size_t OutputStride, - size_t OutputCount, - const float* Bias, - unsigned KernelFlags - ); - -} - -#endif diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp index 460846a77f800..28293f49013f5 100644 --- a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -52,11 +52,9 @@ void const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; - // Initialize negative infinity vector for out-of-bounds values using MLAS intrinsics const MLAS_FLOAT32X4 NegInfVector = MlasBroadcastFloat32x4(-std::numeric_limits::infinity()); for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { - // Initialize maximum values to negative infinity using MLAS intrinsics MLAS_FLOAT32X4 MaxVector = NegInfVector; for (size_t kh = 0; kh < KernelHeight; kh++) { @@ -69,10 +67,8 @@ void MLAS_FLOAT32X4 InputVector; if (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end) { - // All elements are within bounds, load directly using MLAS intrinsics InputVector = MlasLoadFloat32x4(input_ptr); } else { - // Some elements might be out of bounds, load individually std::vector values(BlockSize); for (size_t i = 0; i < BlockSize; i++) { const float* element_ptr = input_ptr + i; @@ -85,17 +81,14 @@ void InputVector = MlasLoadFloat32x4(values.data()); } - // Update maximum using MLAS intrinsics MaxVector = MlasMaximumFloat32x4(MaxVector, InputVector); } } - // Store the results using MLAS intrinsics MlasStoreFloat32x4(&Output[output_idx * BlockSize], MaxVector); } } -// Helper function for average pooling kernels static void MlasPoolAverageFloatKernelNeonImpl( const float* Input, @@ -122,10 +115,8 @@ MlasPoolAverageFloatKernelNeonImpl( const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; - // Initialize zero vector using MLAS intrinsics const MLAS_FLOAT32X4 ZeroVector = MlasZeroFloat32x4(); - // For include pad, prepare kernel size vector MLAS_FLOAT32X4 KernelSizeVector; if (!ExcludePad) { const float KernelSize = static_cast(ActualKernelSize); @@ -135,10 +126,8 @@ MlasPoolAverageFloatKernelNeonImpl( for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); - // Initialize sum vector using MLAS intrinsics MLAS_FLOAT32X4 SumVector = ZeroVector; - // Track valid count for each element (only needed for exclude pad) std::vector valid_count; if (ExcludePad) { valid_count.resize(BlockSize, 0); @@ -154,21 +143,17 @@ MlasPoolAverageFloatKernelNeonImpl( MLAS_FLOAT32X4 InputVector; - // Determine if we can do a fast vector load bool can_fast_load = ExcludePad ? (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end) : (is_main_region || (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end)); if (can_fast_load) { - // All elements are within bounds (or in main region for include pad) InputVector = MlasLoadFloat32x4(input_ptr); if (ExcludePad) { - // All elements are valid for exclude pad for (size_t i = 0; i < BlockSize; i++) { valid_count[i]++; } } } else { - // Some elements might be out of bounds, handle individually std::vector values(BlockSize); for (size_t i = 0; i < BlockSize; i++) { const float* element_ptr = input_ptr + i; @@ -181,21 +166,18 @@ MlasPoolAverageFloatKernelNeonImpl( valid_count[i]++; } } else { - values[i] = 0.0f; // Padding values are treated as 0 + values[i] = 0.0f; } } InputVector = MlasLoadFloat32x4(values.data()); } - // Add to sum using MLAS intrinsics SumVector = MlasAddFloat32x4(SumVector, InputVector); } } - // Compute final result based on pooling type MLAS_FLOAT32X4 ResultVector; if (ExcludePad) { - // Exclude pad: divide by valid count for each element std::vector results(BlockSize); MlasStoreFloat32x4(results.data(), SumVector); @@ -209,11 +191,9 @@ MlasPoolAverageFloatKernelNeonImpl( ResultVector = MlasLoadFloat32x4(results.data()); } else { - // Include pad: divide by total kernel size ResultVector = MlasDivideFloat32x4(SumVector, KernelSizeVector); } - // Store the results using MLAS intrinsics MlasStoreFloat32x4(&Output[output_idx * BlockSize], ResultVector); } } From 16be947d07e8c6709fd5753af1c0793f6b6fc05a Mon Sep 17 00:00:00 2001 From: Rohan Date: Fri, 11 Jul 2025 19:27:25 +0000 Subject: [PATCH 15/26] Minor modifications to pass more tests --- .../core/mlas/lib/spool_kernel_neon.cpp | 49 +++++++------------ 1 file changed, 17 insertions(+), 32 deletions(-) diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp index 28293f49013f5..ecfb83142f7b6 100644 --- a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -49,24 +49,25 @@ void const size_t DilationWidthElements = DilationWidth / sizeof(float); const size_t InputWidthElements = InputWidth / sizeof(float); const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); - const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; - const MLAS_FLOAT32X4 NegInfVector = MlasBroadcastFloat32x4(-std::numeric_limits::infinity()); + const float MaxPaddingValue = std::numeric_limits::lowest(); + const MLAS_FLOAT32X4 MaxPaddingVector = MlasBroadcastFloat32x4(MaxPaddingValue); for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { - MLAS_FLOAT32X4 MaxVector = NegInfVector; + MLAS_FLOAT32X4 MaxVector = MaxPaddingVector; for (size_t kh = 0; kh < KernelHeight; kh++) { + const float* row_start = InputBase + kh * DilatedInputWidthElements; + const float* row_end = row_start + InputWidthElements; + for (size_t kw = 0; kw < KernelWidth; kw++) { const float* input_ptr = Input + output_idx * StrideWidthElements + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float* row_start = InputBase + kh * DilatedInputWidthElements; - const float* row_end = row_start + InputWidthElements; - MLAS_FLOAT32X4 InputVector; - if (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end) { + + if (input_ptr >= row_start && (input_ptr + BlockSize) <= row_end) { InputVector = MlasLoadFloat32x4(input_ptr); } else { std::vector values(BlockSize); @@ -75,7 +76,7 @@ void if (element_ptr >= row_start && element_ptr < row_end) { values[i] = *element_ptr; } else { - values[i] = -std::numeric_limits::infinity(); + values[i] = MaxPaddingValue; } } InputVector = MlasLoadFloat32x4(values.data()); @@ -112,20 +113,11 @@ MlasPoolAverageFloatKernelNeonImpl( const size_t DilationWidthElements = DilationWidth / sizeof(float); const size_t InputWidthElements = InputWidth / sizeof(float); const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); - const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; const MLAS_FLOAT32X4 ZeroVector = MlasZeroFloat32x4(); - MLAS_FLOAT32X4 KernelSizeVector; - if (!ExcludePad) { - const float KernelSize = static_cast(ActualKernelSize); - KernelSizeVector = MlasBroadcastFloat32x4(KernelSize); - } - for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { - bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); - MLAS_FLOAT32X4 SumVector = ZeroVector; std::vector valid_count; @@ -134,18 +126,16 @@ MlasPoolAverageFloatKernelNeonImpl( } for (size_t kh = 0; kh < KernelHeight; kh++) { + const float* row_start = InputBase + kh * DilatedInputWidthElements; + const float* row_end = row_start + InputWidthElements; + for (size_t kw = 0; kw < KernelWidth; kw++) { const float* input_ptr = Input + output_idx * StrideWidthElements + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float* row_start = InputBase + kh * DilatedInputWidthElements; - const float* row_end = row_start + InputWidthElements; - MLAS_FLOAT32X4 InputVector; - bool can_fast_load = ExcludePad ? (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end) : (is_main_region || (input_ptr >= row_start && (input_ptr + BlockSize - 1) < row_end)); - - if (can_fast_load) { + if (input_ptr >= row_start && (input_ptr + BlockSize) <= row_end) { InputVector = MlasLoadFloat32x4(input_ptr); if (ExcludePad) { @@ -157,10 +147,7 @@ MlasPoolAverageFloatKernelNeonImpl( std::vector values(BlockSize); for (size_t i = 0; i < BlockSize; i++) { const float* element_ptr = input_ptr + i; - - bool is_valid = ExcludePad ? (element_ptr >= row_start && element_ptr < row_end) : (is_main_region || (element_ptr >= row_start && element_ptr < row_end)); - - if (is_valid) { + if (element_ptr >= row_start && element_ptr < row_end) { values[i] = *element_ptr; if (ExcludePad) { valid_count[i]++; @@ -182,15 +169,13 @@ MlasPoolAverageFloatKernelNeonImpl( MlasStoreFloat32x4(results.data(), SumVector); for (size_t i = 0; i < BlockSize; i++) { - if (valid_count[i] > 0) { - results[i] = results[i] / static_cast(valid_count[i]); - } else { - results[i] = 0.0f; - } + results[i] = results[i] / static_cast(valid_count[i]); } ResultVector = MlasLoadFloat32x4(results.data()); } else { + const float KernelSize = static_cast(ActualKernelSize); + const MLAS_FLOAT32X4 KernelSizeVector = MlasBroadcastFloat32x4(KernelSize); ResultVector = MlasDivideFloat32x4(SumVector, KernelSizeVector); } From f7d971d3007ae299d8377c5284af752fc829596b Mon Sep 17 00:00:00 2001 From: Rohan Date: Tue, 15 Jul 2025 18:09:34 +0000 Subject: [PATCH 16/26] Remove unnecessary code & formatting changes --- onnxruntime/core/mlas/lib/sconv.h | 106 +----------------- .../core/mlas/lib/sconv_kernel_neon.cpp | 46 ++++---- 2 files changed, 28 insertions(+), 124 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sconv.h b/onnxruntime/core/mlas/lib/sconv.h index 9cdd45abdc9ec..d9f29f0492423 100644 --- a/onnxruntime/core/mlas/lib/sconv.h +++ b/onnxruntime/core/mlas/lib/sconv.h @@ -35,105 +35,7 @@ Module Name: // Define the convolution kernel flags. // -#define MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT 0x00000001 -#define MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION 0x00000002 -#define MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION 0x00000004 -#define MLAS_CONV_KERNEL_FLAG_OTHER_ACTIVATION 0x00000008 - -// -// Define the prototypes of the NEON convolution kernels. -// - -#if defined(__aarch64__) || defined(_M_ARM64) - -extern "C" { - -void -MLASCALL -MlasConvNchwFloatKernelNeon( - const float* Input, - const float* Filter, - float* Output, - size_t StrideWidth, - size_t DilationWidth, - size_t FilterCount, - size_t InputStride, - size_t FilterStride, - size_t OutputStride, - size_t KernelHeight, - size_t KernelWidth, - const float* InputBase, - size_t InputWidth, - size_t DilatedInputWidth, - size_t OutputCountLeftPad, - size_t OutputCount, - size_t OutputCountRightPad, - const float* Bias, - unsigned KernelFlags - ); - -void -MLASCALL -MlasConvNchwcFloatKernelNeon( - const float* Input, - const float* Filter, - float* Output, - size_t StrideWidth, - size_t DilationWidth, - size_t FilterCount, - size_t InputStride, - size_t FilterStride, - size_t OutputStride, - size_t KernelHeight, - size_t KernelWidth, - const float* InputBase, - size_t InputWidth, - size_t DilatedInputWidth, - size_t OutputCountLeftPad, - size_t OutputCount, - size_t OutputCountRightPad, - const float* Bias, - unsigned KernelFlags - ); - -void -MLASCALL -MlasConvDepthwiseFloatKernelNeon( - const float* Input, - const float* Filter, - float* Output, - size_t StrideWidth, - size_t DilationWidth, - size_t InputStride, - size_t KernelHeight, - size_t KernelWidth, - const float* InputBase, - size_t InputWidth, - size_t DilatedInputWidth, - size_t OutputCountLeftPad, - size_t OutputCount, - size_t OutputCountRightPad, - const float* Bias, - unsigned KernelFlags - ); - -void -MLASCALL -MlasConvPointwiseFloatKernelNeon( - const float* Input, - const float* Filter, - float* Output, - size_t StrideWidth, - size_t InputChannels, - size_t FilterCount, - size_t InputStride, - size_t FilterStride, - size_t OutputStride, - size_t OutputCount, - const float* Bias, - unsigned KernelFlags - ); - -} - -#endif +#define MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT 0x00000001 +#define MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION 0x00000002 +#define MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION 0x00000004 +#define MLAS_CONV_KERNEL_FLAG_OTHER_ACTIVATION 0x00000008 \ No newline at end of file diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index bc27e0b413ead..69a46aad980c6 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -21,7 +21,6 @@ Module Name: #include #include -#include "arm_neon.h" #include "mlasi.h" void @@ -53,7 +52,7 @@ void const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; const size_t BlockSize = MlasNchwcGetBlockSize(); - const float32x4_t ZeroVector = vdupq_n_f32(0.0f); + const float32x4_t ZeroVector = MlasBroadcastFloat32x4(0.0f); const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); @@ -78,12 +77,12 @@ void if (AccumulateOutput) { Accumulator = MlasLoadFloat32x4(&output[output_idx * BlockSize]); } else { - Accumulator = vdupq_n_f32(0.0f); + Accumulator = MlasBroadcastFloat32x4(0.0f); } if (BiasAddition) { const float32x4_t BiasVector = MlasLoadFloat32x4(&Bias[filterSetBlock * BlockSize]); - Accumulator = vaddq_f32(Accumulator, BiasVector); + Accumulator = MlasAddFloat32x4(Accumulator, BiasVector); } for (size_t kh = 0; kh < KernelHeight; kh++) { @@ -101,7 +100,7 @@ void input_value = 0.0f; } - const float32x4_t InputVector = vdupq_n_f32(input_value); + const float32x4_t InputVector = MlasBroadcastFloat32x4(input_value); size_t kernel_base_pos = kh * KernelWidth + kw; @@ -153,7 +152,7 @@ void const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; const size_t BlockSize = MlasNchwcGetBlockSize(); - const float32x4_t ZeroVector = vdupq_n_f32(0.0f); + const float32x4_t ZeroVector = MlasBroadcastFloat32x4(0.0f); const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); @@ -178,12 +177,12 @@ void if (AccumulateOutput) { Accumulator = MlasLoadFloat32x4(&output[output_idx * BlockSize]); } else { - Accumulator = vdupq_n_f32(0.0f); + Accumulator = MlasBroadcastFloat32x4(0.0f); } if (BiasAddition) { const float32x4_t BiasVector = MlasLoadFloat32x4(&Bias[filterSetBlock * BlockSize]); - Accumulator = vaddq_f32(Accumulator, BiasVector); + Accumulator = MlasAddFloat32x4(Accumulator, BiasVector); } for (size_t kh = 0; kh < KernelHeight; kh++) { @@ -203,7 +202,7 @@ void input_value = 0.0f; } - const float32x4_t InputVector = vdupq_n_f32(input_value); + const float32x4_t InputVector = MlasBroadcastFloat32x4(input_value); size_t kernel_base_pos = kh * (KernelWidth * BlockSize * BlockSize) + kw * (BlockSize * BlockSize) + @@ -259,7 +258,7 @@ void const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; const size_t BlockSize = MlasNchwcGetBlockSize(); - const float32x4_t ZeroVector = vdupq_n_f32(0.0f); + const float32x4_t ZeroVector = MlasBroadcastFloat32x4(0.0f); const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); @@ -279,10 +278,11 @@ void if (AccumulateOutput) { Accumulator = MlasLoadFloat32x4(&Output[output_idx * BlockSize]); - } else if (BiasAddition) { - Accumulator = MlasLoadFloat32x4(Bias); } else { - Accumulator = vdupq_n_f32(0.0f); + Accumulator = MlasBroadcastFloat32x4(0.0f); + } + if (BiasAddition) { + Accumulator = MlasAddFloat32x4(Accumulator, MlasLoadFloat32x4(Bias)); } for (size_t kh = 0; kh < KernelHeight; kh++) { @@ -361,25 +361,27 @@ void const size_t OutputStrideElements = OutputStride / sizeof(float); const size_t BlockSize = MlasNchwcGetBlockSize(); - const float32x4_t ZeroVector = vdupq_n_f32(0.0f); + const float32x4_t ZeroVector = MlasBroadcastFloat32x4(0.0f); - for (size_t i = 0; i < OutputCount; i++) { + for (size_t output_idx = 0; output_idx < OutputCount; output_idx++) { for (size_t f = 0; f < FilterCount; f++) { const float* filter = Filter + f * FilterStrideElements; float* output = Output + f * OutputStrideElements; float32x4_t Accumulator; if (AccumulateOutput) { - Accumulator = MlasLoadFloat32x4(&output[i * BlockSize]); - } else if (BiasAddition) { - Accumulator = MlasLoadFloat32x4(&Bias[f * BlockSize]); + Accumulator = MlasLoadFloat32x4(&output[output_idx * BlockSize]); } else { - Accumulator = vdupq_n_f32(0.0f); + Accumulator = MlasBroadcastFloat32x4(0.0f); + } + if (BiasAddition) { + const float32x4_t BiasVector = MlasLoadFloat32x4(&Bias[f * BlockSize]); + Accumulator = MlasAddFloat32x4(Accumulator, BiasVector); } for (size_t c = 0; c < InputChannels; c++) { - const float* input_ptr = Input + c * InputStrideElements + i * StrideWidthElements; + const float* input_ptr = Input + c * InputStrideElements + output_idx * StrideWidthElements; for (size_t input_b = 0; input_b < BlockSize; input_b++) { const float input_value = input_ptr[input_b]; - const float32x4_t InputVector = vdupq_n_f32(input_value); + const float32x4_t InputVector = MlasBroadcastFloat32x4(input_value); const float* filter_ptr = filter + (c * BlockSize + input_b) * BlockSize; const float32x4_t FilterVector = MlasLoadFloat32x4(filter_ptr); Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); @@ -388,7 +390,7 @@ void if (ReluActivation) { Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); } - MlasStoreFloat32x4(&output[i * BlockSize], Accumulator); + MlasStoreFloat32x4(&output[output_idx * BlockSize], Accumulator); } } } From 0ff394cde8910b3e6205e862ca8f066621369fd6 Mon Sep 17 00:00:00 2001 From: Rohan Date: Tue, 15 Jul 2025 20:01:57 +0000 Subject: [PATCH 17/26] Refactor to share some code --- .../core/mlas/lib/sconv_kernel_neon.cpp | 196 ++++++++++-------- 1 file changed, 109 insertions(+), 87 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index 69a46aad980c6..61f28d096234a 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -23,9 +23,11 @@ Module Name: #include "mlasi.h" +// Common implementation for NCHW and NCHWC convolution kernels +template void MLASCALL - MlasConvNchwFloatKernelNeon( + MlasConvFloatKernelNeonImpl( const float* Input, const float* Filter, float* Output, @@ -90,23 +92,50 @@ void const float* input_base = Input + output_idx * StrideWidthElements + kh * DilatedInputWidthElements + kw * DilationWidthElements; - const float* input_row_start = InputBase + kh * DilatedInputWidthElements; - const float* input_row_end = input_row_start + InputWidthElements; + if (IsNchwcFormat) { + // NCHWC format - process each element in the block + for (size_t filterBlock = 0; filterBlock < BlockSize; filterBlock++) { + const float* input_element = input_base + filterBlock; + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; + + float input_value; + if (is_main_region || (input_element >= input_row_start && input_element < input_row_end)) { + input_value = *input_element; + } else { + input_value = 0.0f; + } + + const float32x4_t InputVector = MlasBroadcastFloat32x4(input_value); + + size_t kernel_base_pos = kh * (KernelWidth * BlockSize * BlockSize) + + kw * (BlockSize * BlockSize) + + filterBlock * BlockSize; - float input_value; - if (is_main_region || (input_base >= input_row_start && input_base < input_row_end)) { - input_value = *input_base; + const float32x4_t FilterVector = MlasLoadFloat32x4(&filter[kernel_base_pos]); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } } else { - input_value = 0.0f; - } + // NCHW format - simpler processing + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; - const float32x4_t InputVector = MlasBroadcastFloat32x4(input_value); + float input_value; + if (is_main_region || (input_base >= input_row_start && input_base < input_row_end)) { + input_value = *input_base; + } else { + input_value = 0.0f; + } - size_t kernel_base_pos = kh * KernelWidth + kw; + const float32x4_t InputVector = MlasBroadcastFloat32x4(input_value); - const float32x4_t FilterVector = MlasLoadFloat32x4(&filter[kernel_base_pos * BlockSize]); + size_t kernel_base_pos = kh * KernelWidth + kw; - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + const float32x4_t FilterVector = MlasLoadFloat32x4(&filter[kernel_base_pos * BlockSize]); + + Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + } } } @@ -119,6 +148,53 @@ void } } +void + MLASCALL + MlasConvNchwFloatKernelNeon( + const float* Input, + const float* Filter, + float* Output, + size_t StrideWidth, + size_t DilationWidth, + size_t FilterCount, + size_t InputStride, + size_t FilterStride, + size_t OutputStride, + size_t KernelHeight, + size_t KernelWidth, + const float* InputBase, + size_t InputWidth, + size_t DilatedInputWidth, + size_t OutputCountLeftPad, + size_t OutputCount, + size_t OutputCountRightPad, + const float* Bias, + unsigned KernelFlags + ) +{ + MlasConvFloatKernelNeonImpl( + Input, + Filter, + Output, + StrideWidth, + DilationWidth, + FilterCount, + InputStride, + FilterStride, + OutputStride, + KernelHeight, + KernelWidth, + InputBase, + InputWidth, + DilatedInputWidth, + OutputCountLeftPad, + OutputCount, + OutputCountRightPad, + Bias, + KernelFlags + ); +} + // // Implementation of MlasConvNchwcFloatKernelNeon // @@ -147,81 +223,27 @@ void unsigned KernelFlags ) { - const bool AccumulateOutput = (KernelFlags & MLAS_CONV_KERNEL_FLAG_ACCUMULATE_OUTPUT) != 0; - const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; - const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; - - const size_t BlockSize = MlasNchwcGetBlockSize(); - const float32x4_t ZeroVector = MlasBroadcastFloat32x4(0.0f); - - const size_t StrideWidthElements = StrideWidth / sizeof(float); - const size_t DilationWidthElements = DilationWidth / sizeof(float); - const size_t FilterStrideElements = FilterStride / sizeof(float); - const size_t OutputStrideElements = OutputStride / sizeof(float); - const size_t InputWidthElements = InputWidth / sizeof(float); - const size_t DilatedInputWidthElements = DilatedInputWidth / sizeof(float); - - (void)InputStride; - - const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; - - for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { - bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); - - for (size_t filterSetBlock = 0; filterSetBlock < FilterCount; filterSetBlock++) { - const float* filter = Filter + filterSetBlock * FilterStrideElements; - float* output = Output + filterSetBlock * OutputStrideElements; - - float32x4_t Accumulator; - - if (AccumulateOutput) { - Accumulator = MlasLoadFloat32x4(&output[output_idx * BlockSize]); - } else { - Accumulator = MlasBroadcastFloat32x4(0.0f); - } - - if (BiasAddition) { - const float32x4_t BiasVector = MlasLoadFloat32x4(&Bias[filterSetBlock * BlockSize]); - Accumulator = MlasAddFloat32x4(Accumulator, BiasVector); - } - - for (size_t kh = 0; kh < KernelHeight; kh++) { - for (size_t kw = 0; kw < KernelWidth; kw++) { - const float* input_base = Input + output_idx * StrideWidthElements + - kh * DilatedInputWidthElements + kw * DilationWidthElements; - - for (size_t filterBlock = 0; filterBlock < BlockSize; filterBlock++) { - const float* input_element = input_base + filterBlock; - const float* input_row_start = InputBase + kh * DilatedInputWidthElements; - const float* input_row_end = input_row_start + InputWidthElements; - - float input_value; - if (is_main_region || (input_element >= input_row_start && input_element < input_row_end)) { - input_value = *input_element; - } else { - input_value = 0.0f; - } - - const float32x4_t InputVector = MlasBroadcastFloat32x4(input_value); - - size_t kernel_base_pos = kh * (KernelWidth * BlockSize * BlockSize) + - kw * (BlockSize * BlockSize) + - filterBlock * BlockSize; - - const float32x4_t FilterVector = MlasLoadFloat32x4(&filter[kernel_base_pos]); - - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); - } - } - } - - if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); - } - - MlasStoreFloat32x4(&output[output_idx * BlockSize], Accumulator); - } - } + MlasConvFloatKernelNeonImpl( + Input, + Filter, + Output, + StrideWidth, + DilationWidth, + FilterCount, + InputStride, + FilterStride, + OutputStride, + KernelHeight, + KernelWidth, + InputBase, + InputWidth, + DilatedInputWidth, + OutputCountLeftPad, + OutputCount, + OutputCountRightPad, + Bias, + KernelFlags + ); } // From bd2b6c44fc672cbbf6a9f50cb5d99532582a9683 Mon Sep 17 00:00:00 2001 From: Rohan Date: Thu, 17 Jul 2025 20:36:35 +0000 Subject: [PATCH 18/26] Change block size to 16 --- onnxruntime/core/mlas/lib/platform.cpp | 2 +- .../core/mlas/lib/sconv_kernel_neon.cpp | 217 +++++++++++++++--- 2 files changed, 183 insertions(+), 36 deletions(-) diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp index 9e1b0906eeb93..7a84da0072c6c 100644 --- a/onnxruntime/core/mlas/lib/platform.cpp +++ b/onnxruntime/core/mlas/lib/platform.cpp @@ -565,7 +565,7 @@ Return Value: this->PoolFloatKernel[MlasMaximumPooling] = MlasPoolMaximumFloatKernelNeon; this->PoolFloatKernel[MlasAveragePoolingExcludePad] = MlasPoolAverageExcludePadFloatKernelNeon; this->PoolFloatKernel[MlasAveragePoolingIncludePad] = MlasPoolAverageIncludePadFloatKernelNeon; - this->NchwcBlockSize = 4; // What is it supposed to be? + this->NchwcBlockSize = 16; // What is it supposed to be? // // Check if the processor supports ASIMD dot product instructions. diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index 61f28d096234a..962ec1ba7057c 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -74,17 +74,30 @@ void const float* filter = Filter + filterSetBlock * FilterStrideElements; float* output = Output + filterSetBlock * OutputStrideElements; - float32x4_t Accumulator; + float32x4_t Accumulator0, Accumulator1, Accumulator2, Accumulator3; if (AccumulateOutput) { - Accumulator = MlasLoadFloat32x4(&output[output_idx * BlockSize]); + Accumulator0 = MlasLoadFloat32x4(&output[output_idx * BlockSize]); + Accumulator1 = MlasLoadFloat32x4(&output[output_idx * BlockSize + 4]); + Accumulator2 = MlasLoadFloat32x4(&output[output_idx * BlockSize + 8]); + Accumulator3 = MlasLoadFloat32x4(&output[output_idx * BlockSize + 12]); } else { - Accumulator = MlasBroadcastFloat32x4(0.0f); + Accumulator0 = MlasBroadcastFloat32x4(0.0f); + Accumulator1 = MlasBroadcastFloat32x4(0.0f); + Accumulator2 = MlasBroadcastFloat32x4(0.0f); + Accumulator3 = MlasBroadcastFloat32x4(0.0f); } if (BiasAddition) { - const float32x4_t BiasVector = MlasLoadFloat32x4(&Bias[filterSetBlock * BlockSize]); - Accumulator = MlasAddFloat32x4(Accumulator, BiasVector); + const float32x4_t BiasVector0 = MlasLoadFloat32x4(&Bias[filterSetBlock * BlockSize]); + const float32x4_t BiasVector1 = MlasLoadFloat32x4(&Bias[filterSetBlock * BlockSize + 4]); + const float32x4_t BiasVector2 = MlasLoadFloat32x4(&Bias[filterSetBlock * BlockSize + 8]); + const float32x4_t BiasVector3 = MlasLoadFloat32x4(&Bias[filterSetBlock * BlockSize + 12]); + + Accumulator0 = MlasAddFloat32x4(Accumulator0, BiasVector0); + Accumulator1 = MlasAddFloat32x4(Accumulator1, BiasVector1); + Accumulator2 = MlasAddFloat32x4(Accumulator2, BiasVector2); + Accumulator3 = MlasAddFloat32x4(Accumulator3, BiasVector3); } for (size_t kh = 0; kh < KernelHeight; kh++) { @@ -93,7 +106,6 @@ void kh * DilatedInputWidthElements + kw * DilationWidthElements; if (IsNchwcFormat) { - // NCHWC format - process each element in the block for (size_t filterBlock = 0; filterBlock < BlockSize; filterBlock++) { const float* input_element = input_base + filterBlock; const float* input_row_start = InputBase + kh * DilatedInputWidthElements; @@ -112,12 +124,17 @@ void kw * (BlockSize * BlockSize) + filterBlock * BlockSize; - const float32x4_t FilterVector = MlasLoadFloat32x4(&filter[kernel_base_pos]); + const float32x4_t FilterVector0 = MlasLoadFloat32x4(&filter[kernel_base_pos]); + const float32x4_t FilterVector1 = MlasLoadFloat32x4(&filter[kernel_base_pos + 4]); + const float32x4_t FilterVector2 = MlasLoadFloat32x4(&filter[kernel_base_pos + 8]); + const float32x4_t FilterVector3 = MlasLoadFloat32x4(&filter[kernel_base_pos + 12]); - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + Accumulator0 = MlasMultiplyAddFloat32x4(InputVector, FilterVector0, Accumulator0); + Accumulator1 = MlasMultiplyAddFloat32x4(InputVector, FilterVector1, Accumulator1); + Accumulator2 = MlasMultiplyAddFloat32x4(InputVector, FilterVector2, Accumulator2); + Accumulator3 = MlasMultiplyAddFloat32x4(InputVector, FilterVector3, Accumulator3); } } else { - // NCHW format - simpler processing const float* input_row_start = InputBase + kh * DilatedInputWidthElements; const float* input_row_end = input_row_start + InputWidthElements; @@ -132,18 +149,30 @@ void size_t kernel_base_pos = kh * KernelWidth + kw; - const float32x4_t FilterVector = MlasLoadFloat32x4(&filter[kernel_base_pos * BlockSize]); + const float32x4_t FilterVector0 = MlasLoadFloat32x4(&filter[kernel_base_pos * BlockSize]); + const float32x4_t FilterVector1 = MlasLoadFloat32x4(&filter[kernel_base_pos * BlockSize + 4]); + const float32x4_t FilterVector2 = MlasLoadFloat32x4(&filter[kernel_base_pos * BlockSize + 8]); + const float32x4_t FilterVector3 = MlasLoadFloat32x4(&filter[kernel_base_pos * BlockSize + 12]); - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + Accumulator0 = MlasMultiplyAddFloat32x4(InputVector, FilterVector0, Accumulator0); + Accumulator1 = MlasMultiplyAddFloat32x4(InputVector, FilterVector1, Accumulator1); + Accumulator2 = MlasMultiplyAddFloat32x4(InputVector, FilterVector2, Accumulator2); + Accumulator3 = MlasMultiplyAddFloat32x4(InputVector, FilterVector3, Accumulator3); } } } if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + Accumulator0 = MlasMaximumFloat32x4(Accumulator0, ZeroVector); + Accumulator1 = MlasMaximumFloat32x4(Accumulator1, ZeroVector); + Accumulator2 = MlasMaximumFloat32x4(Accumulator2, ZeroVector); + Accumulator3 = MlasMaximumFloat32x4(Accumulator3, ZeroVector); } - MlasStoreFloat32x4(&output[output_idx * BlockSize], Accumulator); + MlasStoreFloat32x4(&output[output_idx * BlockSize], Accumulator0); + MlasStoreFloat32x4(&output[output_idx * BlockSize + 4], Accumulator1); + MlasStoreFloat32x4(&output[output_idx * BlockSize + 8], Accumulator2); + MlasStoreFloat32x4(&output[output_idx * BlockSize + 12], Accumulator3); } } } @@ -296,15 +325,30 @@ void for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { bool is_main_region = (output_idx >= OutputCountLeftPad && output_idx < OutputCountLeftPad + OutputCount); - float32x4_t Accumulator; + float32x4_t Accumulator0, Accumulator1, Accumulator2, Accumulator3; if (AccumulateOutput) { - Accumulator = MlasLoadFloat32x4(&Output[output_idx * BlockSize]); + Accumulator0 = MlasLoadFloat32x4(&Output[output_idx * BlockSize]); + Accumulator1 = MlasLoadFloat32x4(&Output[output_idx * BlockSize + 4]); + Accumulator2 = MlasLoadFloat32x4(&Output[output_idx * BlockSize + 8]); + Accumulator3 = MlasLoadFloat32x4(&Output[output_idx * BlockSize + 12]); } else { - Accumulator = MlasBroadcastFloat32x4(0.0f); + Accumulator0 = MlasBroadcastFloat32x4(0.0f); + Accumulator1 = MlasBroadcastFloat32x4(0.0f); + Accumulator2 = MlasBroadcastFloat32x4(0.0f); + Accumulator3 = MlasBroadcastFloat32x4(0.0f); } + if (BiasAddition) { - Accumulator = MlasAddFloat32x4(Accumulator, MlasLoadFloat32x4(Bias)); + const float32x4_t BiasVector0 = MlasLoadFloat32x4(Bias); + const float32x4_t BiasVector1 = MlasLoadFloat32x4(Bias + 4); + const float32x4_t BiasVector2 = MlasLoadFloat32x4(Bias + 8); + const float32x4_t BiasVector3 = MlasLoadFloat32x4(Bias + 12); + + Accumulator0 = MlasAddFloat32x4(Accumulator0, BiasVector0); + Accumulator1 = MlasAddFloat32x4(Accumulator1, BiasVector1); + Accumulator2 = MlasAddFloat32x4(Accumulator2, BiasVector2); + Accumulator3 = MlasAddFloat32x4(Accumulator3, BiasVector3); } for (size_t kh = 0; kh < KernelHeight; kh++) { @@ -314,13 +358,12 @@ void const float* input_base = Input + output_idx * StrideWidthElements + kh * DilatedInputWidthElements + kw * DilationWidthElements; - float32x4_t InputVector; - + float32x4_t InputVector0; if (is_main_region) { - InputVector = MlasLoadFloat32x4(input_base); + InputVector0 = MlasLoadFloat32x4(input_base); } else { float input_values[4]; - for (size_t i = 0; i < BlockSize; i++) { + for (size_t i = 0; i < 4; i++) { const float* input_element = input_base + i; const float* input_row_start = InputBase + kh * DilatedInputWidthElements; const float* input_row_end = input_row_start + InputWidthElements; @@ -331,20 +374,89 @@ void input_values[i] = 0.0f; } } - InputVector = MlasLoadFloat32x4(input_values); + InputVector0 = MlasLoadFloat32x4(input_values); + } + + float32x4_t InputVector1; + if (is_main_region) { + InputVector1 = MlasLoadFloat32x4(input_base + 4); + } else { + float input_values[4]; + for (size_t i = 0; i < 4; i++) { + const float* input_element = input_base + 4 + i; + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; + + if (input_element >= input_row_start && input_element < input_row_end) { + input_values[i] = *input_element; + } else { + input_values[i] = 0.0f; + } + } + InputVector1 = MlasLoadFloat32x4(input_values); + } + + float32x4_t InputVector2; + if (is_main_region) { + InputVector2 = MlasLoadFloat32x4(input_base + 8); + } else { + float input_values[4]; + for (size_t i = 0; i < 4; i++) { + const float* input_element = input_base + 8 + i; + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; + + if (input_element >= input_row_start && input_element < input_row_end) { + input_values[i] = *input_element; + } else { + input_values[i] = 0.0f; + } + } + InputVector2 = MlasLoadFloat32x4(input_values); } - const float32x4_t FilterVector = MlasLoadFloat32x4(&Filter[kernel_pos * BlockSize]); + float32x4_t InputVector3; + if (is_main_region) { + InputVector3 = MlasLoadFloat32x4(input_base + 12); + } else { + float input_values[4]; + for (size_t i = 0; i < 4; i++) { + const float* input_element = input_base + 12 + i; + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + if (input_element >= input_row_start && input_element < input_row_end) { + input_values[i] = *input_element; + } else { + input_values[i] = 0.0f; + } + } + InputVector3 = MlasLoadFloat32x4(input_values); + } + + const float32x4_t FilterVector0 = MlasLoadFloat32x4(&Filter[kernel_pos * BlockSize]); + const float32x4_t FilterVector1 = MlasLoadFloat32x4(&Filter[kernel_pos * BlockSize + 4]); + const float32x4_t FilterVector2 = MlasLoadFloat32x4(&Filter[kernel_pos * BlockSize + 8]); + const float32x4_t FilterVector3 = MlasLoadFloat32x4(&Filter[kernel_pos * BlockSize + 12]); + + Accumulator0 = MlasMultiplyAddFloat32x4(InputVector0, FilterVector0, Accumulator0); + Accumulator1 = MlasMultiplyAddFloat32x4(InputVector1, FilterVector1, Accumulator1); + Accumulator2 = MlasMultiplyAddFloat32x4(InputVector2, FilterVector2, Accumulator2); + Accumulator3 = MlasMultiplyAddFloat32x4(InputVector3, FilterVector3, Accumulator3); } } if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + Accumulator0 = MlasMaximumFloat32x4(Accumulator0, ZeroVector); + Accumulator1 = MlasMaximumFloat32x4(Accumulator1, ZeroVector); + Accumulator2 = MlasMaximumFloat32x4(Accumulator2, ZeroVector); + Accumulator3 = MlasMaximumFloat32x4(Accumulator3, ZeroVector); } - MlasStoreFloat32x4(&Output[output_idx * BlockSize], Accumulator); + MlasStoreFloat32x4(&Output[output_idx * BlockSize], Accumulator0); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 4], Accumulator1); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 8], Accumulator2); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 12], Accumulator3); } } @@ -389,30 +501,65 @@ void for (size_t f = 0; f < FilterCount; f++) { const float* filter = Filter + f * FilterStrideElements; float* output = Output + f * OutputStrideElements; - float32x4_t Accumulator; + + float32x4_t Accumulator0, Accumulator1, Accumulator2, Accumulator3; + if (AccumulateOutput) { - Accumulator = MlasLoadFloat32x4(&output[output_idx * BlockSize]); + Accumulator0 = MlasLoadFloat32x4(&output[output_idx * BlockSize]); + Accumulator1 = MlasLoadFloat32x4(&output[output_idx * BlockSize + 4]); + Accumulator2 = MlasLoadFloat32x4(&output[output_idx * BlockSize + 8]); + Accumulator3 = MlasLoadFloat32x4(&output[output_idx * BlockSize + 12]); } else { - Accumulator = MlasBroadcastFloat32x4(0.0f); + Accumulator0 = MlasBroadcastFloat32x4(0.0f); + Accumulator1 = MlasBroadcastFloat32x4(0.0f); + Accumulator2 = MlasBroadcastFloat32x4(0.0f); + Accumulator3 = MlasBroadcastFloat32x4(0.0f); } + if (BiasAddition) { - const float32x4_t BiasVector = MlasLoadFloat32x4(&Bias[f * BlockSize]); - Accumulator = MlasAddFloat32x4(Accumulator, BiasVector); + const float32x4_t BiasVector0 = MlasLoadFloat32x4(&Bias[f * BlockSize]); + const float32x4_t BiasVector1 = MlasLoadFloat32x4(&Bias[f * BlockSize + 4]); + const float32x4_t BiasVector2 = MlasLoadFloat32x4(&Bias[f * BlockSize + 8]); + const float32x4_t BiasVector3 = MlasLoadFloat32x4(&Bias[f * BlockSize + 12]); + + Accumulator0 = MlasAddFloat32x4(Accumulator0, BiasVector0); + Accumulator1 = MlasAddFloat32x4(Accumulator1, BiasVector1); + Accumulator2 = MlasAddFloat32x4(Accumulator2, BiasVector2); + Accumulator3 = MlasAddFloat32x4(Accumulator3, BiasVector3); } + for (size_t c = 0; c < InputChannels; c++) { const float* input_ptr = Input + c * InputStrideElements + output_idx * StrideWidthElements; + for (size_t input_b = 0; input_b < BlockSize; input_b++) { const float input_value = input_ptr[input_b]; const float32x4_t InputVector = MlasBroadcastFloat32x4(input_value); + const float* filter_ptr = filter + (c * BlockSize + input_b) * BlockSize; - const float32x4_t FilterVector = MlasLoadFloat32x4(filter_ptr); - Accumulator = MlasMultiplyAddFloat32x4(InputVector, FilterVector, Accumulator); + + const float32x4_t FilterVector0 = MlasLoadFloat32x4(filter_ptr); + const float32x4_t FilterVector1 = MlasLoadFloat32x4(filter_ptr + 4); + const float32x4_t FilterVector2 = MlasLoadFloat32x4(filter_ptr + 8); + const float32x4_t FilterVector3 = MlasLoadFloat32x4(filter_ptr + 12); + + Accumulator0 = MlasMultiplyAddFloat32x4(InputVector, FilterVector0, Accumulator0); + Accumulator1 = MlasMultiplyAddFloat32x4(InputVector, FilterVector1, Accumulator1); + Accumulator2 = MlasMultiplyAddFloat32x4(InputVector, FilterVector2, Accumulator2); + Accumulator3 = MlasMultiplyAddFloat32x4(InputVector, FilterVector3, Accumulator3); } } + if (ReluActivation) { - Accumulator = MlasMaximumFloat32x4(Accumulator, ZeroVector); + Accumulator0 = MlasMaximumFloat32x4(Accumulator0, ZeroVector); + Accumulator1 = MlasMaximumFloat32x4(Accumulator1, ZeroVector); + Accumulator2 = MlasMaximumFloat32x4(Accumulator2, ZeroVector); + Accumulator3 = MlasMaximumFloat32x4(Accumulator3, ZeroVector); } - MlasStoreFloat32x4(&output[output_idx * BlockSize], Accumulator); + + MlasStoreFloat32x4(&output[output_idx * BlockSize], Accumulator0); + MlasStoreFloat32x4(&output[output_idx * BlockSize + 4], Accumulator1); + MlasStoreFloat32x4(&output[output_idx * BlockSize + 8], Accumulator2); + MlasStoreFloat32x4(&output[output_idx * BlockSize + 12], Accumulator3); } } } From 2b78377602d747850966d6b605167f750551a41b Mon Sep 17 00:00:00 2001 From: Rohan Date: Thu, 17 Jul 2025 21:31:52 +0000 Subject: [PATCH 19/26] Update pooling algorithm for block size 16 --- .../core/mlas/lib/spool_kernel_neon.cpp | 96 +++++++++++++++---- 1 file changed, 75 insertions(+), 21 deletions(-) diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp index ecfb83142f7b6..05198c0464cec 100644 --- a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -52,10 +52,14 @@ void const size_t TotalOutputCount = OutputCountLeftPad + OutputCount + OutputCountRightPad; const float MaxPaddingValue = std::numeric_limits::lowest(); + const MLAS_FLOAT32X4 MaxPaddingVector = MlasBroadcastFloat32x4(MaxPaddingValue); for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { - MLAS_FLOAT32X4 MaxVector = MaxPaddingVector; + MLAS_FLOAT32X4 MaxVector0 = MaxPaddingVector; + MLAS_FLOAT32X4 MaxVector1 = MaxPaddingVector; + MLAS_FLOAT32X4 MaxVector2 = MaxPaddingVector; + MLAS_FLOAT32X4 MaxVector3 = MaxPaddingVector; for (size_t kh = 0; kh < KernelHeight; kh++) { const float* row_start = InputBase + kh * DilatedInputWidthElements; @@ -65,10 +69,16 @@ void const float* input_ptr = Input + output_idx * StrideWidthElements + kh * DilatedInputWidthElements + kw * DilationWidthElements; - MLAS_FLOAT32X4 InputVector; - if (input_ptr >= row_start && (input_ptr + BlockSize) <= row_end) { - InputVector = MlasLoadFloat32x4(input_ptr); + MLAS_FLOAT32X4 InputVector0 = MlasLoadFloat32x4(input_ptr); + MLAS_FLOAT32X4 InputVector1 = MlasLoadFloat32x4(input_ptr + 4); + MLAS_FLOAT32X4 InputVector2 = MlasLoadFloat32x4(input_ptr + 8); + MLAS_FLOAT32X4 InputVector3 = MlasLoadFloat32x4(input_ptr + 12); + + MaxVector0 = MlasMaximumFloat32x4(MaxVector0, InputVector0); + MaxVector1 = MlasMaximumFloat32x4(MaxVector1, InputVector1); + MaxVector2 = MlasMaximumFloat32x4(MaxVector2, InputVector2); + MaxVector3 = MlasMaximumFloat32x4(MaxVector3, InputVector3); } else { std::vector values(BlockSize); for (size_t i = 0; i < BlockSize; i++) { @@ -79,14 +89,24 @@ void values[i] = MaxPaddingValue; } } - InputVector = MlasLoadFloat32x4(values.data()); - } - MaxVector = MlasMaximumFloat32x4(MaxVector, InputVector); + MLAS_FLOAT32X4 InputVector0 = MlasLoadFloat32x4(&values[0]); + MLAS_FLOAT32X4 InputVector1 = MlasLoadFloat32x4(&values[4]); + MLAS_FLOAT32X4 InputVector2 = MlasLoadFloat32x4(&values[8]); + MLAS_FLOAT32X4 InputVector3 = MlasLoadFloat32x4(&values[12]); + + MaxVector0 = MlasMaximumFloat32x4(MaxVector0, InputVector0); + MaxVector1 = MlasMaximumFloat32x4(MaxVector1, InputVector1); + MaxVector2 = MlasMaximumFloat32x4(MaxVector2, InputVector2); + MaxVector3 = MlasMaximumFloat32x4(MaxVector3, InputVector3); + } } } - MlasStoreFloat32x4(&Output[output_idx * BlockSize], MaxVector); + MlasStoreFloat32x4(&Output[output_idx * BlockSize], MaxVector0); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 4], MaxVector1); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 8], MaxVector2); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 12], MaxVector3); } } @@ -118,7 +138,10 @@ MlasPoolAverageFloatKernelNeonImpl( const MLAS_FLOAT32X4 ZeroVector = MlasZeroFloat32x4(); for (size_t output_idx = 0; output_idx < TotalOutputCount; output_idx++) { - MLAS_FLOAT32X4 SumVector = ZeroVector; + MLAS_FLOAT32X4 SumVector0 = ZeroVector; + MLAS_FLOAT32X4 SumVector1 = ZeroVector; + MLAS_FLOAT32X4 SumVector2 = ZeroVector; + MLAS_FLOAT32X4 SumVector3 = ZeroVector; std::vector valid_count; if (ExcludePad) { @@ -133,10 +156,16 @@ MlasPoolAverageFloatKernelNeonImpl( const float* input_ptr = Input + output_idx * StrideWidthElements + kh * DilatedInputWidthElements + kw * DilationWidthElements; - MLAS_FLOAT32X4 InputVector; - if (input_ptr >= row_start && (input_ptr + BlockSize) <= row_end) { - InputVector = MlasLoadFloat32x4(input_ptr); + MLAS_FLOAT32X4 InputVector0 = MlasLoadFloat32x4(input_ptr); + MLAS_FLOAT32X4 InputVector1 = MlasLoadFloat32x4(input_ptr + 4); + MLAS_FLOAT32X4 InputVector2 = MlasLoadFloat32x4(input_ptr + 8); + MLAS_FLOAT32X4 InputVector3 = MlasLoadFloat32x4(input_ptr + 12); + + SumVector0 = MlasAddFloat32x4(SumVector0, InputVector0); + SumVector1 = MlasAddFloat32x4(SumVector1, InputVector1); + SumVector2 = MlasAddFloat32x4(SumVector2, InputVector2); + SumVector3 = MlasAddFloat32x4(SumVector3, InputVector3); if (ExcludePad) { for (size_t i = 0; i < BlockSize; i++) { @@ -156,30 +185,55 @@ MlasPoolAverageFloatKernelNeonImpl( values[i] = 0.0f; } } - InputVector = MlasLoadFloat32x4(values.data()); - } - SumVector = MlasAddFloat32x4(SumVector, InputVector); + MLAS_FLOAT32X4 InputVector0 = MlasLoadFloat32x4(&values[0]); + MLAS_FLOAT32X4 InputVector1 = MlasLoadFloat32x4(&values[4]); + MLAS_FLOAT32X4 InputVector2 = MlasLoadFloat32x4(&values[8]); + MLAS_FLOAT32X4 InputVector3 = MlasLoadFloat32x4(&values[12]); + + SumVector0 = MlasAddFloat32x4(SumVector0, InputVector0); + SumVector1 = MlasAddFloat32x4(SumVector1, InputVector1); + SumVector2 = MlasAddFloat32x4(SumVector2, InputVector2); + SumVector3 = MlasAddFloat32x4(SumVector3, InputVector3); + } } } - MLAS_FLOAT32X4 ResultVector; if (ExcludePad) { std::vector results(BlockSize); - MlasStoreFloat32x4(results.data(), SumVector); + + MlasStoreFloat32x4(&results[0], SumVector0); + MlasStoreFloat32x4(&results[4], SumVector1); + MlasStoreFloat32x4(&results[8], SumVector2); + MlasStoreFloat32x4(&results[12], SumVector3); for (size_t i = 0; i < BlockSize; i++) { results[i] = results[i] / static_cast(valid_count[i]); } - ResultVector = MlasLoadFloat32x4(results.data()); + MLAS_FLOAT32X4 ResultVector0 = MlasLoadFloat32x4(&results[0]); + MLAS_FLOAT32X4 ResultVector1 = MlasLoadFloat32x4(&results[4]); + MLAS_FLOAT32X4 ResultVector2 = MlasLoadFloat32x4(&results[8]); + MLAS_FLOAT32X4 ResultVector3 = MlasLoadFloat32x4(&results[12]); + + MlasStoreFloat32x4(&Output[output_idx * BlockSize], ResultVector0); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 4], ResultVector1); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 8], ResultVector2); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 12], ResultVector3); } else { const float KernelSize = static_cast(ActualKernelSize); const MLAS_FLOAT32X4 KernelSizeVector = MlasBroadcastFloat32x4(KernelSize); - ResultVector = MlasDivideFloat32x4(SumVector, KernelSizeVector); - } - MlasStoreFloat32x4(&Output[output_idx * BlockSize], ResultVector); + MLAS_FLOAT32X4 ResultVector0 = MlasDivideFloat32x4(SumVector0, KernelSizeVector); + MLAS_FLOAT32X4 ResultVector1 = MlasDivideFloat32x4(SumVector1, KernelSizeVector); + MLAS_FLOAT32X4 ResultVector2 = MlasDivideFloat32x4(SumVector2, KernelSizeVector); + MLAS_FLOAT32X4 ResultVector3 = MlasDivideFloat32x4(SumVector3, KernelSizeVector); + + MlasStoreFloat32x4(&Output[output_idx * BlockSize], ResultVector0); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 4], ResultVector1); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 8], ResultVector2); + MlasStoreFloat32x4(&Output[output_idx * BlockSize + 12], ResultVector3); + } } } From ee9b9431c3b6072f5c22e686a89898be44c1fa44 Mon Sep 17 00:00:00 2001 From: Rohan Date: Tue, 29 Jul 2025 16:42:53 +0000 Subject: [PATCH 20/26] Remove comment --- onnxruntime/core/mlas/lib/platform.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp index 7a84da0072c6c..3ebde9d4d65a8 100644 --- a/onnxruntime/core/mlas/lib/platform.cpp +++ b/onnxruntime/core/mlas/lib/platform.cpp @@ -565,7 +565,7 @@ Return Value: this->PoolFloatKernel[MlasMaximumPooling] = MlasPoolMaximumFloatKernelNeon; this->PoolFloatKernel[MlasAveragePoolingExcludePad] = MlasPoolAverageExcludePadFloatKernelNeon; this->PoolFloatKernel[MlasAveragePoolingIncludePad] = MlasPoolAverageIncludePadFloatKernelNeon; - this->NchwcBlockSize = 16; // What is it supposed to be? + this->NchwcBlockSize = 16; // // Check if the processor supports ASIMD dot product instructions. From 23425e8e231b92ad5631416e0aaee1d73be747fb Mon Sep 17 00:00:00 2001 From: Rohan Date: Tue, 9 Sep 2025 22:28:28 +0000 Subject: [PATCH 21/26] Add correct header and refactor kernels to share code. --- onnxruntime/core/mlas/lib/mlasi.h | 4 +- onnxruntime/core/mlas/lib/sconv.h | 20 +-- .../core/mlas/lib/sconv_kernel_neon.cpp | 120 ++++++------------ .../core/mlas/lib/spool_kernel_neon.cpp | 8 +- 4 files changed, 45 insertions(+), 107 deletions(-) diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h index 4a9b2abca5a73..3612f8ae69aa0 100644 --- a/onnxruntime/core/mlas/lib/mlasi.h +++ b/onnxruntime/core/mlas/lib/mlasi.h @@ -949,6 +949,8 @@ extern "C" { #if defined(__aarch64__) && defined(__linux__) MLAS_SBGEMM_FLOAT_KERNEL MlasSbgemmKernelZero; MLAS_SBGEMM_FLOAT_KERNEL MlasSbgemmKernelAdd; +#endif +#if defined(MLAS_TARGET_ARM64) MLAS_CONV_FLOAT_KERNEL MlasConvNchwFloatKernelNeon; MLAS_CONV_FLOAT_KERNEL MlasConvNchwcFloatKernelNeon; MLAS_CONV_DEPTHWISE_FLOAT_KERNEL MlasConvDepthwiseFloatKernelNeon; @@ -1339,8 +1341,6 @@ struct MLAS_PLATFORM { const MLAS_GEMM_QUANT_DISPATCH* GemmU8U8Dispatch; const MLAS_GEMM_QUANT_DISPATCH* GemmU8S8Dispatch; const MLAS_GEMM_QUANT_DISPATCH* GemmS8S8Dispatch; -#endif -#if defined(__aarch64__) && defined(__linux__) MLAS_CONV_FLOAT_KERNEL* ConvNchwFloatKernel; MLAS_CONV_FLOAT_KERNEL* ConvNchwcFloatKernel; MLAS_CONV_DEPTHWISE_FLOAT_KERNEL* ConvDepthwiseFloatKernel; diff --git a/onnxruntime/core/mlas/lib/sconv.h b/onnxruntime/core/mlas/lib/sconv.h index d9f29f0492423..94e657638975a 100644 --- a/onnxruntime/core/mlas/lib/sconv.h +++ b/onnxruntime/core/mlas/lib/sconv.h @@ -10,27 +10,11 @@ Module Name: Abstract: - This module contains the private data structures and procedure prototypes - for the single precision convolution operation. + This module defines convolution kernel flags for configuring convolution + operations including output accumulation, bias addition, and activations. --*/ -#pragma once - -#include - -// -// Define the calling convention for MLAS functions. -// - -#ifndef MLASCALL -#if defined(_WIN32) && !defined(_WIN64) -#define MLASCALL __stdcall -#else -#define MLASCALL -#endif -#endif - // // Define the convolution kernel flags. // diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index 962ec1ba7057c..26572e3479d41 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -14,14 +14,10 @@ Module Name: --*/ +#include "mlasi.h" #include "sconv.h" -#if defined(__aarch64__) || defined(_M_ARM64) - -#include -#include - -#include "mlasi.h" +#if defined(MLAS_TARGET_ARM64) // Common implementation for NCHW and NCHWC convolution kernels template @@ -275,6 +271,39 @@ void ); } +// +// Helper function to load input vector with bounds checking +// +static inline float32x4_t +LoadInputVectorWithBounds( + const float* input_base, + size_t offset, + bool is_main_region, + const float* InputBase, + size_t kh, + size_t DilatedInputWidthElements, + size_t InputWidthElements +) +{ + if (is_main_region) { + return MlasLoadFloat32x4(input_base + offset); + } else { + float input_values[4]; + for (size_t i = 0; i < 4; i++) { + const float* input_element = input_base + offset + i; + const float* input_row_start = InputBase + kh * DilatedInputWidthElements; + const float* input_row_end = input_row_start + InputWidthElements; + + if (input_element >= input_row_start && input_element < input_row_end) { + input_values[i] = *input_element; + } else { + input_values[i] = 0.0f; + } + } + return MlasLoadFloat32x4(input_values); + } +} + // // Implementation of MlasConvDepthwiseFloatKernelNeon // @@ -358,81 +387,10 @@ void const float* input_base = Input + output_idx * StrideWidthElements + kh * DilatedInputWidthElements + kw * DilationWidthElements; - float32x4_t InputVector0; - if (is_main_region) { - InputVector0 = MlasLoadFloat32x4(input_base); - } else { - float input_values[4]; - for (size_t i = 0; i < 4; i++) { - const float* input_element = input_base + i; - const float* input_row_start = InputBase + kh * DilatedInputWidthElements; - const float* input_row_end = input_row_start + InputWidthElements; - - if (input_element >= input_row_start && input_element < input_row_end) { - input_values[i] = *input_element; - } else { - input_values[i] = 0.0f; - } - } - InputVector0 = MlasLoadFloat32x4(input_values); - } - - float32x4_t InputVector1; - if (is_main_region) { - InputVector1 = MlasLoadFloat32x4(input_base + 4); - } else { - float input_values[4]; - for (size_t i = 0; i < 4; i++) { - const float* input_element = input_base + 4 + i; - const float* input_row_start = InputBase + kh * DilatedInputWidthElements; - const float* input_row_end = input_row_start + InputWidthElements; - - if (input_element >= input_row_start && input_element < input_row_end) { - input_values[i] = *input_element; - } else { - input_values[i] = 0.0f; - } - } - InputVector1 = MlasLoadFloat32x4(input_values); - } - - float32x4_t InputVector2; - if (is_main_region) { - InputVector2 = MlasLoadFloat32x4(input_base + 8); - } else { - float input_values[4]; - for (size_t i = 0; i < 4; i++) { - const float* input_element = input_base + 8 + i; - const float* input_row_start = InputBase + kh * DilatedInputWidthElements; - const float* input_row_end = input_row_start + InputWidthElements; - - if (input_element >= input_row_start && input_element < input_row_end) { - input_values[i] = *input_element; - } else { - input_values[i] = 0.0f; - } - } - InputVector2 = MlasLoadFloat32x4(input_values); - } - - float32x4_t InputVector3; - if (is_main_region) { - InputVector3 = MlasLoadFloat32x4(input_base + 12); - } else { - float input_values[4]; - for (size_t i = 0; i < 4; i++) { - const float* input_element = input_base + 12 + i; - const float* input_row_start = InputBase + kh * DilatedInputWidthElements; - const float* input_row_end = input_row_start + InputWidthElements; - - if (input_element >= input_row_start && input_element < input_row_end) { - input_values[i] = *input_element; - } else { - input_values[i] = 0.0f; - } - } - InputVector3 = MlasLoadFloat32x4(input_values); - } + float32x4_t InputVector0 = LoadInputVectorWithBounds(input_base, 0, is_main_region, InputBase, kh, DilatedInputWidthElements, InputWidthElements); + float32x4_t InputVector1 = LoadInputVectorWithBounds(input_base, 4, is_main_region, InputBase, kh, DilatedInputWidthElements, InputWidthElements); + float32x4_t InputVector2 = LoadInputVectorWithBounds(input_base, 8, is_main_region, InputBase, kh, DilatedInputWidthElements, InputWidthElements); + float32x4_t InputVector3 = LoadInputVectorWithBounds(input_base, 12, is_main_region, InputBase, kh, DilatedInputWidthElements, InputWidthElements); const float32x4_t FilterVector0 = MlasLoadFloat32x4(&Filter[kernel_pos * BlockSize]); const float32x4_t FilterVector1 = MlasLoadFloat32x4(&Filter[kernel_pos * BlockSize + 4]); diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp index 05198c0464cec..ef0d9b6a5ad28 100644 --- a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -14,14 +14,10 @@ Module Name: --*/ -#if defined(__aarch64__) || defined(_M_ARM64) - -#include -#include -#include - #include "mlasi.h" +#if defined(MLAS_TARGET_ARM64) + void MLASCALL MlasPoolMaximumFloatKernelNeon( From 7000e9fe52bacb6ce05b31033b61167670b0cfbc Mon Sep 17 00:00:00 2001 From: Rohan Date: Tue, 9 Sep 2025 23:09:03 +0000 Subject: [PATCH 22/26] Address Copilot comments --- onnxruntime/core/mlas/lib/spool_kernel_neon.cpp | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp index ef0d9b6a5ad28..473348ffc16dd 100644 --- a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -76,7 +76,7 @@ void MaxVector2 = MlasMaximumFloat32x4(MaxVector2, InputVector2); MaxVector3 = MlasMaximumFloat32x4(MaxVector3, InputVector3); } else { - std::vector values(BlockSize); + float values[BlockSize]; for (size_t i = 0; i < BlockSize; i++) { const float* element_ptr = input_ptr + i; if (element_ptr >= row_start && element_ptr < row_end) { @@ -169,7 +169,7 @@ MlasPoolAverageFloatKernelNeonImpl( } } } else { - std::vector values(BlockSize); + float values[BlockSize]; for (size_t i = 0; i < BlockSize; i++) { const float* element_ptr = input_ptr + i; if (element_ptr >= row_start && element_ptr < row_end) { @@ -196,7 +196,7 @@ MlasPoolAverageFloatKernelNeonImpl( } if (ExcludePad) { - std::vector results(BlockSize); + float results[BlockSize]; MlasStoreFloat32x4(&results[0], SumVector0); MlasStoreFloat32x4(&results[4], SumVector1); From c5c3f051ff3686276a46d27ddbaa285141143aa7 Mon Sep 17 00:00:00 2001 From: Rohan Date: Tue, 9 Sep 2025 23:09:18 +0000 Subject: [PATCH 23/26] Extend kernels to Windows & Apple --- cmake/onnxruntime_mlas.cmake | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index c8cb173f635c0..0916e8569e036 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -108,6 +108,8 @@ function(setup_mlas_source_for_windows) ${MLAS_SRC_DIR}/eltwise_kernel_neon.h ${MLAS_SRC_DIR}/eltwise_kernel_neon.cpp ${MLAS_SRC_DIR}/eltwise_kernel_neon_fp16.cpp + ${MLAS_SRC_DIR}/sconv_kernel_neon.cpp + ${MLAS_SRC_DIR}/spool_kernel_neon.cpp ) set(mlas_platform_preprocess_srcs @@ -429,6 +431,8 @@ else() ${MLAS_SRC_DIR}/softmax_kernel_neon.cpp ${MLAS_SRC_DIR}/eltwise_kernel_neon.h ${MLAS_SRC_DIR}/eltwise_kernel_neon.cpp + ${MLAS_SRC_DIR}/sconv_kernel_neon.cpp + ${MLAS_SRC_DIR}/spool_kernel_neon.cpp ) if (onnxruntime_USE_KLEIDIAI) setup_kleidiai() @@ -449,8 +453,6 @@ else() ${MLAS_SRC_DIR}/qgemm_kernel_smmla.cpp ${MLAS_SRC_DIR}/qgemm_kernel_ummla.cpp ${MLAS_SRC_DIR}/sbgemm_kernel_neon.cpp - ${MLAS_SRC_DIR}/sconv_kernel_neon.cpp - ${MLAS_SRC_DIR}/spool_kernel_neon.cpp ${MLAS_SRC_DIR}/cast_kernel_neon.cpp ${MLAS_SRC_DIR}/hqnbitgemm_kernel_neon_fp16.cpp ${MLAS_SRC_DIR}/rotary_embedding_kernel_neon_fp16.cpp From 506bf053ad20d366e08ee5b9b62cdefe74da01fd Mon Sep 17 00:00:00 2001 From: Rohan Date: Wed, 10 Sep 2025 00:25:26 +0000 Subject: [PATCH 24/26] Hardcode BlockSize to 16 and add it to the header. --- onnxruntime/core/mlas/lib/sconv.h | 7 +++++++ onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp | 5 ++--- onnxruntime/core/mlas/lib/spool_kernel_neon.cpp | 6 +++--- 3 files changed, 12 insertions(+), 6 deletions(-) diff --git a/onnxruntime/core/mlas/lib/sconv.h b/onnxruntime/core/mlas/lib/sconv.h index 94e657638975a..ed2beda6d65f0 100644 --- a/onnxruntime/core/mlas/lib/sconv.h +++ b/onnxruntime/core/mlas/lib/sconv.h @@ -15,6 +15,13 @@ Module Name: --*/ +/* + The MLAS_NEON_BLOCK_SIZE has to be the equal to the NchwcBlockSize in platform.cpp. + Refer to the discussion in https://github.com/microsoft/onnxruntime/pull/25580. +*/ + +constexpr size_t MLAS_NEON_BLOCK_SIZE = 16; + // // Define the convolution kernel flags. // diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index 26572e3479d41..05fd6b07905ef 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -19,6 +19,8 @@ Module Name: #if defined(MLAS_TARGET_ARM64) +constexpr size_t BlockSize = MLAS_NEON_BLOCK_SIZE; + // Common implementation for NCHW and NCHWC convolution kernels template void @@ -49,7 +51,6 @@ void const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; - const size_t BlockSize = MlasNchwcGetBlockSize(); const float32x4_t ZeroVector = MlasBroadcastFloat32x4(0.0f); const size_t StrideWidthElements = StrideWidth / sizeof(float); @@ -337,7 +338,6 @@ void const bool BiasAddition = (KernelFlags & MLAS_CONV_KERNEL_FLAG_BIAS_ADDITION) != 0; const bool ReluActivation = (KernelFlags & MLAS_CONV_KERNEL_FLAG_RELU_ACTIVATION) != 0; - const size_t BlockSize = MlasNchwcGetBlockSize(); const float32x4_t ZeroVector = MlasBroadcastFloat32x4(0.0f); const size_t StrideWidthElements = StrideWidth / sizeof(float); @@ -452,7 +452,6 @@ void const size_t FilterStrideElements = FilterStride / sizeof(float); const size_t OutputStrideElements = OutputStride / sizeof(float); - const size_t BlockSize = MlasNchwcGetBlockSize(); const float32x4_t ZeroVector = MlasBroadcastFloat32x4(0.0f); for (size_t output_idx = 0; output_idx < OutputCount; output_idx++) { diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp index 473348ffc16dd..afb534f09aee8 100644 --- a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -15,9 +15,12 @@ Module Name: --*/ #include "mlasi.h" +#include "sconv.h" #if defined(MLAS_TARGET_ARM64) +constexpr size_t BlockSize = MLAS_NEON_BLOCK_SIZE; + void MLASCALL MlasPoolMaximumFloatKernelNeon( @@ -39,8 +42,6 @@ void { MLAS_UNREFERENCED_PARAMETER(ActualKernelSize); MLAS_UNREFERENCED_PARAMETER(InputStride); - - const size_t BlockSize = MlasNchwcGetBlockSize(); const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); const size_t InputWidthElements = InputWidth / sizeof(float); @@ -124,7 +125,6 @@ MlasPoolAverageFloatKernelNeonImpl( bool ExcludePad ) { - const size_t BlockSize = MlasNchwcGetBlockSize(); const size_t StrideWidthElements = StrideWidth / sizeof(float); const size_t DilationWidthElements = DilationWidth / sizeof(float); const size_t InputWidthElements = InputWidth / sizeof(float); From fb5fb504880a960974aac855f3df61d80dbdf5c8 Mon Sep 17 00:00:00 2001 From: Rohan Date: Fri, 12 Sep 2025 17:23:46 +0000 Subject: [PATCH 25/26] Increase android build size to 10% higher than the CI-reported size of the NCHWc build --- .github/workflows/android.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/android.yml b/.github/workflows/android.yml index 35d412f88ddbc..7f45c824959f6 100644 --- a/.github/workflows/android.yml +++ b/.github/workflows/android.yml @@ -71,8 +71,8 @@ jobs: run: | set -e -x BINARY_SIZE_THRESHOLD_ARGS="" - echo "Binary size threshold in bytes: 1436672" - BINARY_SIZE_THRESHOLD_ARGS="--threshold_size_in_bytes 1436672" + echo "Binary size threshold in bytes: 1722565" + BINARY_SIZE_THRESHOLD_ARGS="--threshold_size_in_bytes 1722565" # Ensure ANDROID_NDK_HOME is available and get its real path if [ -z "$ANDROID_NDK_HOME" ]; then From fb99f7dc1d3082b822944b504ae8c2479e89b3d8 Mon Sep 17 00:00:00 2001 From: Rohan Date: Fri, 12 Sep 2025 17:24:45 +0000 Subject: [PATCH 26/26] Centralize MLAS_NEON_NCHWC_BLOCK_SIZE --- onnxruntime/core/mlas/lib/mlasi.h | 1 + onnxruntime/core/mlas/lib/platform.cpp | 2 +- onnxruntime/core/mlas/lib/sconv.h | 7 ------- onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp | 6 +----- onnxruntime/core/mlas/lib/spool_kernel_neon.cpp | 7 +------ 5 files changed, 4 insertions(+), 19 deletions(-) diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h index a89dd69f9467f..4cad44a56ba96 100644 --- a/onnxruntime/core/mlas/lib/mlasi.h +++ b/onnxruntime/core/mlas/lib/mlasi.h @@ -1410,6 +1410,7 @@ struct MLAS_PLATFORM { int32_t MaximumThreadCount; #elif defined(MLAS_TARGET_ARM64) static constexpr int32_t MaximumThreadCount = MLAS_MAXIMUM_THREAD_COUNT * 4; + static constexpr size_t MLAS_NEON_NCHWC_BLOCK_SIZE = 16; #else static constexpr int32_t MaximumThreadCount = MLAS_MAXIMUM_THREAD_COUNT; #endif diff --git a/onnxruntime/core/mlas/lib/platform.cpp b/onnxruntime/core/mlas/lib/platform.cpp index 31e45639b8231..923e513ccb07a 100644 --- a/onnxruntime/core/mlas/lib/platform.cpp +++ b/onnxruntime/core/mlas/lib/platform.cpp @@ -565,7 +565,7 @@ Return Value: this->PoolFloatKernel[MlasMaximumPooling] = MlasPoolMaximumFloatKernelNeon; this->PoolFloatKernel[MlasAveragePoolingExcludePad] = MlasPoolAverageExcludePadFloatKernelNeon; this->PoolFloatKernel[MlasAveragePoolingIncludePad] = MlasPoolAverageIncludePadFloatKernelNeon; - this->NchwcBlockSize = 16; + this->NchwcBlockSize = MLAS_NEON_NCHWC_BLOCK_SIZE; // // Check if the processor supports ASIMD dot product instructions. diff --git a/onnxruntime/core/mlas/lib/sconv.h b/onnxruntime/core/mlas/lib/sconv.h index ed2beda6d65f0..94e657638975a 100644 --- a/onnxruntime/core/mlas/lib/sconv.h +++ b/onnxruntime/core/mlas/lib/sconv.h @@ -15,13 +15,6 @@ Module Name: --*/ -/* - The MLAS_NEON_BLOCK_SIZE has to be the equal to the NchwcBlockSize in platform.cpp. - Refer to the discussion in https://github.com/microsoft/onnxruntime/pull/25580. -*/ - -constexpr size_t MLAS_NEON_BLOCK_SIZE = 16; - // // Define the convolution kernel flags. // diff --git a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp index 05fd6b07905ef..3ecad66a32886 100644 --- a/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/sconv_kernel_neon.cpp @@ -17,9 +17,7 @@ Module Name: #include "mlasi.h" #include "sconv.h" -#if defined(MLAS_TARGET_ARM64) - -constexpr size_t BlockSize = MLAS_NEON_BLOCK_SIZE; +constexpr size_t BlockSize = MLAS_PLATFORM::MLAS_NEON_NCHWC_BLOCK_SIZE; // Common implementation for NCHW and NCHWC convolution kernels template @@ -520,5 +518,3 @@ void } } } - -#endif // __aarch64__ || _M_ARM64 \ No newline at end of file diff --git a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp index afb534f09aee8..8cca036d54c3a 100644 --- a/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp +++ b/onnxruntime/core/mlas/lib/spool_kernel_neon.cpp @@ -15,11 +15,8 @@ Module Name: --*/ #include "mlasi.h" -#include "sconv.h" -#if defined(MLAS_TARGET_ARM64) - -constexpr size_t BlockSize = MLAS_NEON_BLOCK_SIZE; +constexpr size_t BlockSize = MLAS_PLATFORM::MLAS_NEON_NCHWC_BLOCK_SIZE; void MLASCALL @@ -290,5 +287,3 @@ void false // ExcludePad = false ); } - -#endif // __aarch64__ || _M_ARM64 \ No newline at end of file