diff --git a/cmake/onnxruntime_mlas.cmake b/cmake/onnxruntime_mlas.cmake index 1b6bc1713c175..fd33d430dbd64 100644 --- a/cmake/onnxruntime_mlas.cmake +++ b/cmake/onnxruntime_mlas.cmake @@ -375,6 +375,7 @@ function(setup_kleidiai) ${MLAS_SRC_DIR}/kleidiai/convolve_kleidiai.cpp ${MLAS_SRC_DIR}/kleidiai/qgemm_kleidiai.cpp ${MLAS_SRC_DIR}/kleidiai/qnbitgemm_kleidiai.cpp + ${MLAS_SRC_DIR}/kleidiai/dw_conv_kleidiai.cpp ) target_link_libraries(onnxruntime_mlas PRIVATE kleidiai) list(APPEND onnxruntime_EXTERNAL_LIBRARIES kleidiai) diff --git a/onnxruntime/core/mlas/lib/convolve.cpp b/onnxruntime/core/mlas/lib/convolve.cpp index d5d8560f8c540..6d6eca2c0a98f 100644 --- a/onnxruntime/core/mlas/lib/convolve.cpp +++ b/onnxruntime/core/mlas/lib/convolve.cpp @@ -1,6 +1,7 @@ /*++ Copyright (c) Microsoft Corporation. All rights reserved. +SPDX-FileCopyrightText: Copyright 2026 Arm Limited and/or its affiliates Licensed under the MIT License. @@ -15,6 +16,9 @@ Module Name: --*/ #include "mlasi.h" +#if defined(USE_KLEIDIAI) && defined(MLAS_TARGET_ARM64) +#include "kleidiai/mlasi_kleidiai.h" +#endif #if defined(BUILD_MLAS_NO_ONNXRUNTIME) // Standalone MLAS builds don't have access to the ORT-internal SafeInt // wrapper; fall back to the SafeInt.hpp header directly (its default @@ -1351,6 +1355,7 @@ static constexpr size_t ComputeChannelsLastConvOutSize(size_t input, size_t kern return 0; } + #endif } // namespace @@ -1453,22 +1458,34 @@ MlasConvSupportsDepthwiseChannelsLast2DFloatKernel( MLAS_UNREFERENCED_PARAMETER(Beta); return false; #else - MLAS_UNREFERENCED_PARAMETER(Dimensions); - MLAS_UNREFERENCED_PARAMETER(BatchCount); - MLAS_UNREFERENCED_PARAMETER(GroupCount); - MLAS_UNREFERENCED_PARAMETER(InputChannelsPerGroup); - MLAS_UNREFERENCED_PARAMETER(InputShape); - MLAS_UNREFERENCED_PARAMETER(KernelShape); - MLAS_UNREFERENCED_PARAMETER(DilationShape); - MLAS_UNREFERENCED_PARAMETER(Padding); - MLAS_UNREFERENCED_PARAMETER(StrideShape); - MLAS_UNREFERENCED_PARAMETER(FilterCount); - MLAS_UNREFERENCED_PARAMETER(Beta); + // Channels-last float convolution is only implemented by the KleidiAI + // override. The generic MLAS convolution path assumes NCHW layout. + if (GetMlasPlatform().MlasConvPrepareOverride == nullptr || + GetMlasPlatform().MlasConvOverride == nullptr) { + return false; + } - // TODO: enable only for shapes supported by the dedicated - // depthwise kernel. Until then, keep depthwise/grouped convolutions out of - // the Arm® KleidiAI™ NHWC path. - return false; + if (Dimensions != 2) { + return false; + } + + MLAS_CONV_PARAMETERS parameters{}; + parameters.Dimensions = Dimensions; + parameters.BatchCount = BatchCount; + parameters.GroupCount = GroupCount; + parameters.InputChannels = InputChannelsPerGroup; + parameters.FilterCount = FilterCount; + parameters.Beta = Beta; + for (size_t dim = 0; dim < Dimensions; ++dim) { + parameters.InputShape[dim] = InputShape[dim]; + parameters.KernelShape[dim] = KernelShape[dim]; + parameters.DilationShape[dim] = DilationShape[dim]; + parameters.Padding[dim] = Padding[dim]; + parameters.Padding[dim + Dimensions] = Padding[dim + Dimensions]; + parameters.StrideShape[dim] = StrideShape[dim]; + } + + return ArmKleidiAI::DepthwiseConvKleidiAISupported(¶meters); #endif } diff --git a/onnxruntime/core/mlas/lib/kleidiai/convolve_kleidiai.cpp b/onnxruntime/core/mlas/lib/kleidiai/convolve_kleidiai.cpp index 0fe529d9193a5..7a84e4a7d6f37 100644 --- a/onnxruntime/core/mlas/lib/kleidiai/convolve_kleidiai.cpp +++ b/onnxruntime/core/mlas/lib/kleidiai/convolve_kleidiai.cpp @@ -123,7 +123,17 @@ static size_t ComputeMlasWorkingBufferSize(const size_t co, return m * co; } -static bool CheckCapabilitiesSme(const MLAS_CONV_PARAMETERS* Parameters) { +static bool CheckIgemmRouteCapabilities(const MLAS_CONV_PARAMETERS* Parameters, + const ArmKleidiAI::ConvRouteSelection& route_selection) { + if (route_selection.route != ArmKleidiAI::ConvRoute::IGemm) { + if (route_selection.route == ArmKleidiAI::ConvRoute::SGemmFallback) { + KLEIDIAI_DEBUG_LOG("CheckIgemmRouteCapabilities returning false to prefer SGEMM-backed conv path."); + } else { + KLEIDIAI_DEBUG_LOG("CheckIgemmRouteCapabilities returning false on non-IGEMM route."); + } + return false; + } + // Grouped support in this override is only implemented for channels-last // layout. The generic grouped path still assumes contiguous per-group CHW. if (Parameters->GroupCount > 1 && !Parameters->ChannelsLast) { @@ -141,35 +151,24 @@ static bool CheckCapabilitiesSme(const MLAS_CONV_PARAMETERS* Parameters) { Parameters->StrideShape, Parameters->FilterCount, Parameters->Beta)) { - KLEIDIAI_DEBUG_LOG("CheckCapabilitiesSme returning false on shared capability checks."); + KLEIDIAI_DEBUG_LOG("CheckIgemmRouteCapabilities returning false on shared capability checks."); return false; } - const auto route_selection = ArmKleidiAI::SelectConvRoute(Parameters); - const auto route = route_selection.route; - - if (route == ArmKleidiAI::ConvRoute::IGemm) { - // ensure LHS packed buffer size is non-zero - const size_t d_kh = route_selection.effective_kernel_h; - const size_t d_kw = route_selection.effective_kernel_w; - const size_t m_step = imatmul_conv.ukernel.get_m_step(); + // ensure LHS packed buffer size is non-zero + const size_t d_kh = route_selection.effective_kernel_h; + const size_t d_kw = route_selection.effective_kernel_w; + const size_t m_step = imatmul_conv.ukernel.get_m_step(); - const size_t bytes_per_m_step = kai_get_lhs_packed_size_lhs_imatmul_pack_x32p2vlx1_x32p_sme( - m_step, d_kh * d_kw, Parameters->InputChannels); + const size_t bytes_per_m_step = kai_get_lhs_packed_size_lhs_imatmul_pack_x32p2vlx1_x32p_sme( + m_step, d_kh * d_kw, Parameters->InputChannels); - if (bytes_per_m_step == 0) { - KLEIDIAI_DEBUG_LOG("CheckCapabilitiesSME returning false on zero LHS packed size"); - return false; - } - return true; + if (bytes_per_m_step == 0) { + KLEIDIAI_DEBUG_LOG("CheckIgemmRouteCapabilities returning false on zero LHS packed size"); + return false; } - if (route == ArmKleidiAI::ConvRoute::SGemmFallback) { - KLEIDIAI_DEBUG_LOG("CheckCapabilitiesSme returning false to prefer SGEMM-backed conv path."); - } else { - KLEIDIAI_DEBUG_LOG("CheckCapabilitiesSme returning false on functional or optimization checks."); - } - return false; + return true; } //General purpose axis swapping @@ -759,7 +758,13 @@ ArmKleidiAI::MlasConvPrepare(MLAS_CONV_PARAMETERS* Parameters, Parameters->ThreadCount = MlasGetMaximumThreadCount(ThreadPool); - if(!CheckCapabilitiesSme(Parameters)){ + const auto route_selection = ArmKleidiAI::SelectConvRoute(Parameters); + if (route_selection.route == ArmKleidiAI::ConvRoute::Depthwise) { + *WorkingBufferSize = 0; + return true; + } + + if (!CheckIgemmRouteCapabilities(Parameters, route_selection)) { return false; } @@ -791,26 +796,61 @@ ArmKleidiAI::MlasConv( return false; } - if(!CheckCapabilitiesSme(Parameters)){ - // Fallback to Default Mlas + const auto route_selection = ArmKleidiAI::SelectConvRoute(Parameters); + + if (route_selection.route == ArmKleidiAI::ConvRoute::Depthwise) { + if (DepthwiseConvKleidiAI(Parameters->BatchCount, + Parameters->InputShape[0], + Parameters->InputShape[1], + Parameters->GroupCount * Parameters->InputChannels, + Parameters->KernelShape[0], + Parameters->KernelShape[1], + Parameters->Padding[0], + Parameters->Padding[1], + Parameters->Padding[2], + Parameters->Padding[3], + Parameters->ChannelsLast, + Input, + Filter, + Bias, + Parameters->FilterIsPacked ? Parameters->PackedFilter : nullptr, + Output, + -std::numeric_limits::max(), + std::numeric_limits::max(), + ThreadPool)) { + const size_t activation_rows = Parameters->ChannelsLast + ? Parameters->OutputSize + : Parameters->GroupCount * Parameters->FilterCount; + const size_t activation_cols = Parameters->ChannelsLast + ? Parameters->GroupCount * Parameters->FilterCount + : Parameters->OutputSize; + MlasActivation(Parameters->Activation, Output, nullptr, activation_rows, activation_cols, activation_cols); + return true; + } + return false; - }; - ConvolveSme(Parameters->FilterCount, Parameters->InputChannels, // channel out, in - Parameters->InputShape[0], Parameters->InputShape[1], // image dimensions - Parameters->KernelShape[0], Parameters->KernelShape[1], // kernel dimensions - Parameters->StrideShape[0], Parameters->StrideShape[1], // kernel stride dimensions - Parameters->DilationShape[0], Parameters->DilationShape[1], // kernel dilation - Parameters->Padding[0], // image padding - Parameters->GroupCount, // filter groups - Filter, Bias, - reinterpret_cast(Parameters->PackedFilter), - Parameters->PackedFilterGroupStride, - Input, Output, WorkingBuffer, Parameters->ChannelsLast, ThreadPool); - - const bool grouped_channels_last = Parameters->ChannelsLast && Parameters->GroupCount > 1; - const size_t activation_rows = grouped_channels_last ? Parameters->OutputSize : Parameters->FilterCount; - const size_t activation_cols = - grouped_channels_last ? Parameters->GroupCount * Parameters->FilterCount : Parameters->OutputSize; - MlasActivation(Parameters->Activation, Output, nullptr, activation_rows, activation_cols, activation_cols); - return true; + } + + if (CheckIgemmRouteCapabilities(Parameters, route_selection)) { + ConvolveSme(Parameters->FilterCount, Parameters->InputChannels, // channel out, in + Parameters->InputShape[0], Parameters->InputShape[1], // image dimensions + Parameters->KernelShape[0], Parameters->KernelShape[1], // kernel dimensions + Parameters->StrideShape[0], Parameters->StrideShape[1], // kernel stride dimensions + Parameters->DilationShape[0], Parameters->DilationShape[1], // kernel dilation + Parameters->Padding[0], // image padding + Parameters->GroupCount, // filter groups + Filter, Bias, + reinterpret_cast(Parameters->PackedFilter), + Parameters->PackedFilterGroupStride, + Input, Output, WorkingBuffer, Parameters->ChannelsLast, ThreadPool); + + const bool grouped_channels_last = Parameters->ChannelsLast && Parameters->GroupCount > 1; + const size_t activation_rows = grouped_channels_last ? Parameters->OutputSize : Parameters->FilterCount; + const size_t activation_cols = + grouped_channels_last ? Parameters->GroupCount * Parameters->FilterCount : Parameters->OutputSize; + MlasActivation(Parameters->Activation, Output, nullptr, activation_rows, activation_cols, activation_cols); + return true; + } + + return false; } diff --git a/onnxruntime/core/mlas/lib/kleidiai/dw_conv_kleidiai.cpp b/onnxruntime/core/mlas/lib/kleidiai/dw_conv_kleidiai.cpp new file mode 100644 index 0000000000000..fee80b65175a1 --- /dev/null +++ b/onnxruntime/core/mlas/lib/kleidiai/dw_conv_kleidiai.cpp @@ -0,0 +1,456 @@ +// +// SPDX-FileCopyrightText: Copyright 2025-2026 Arm Limited and/or its affiliates +// +// SPDX-License-Identifier: MIT +// + +#include "mlasi_kleidiai.h" + +#include +#include +#include +#include +#include +#include + +#include "kai_ukernel_interface.h" + +#include "kai/ukernels/dwconv/pack/kai_rhs_dwconv_pack_x32p1vlx1b_x32_x32_sme.h" + +extern "C" uint64_t kai_get_sme_vector_length_u8(void); + +namespace ArmKleidiAI { +namespace { + +const KaiF32DepthwiseConvKernel& dwconv = GetKleidiAIDepthwiseConvUKernel(); + +struct DwconvTlsBuffers { + std::vector feature_map_nhwc; + std::vector weights_hwcn; + std::vector weights_packed; + std::vector nhwc_out; + std::vector bias_fallback; + + void ReleaseLargeBuffers() { + ArmKleidiAI::MlasShrinkKleidiAIScratchIfTooLarge(feature_map_nhwc); + ArmKleidiAI::MlasShrinkKleidiAIScratchIfTooLarge(weights_hwcn); + ArmKleidiAI::MlasShrinkKleidiAIScratchIfTooLarge(weights_packed); + ArmKleidiAI::MlasShrinkKleidiAIScratchIfTooLarge(nhwc_out); + ArmKleidiAI::MlasShrinkKleidiAIScratchIfTooLarge(bias_fallback); + } +}; + +struct ScopedKaiDwconvTlsCleanup { + DwconvTlsBuffers& buffers; + + ~ScopedKaiDwconvTlsCleanup() { + buffers.ReleaseLargeBuffers(); + } +}; + +thread_local DwconvTlsBuffers g_dwconv_tls; + +constexpr size_t kDwconvColsPerTile = 4; + +static void ConvertNchwToNhwc(const float* src, + float* dst, + size_t batches, + size_t channels, + size_t height, + size_t width) { + const size_t src_stride_n = channels * height * width; + const size_t src_stride_c = height * width; + const size_t src_stride_h = width; + const size_t dst_stride_n = height * width * channels; + const size_t dst_stride_h = width * channels; + const size_t dst_stride_w = channels; + + for (size_t n = 0; n < batches; ++n) { + for (size_t h = 0; h < height; ++h) { + for (size_t w = 0; w < width; ++w) { + for (size_t c = 0; c < channels; ++c) { + const size_t src_index = n * src_stride_n + c * src_stride_c + h * src_stride_h + w; + const size_t dst_index = n * dst_stride_n + h * dst_stride_h + w * dst_stride_w + c; + dst[dst_index] = src[src_index]; + } + } + } + } +} + +static void ConvertNhwcToNchw(const float* src, + float* dst, + size_t batches, + size_t channels, + size_t height, + size_t width) { + const size_t dst_stride_n = channels * height * width; + const size_t dst_stride_c = height * width; + const size_t dst_stride_h = width; + const size_t src_stride_n = height * width * channels; + const size_t src_stride_h = width * channels; + const size_t src_stride_w = channels; + + for (size_t n = 0; n < batches; ++n) { + for (size_t c = 0; c < channels; ++c) { + for (size_t h = 0; h < height; ++h) { + for (size_t w = 0; w < width; ++w) { + const size_t src_index = n * src_stride_n + h * src_stride_h + w * src_stride_w + c; + const size_t dst_index = n * dst_stride_n + c * dst_stride_c + h * dst_stride_h + w; + dst[dst_index] = src[src_index]; + } + } + } + } +} + +static void ConvertDepthwiseWeightsToHwcn(const float* src, + float* dst, + size_t channels, + size_t filter_height, + size_t filter_width) { + const size_t kernel_size = filter_height * filter_width; + for (size_t c = 0; c < channels; ++c) { + const float* channel_weights = src + c * kernel_size; + for (size_t kh = 0; kh < filter_height; ++kh) { + for (size_t kw = 0; kw < filter_width; ++kw) { + const size_t dst_index = (kh * filter_width + kw) * channels + c; + dst[dst_index] = channel_weights[kh * filter_width + kw]; + } + } + } +} + +static bool TryComputeDepthwiseOutputShape(size_t in_height, + size_t in_width, + size_t filter_height, + size_t filter_width, + size_t pad_top, + size_t pad_left, + size_t pad_bottom, + size_t pad_right, + size_t& out_height, + size_t& out_width) { + const size_t padded_height = in_height + pad_top + pad_bottom; + const size_t padded_width = in_width + pad_left + pad_right; + + if (padded_height < filter_height || padded_width < filter_width) { + return false; + } + + out_height = padded_height + 1 - filter_height; + out_width = padded_width + 1 - filter_width; + return true; +} + +} // namespace + +bool +MLASCALL +DepthwiseConvKleidiAISupported(const MLAS_CONV_PARAMETERS* Parameters) { + if (Parameters == nullptr) { + return false; + } + + if (!UseSME2) { + return false; + } + + if (Parameters->BackendKernelSelectorConfig && !Parameters->BackendKernelSelectorConfig->use_kleidiai) { + return false; + } + + if (Parameters->Dimensions != 2 || Parameters->GroupCount == 0) { + return false; + } + + // The current direct kernel path processes a single batch at a time. + if (Parameters->BatchCount != 1) { + return false; + } + + if (Parameters->Beta != 0.0f) { + return false; + } + + // Depthwise convolution with multiplier 1: one input channel and one filter per group. + if (Parameters->InputChannels != 1 || Parameters->FilterCount != 1) { + return false; + } + + // Kernel specialization is 3x3 with unit stride and dilation. + if (Parameters->KernelShape[0] != 3 || Parameters->KernelShape[1] != 3) { + return false; + } + + if (Parameters->StrideShape[0] != 1 || Parameters->StrideShape[1] != 1) { + return false; + } + + if (Parameters->DilationShape[0] != 1 || Parameters->DilationShape[1] != 1) { + return false; + } + + const bool zero_padding = Parameters->Padding[0] == 0 && Parameters->Padding[1] == 0 && + Parameters->Padding[2] == 0 && Parameters->Padding[3] == 0; + const bool unit_padding = Parameters->Padding[0] == 1 && Parameters->Padding[1] == 1 && + Parameters->Padding[2] == 1 && Parameters->Padding[3] == 1; + if (!zero_padding && !unit_padding) { + return false; + } + + size_t out_height = 0; + size_t out_width = 0; + if (!TryComputeDepthwiseOutputShape(Parameters->InputShape[0], + Parameters->InputShape[1], + Parameters->KernelShape[0], + Parameters->KernelShape[1], + Parameters->Padding[0], + Parameters->Padding[1], + Parameters->Padding[2], + Parameters->Padding[3], + out_height, + out_width)) { + return false; + } + + return out_height >= dwconv.ukernel.get_m_step() && out_width >= kDwconvColsPerTile; +} + +size_t +MLASCALL +MlasDepthwiseConvPackWeightsAndBiasSize(size_t channels, + size_t filter_height, + size_t filter_width) { + const size_t packed_size = + kai_rhs_get_dst_size_dwconv_pack_x32p1vlx1b_x32_x32_sme(filter_height, filter_width, channels); + size_t total_size = 0; + return MlasAddOverflowsSizeT(sizeof(size_t), packed_size, &total_size) ? 0 : total_size; +} + +void +MLASCALL +MlasDepthwiseConvPackWeightsAndBias(size_t channels, + size_t filter_height, + size_t filter_width, + const float* weights, + const float* bias, + void* packed_weights) { + std::vector weights_hwcn(filter_height * filter_width * channels); + ConvertDepthwiseWeightsToHwcn(weights, weights_hwcn.data(), channels, filter_height, filter_width); + + std::vector zero_bias; + const float* bias_data = bias; + if (bias_data == nullptr) { + zero_bias.assign(channels, 0.0f); + bias_data = zero_bias.data(); + } + + const size_t streaming_vector_length = kai_get_sme_vector_length_u8() / sizeof(float); + std::memcpy(packed_weights, &streaming_vector_length, sizeof(streaming_vector_length)); + auto* packed_data = reinterpret_cast(packed_weights) + sizeof(streaming_vector_length); + kai_run_rhs_dwconv_pack_x32p1vlx1b_x32_x32_sme(filter_height, + filter_width, + filter_height, + filter_width, + channels, + weights_hwcn.data(), + bias_data, + packed_data); +} + +bool +MLASCALL +DepthwiseConvKleidiAI(size_t batches, + size_t in_height, + size_t in_width, + size_t channels, + size_t filter_height, + size_t filter_width, + size_t pad_top, + size_t pad_left, + size_t pad_bottom, + size_t pad_right, + bool channels_last, + const float* feature_map, + const float* weights, + const float* bias, + const void* packed_weights, + float* out, + float clamp_min, + float clamp_max, + MLAS_THREADPOOL* thread_pool) { + if (!UseSME2 || feature_map == nullptr || (weights == nullptr && packed_weights == nullptr) || out == nullptr) { + return false; + } + + if (batches != 1 || channels == 0 || filter_height != 3 || filter_width != 3) { + return false; + } + + size_t out_height = 0; + size_t out_width = 0; + if (!TryComputeDepthwiseOutputShape(in_height, + in_width, + filter_height, + filter_width, + pad_top, + pad_left, + pad_bottom, + pad_right, + out_height, + out_width)) { + return false; + } + + const size_t rows_handled = dwconv.ukernel.get_m_step(); + if (out_height < rows_handled || out_width < kDwconvColsPerTile) { + return false; + } + + const uint64_t expected_streaming_vector_length = kai_get_sme_vector_length_u8(); + const size_t expected_packed_vector_length = expected_streaming_vector_length / sizeof(float); + + auto& tls = g_dwconv_tls; + ScopedKaiDwconvTlsCleanup cleanup{tls}; + + const float* feature_map_nhwc = feature_map; + if (!channels_last) { + const size_t input_size = batches * in_height * in_width * channels; + tls.feature_map_nhwc.resize(input_size); + ConvertNchwToNhwc(feature_map, tls.feature_map_nhwc.data(), batches, channels, in_height, in_width); + feature_map_nhwc = tls.feature_map_nhwc.data(); + } + + const std::byte* packed_rhs = nullptr; + if (packed_weights != nullptr) { + size_t packed_vector_length = 0; + std::memcpy(&packed_vector_length, packed_weights, sizeof(packed_vector_length)); + if (packed_vector_length == expected_packed_vector_length) { + packed_rhs = reinterpret_cast(packed_weights) + sizeof(packed_vector_length); + } + } + + if (packed_rhs == nullptr) { + if (weights == nullptr) { + return false; + } + + const size_t weights_size = filter_height * filter_width * channels; + tls.weights_hwcn.resize(weights_size); + ConvertDepthwiseWeightsToHwcn(weights, tls.weights_hwcn.data(), channels, filter_height, filter_width); + + const float* bias_data = bias; + if (bias_data == nullptr) { + tls.bias_fallback.assign(channels, 0.0f); + bias_data = tls.bias_fallback.data(); + } + + const size_t packed_size_bytes = + kai_rhs_get_dst_size_dwconv_pack_x32p1vlx1b_x32_x32_sme(filter_height, filter_width, channels); + tls.weights_packed.resize(packed_size_bytes); + KLEIDIAI_KERNEL_LOG("kai_run_rhs_dwconv_pack_x32p1vlx1b_x32_x32_sme" + << " filter_height=" << filter_height << " filter_width=" << filter_width + << " channels=" << channels); + kai_run_rhs_dwconv_pack_x32p1vlx1b_x32_x32_sme(filter_height, + filter_width, + filter_height, + filter_width, + channels, + tls.weights_hwcn.data(), + bias_data, + tls.weights_packed.data()); + packed_rhs = tls.weights_packed.data(); + } + + float* nhwc_out = out; + if (!channels_last) { + const size_t output_size = batches * out_height * out_width * channels; + tls.nhwc_out.assign(output_size, 0.0f); + nhwc_out = tls.nhwc_out.data(); + } + + const size_t in_row_stride_elements = in_width * channels; + const size_t out_row_stride_elements = out_width * channels; + const size_t row_tile_count = MlasDivRoundup(out_height, rows_handled); + const auto run_row_tiles = [&](size_t first_tile, size_t tile_count) { + for (size_t tile = first_tile; tile < first_tile + tile_count; ++tile) { + const size_t out_row = tile * rows_handled; + const ptrdiff_t start_in_row = static_cast(out_row) - static_cast(pad_top); + const size_t kernel_pad_top = start_in_row < 0 ? static_cast(-start_in_row) : 0; + const size_t in_row = start_in_row < 0 ? 0 : static_cast(start_in_row); + + const size_t rows_to_process = std::min(rows_handled, out_height - out_row); + size_t valid_input_rows = 0; + if (in_row < in_height) { + const size_t max_rows_available = in_height - in_row; + const size_t needed_rows = filter_height + rows_to_process - 1; + valid_input_rows = std::min(max_rows_available, needed_rows); + } + + const float* inptr = feature_map_nhwc + in_row * in_row_stride_elements; + float* outptr = nhwc_out + out_row * out_row_stride_elements; + + KLEIDIAI_KERNEL_LOG(dwconv.name + << " valid_input_rows=" << valid_input_rows + << " valid_dst_rows=" << rows_to_process + << " pad_left=" << pad_left << " pad_top=" << kernel_pad_top); + dwconv.ukernel.run_dwconv(inptr, + packed_rhs, + outptr, + in_row_stride_elements * sizeof(float), + channels * sizeof(float), + out_row_stride_elements * sizeof(float), + channels * sizeof(float), + valid_input_rows, + rows_to_process, + pad_left, + kernel_pad_top, + 0.0f, + clamp_min, + clamp_max); + } + }; + + const size_t maximum_thread_count = + static_cast(std::max(1, MlasGetMaximumThreadCount(thread_pool))); + const double complexity = static_cast(out_height) * static_cast(out_width) * + static_cast(channels) * static_cast(filter_height * filter_width); + const double maximum_threaded_complexity = + static_cast(maximum_thread_count) * static_cast(MLAS_SGEMM_THREAD_COMPLEXITY); + const size_t complexity_thread_count = + complexity >= maximum_threaded_complexity + ? maximum_thread_count + : static_cast(complexity / static_cast(MLAS_SGEMM_THREAD_COMPLEXITY)) + 1; + const size_t thread_count = std::min({row_tile_count, maximum_thread_count, complexity_thread_count}); + + if (thread_count == 1) { + run_row_tiles(0, row_tile_count); + } else { + std::atomic vector_length_mismatch{false}; + MlasTrySimpleParallel(thread_pool, static_cast(thread_count), [&](ptrdiff_t thread_id) { + if (kai_get_sme_vector_length_u8() != expected_streaming_vector_length) { + vector_length_mismatch.store(true, std::memory_order_relaxed); + return; + } + + size_t first_tile = 0; + size_t tile_count = 0; + MlasPartitionWork(thread_id, static_cast(thread_count), row_tile_count, &first_tile, &tile_count); + run_row_tiles(first_tile, tile_count); + }); + + if (vector_length_mismatch.load(std::memory_order_relaxed)) { + // Matching workers may have written disjoint rows, so overwrite the complete output on the caller thread. + run_row_tiles(0, row_tile_count); + } + } + + if (!channels_last) { + ConvertNhwcToNchw(nhwc_out, out, batches, channels, out_height, out_width); + } + + return true; +} + +} // namespace ArmKleidiAI diff --git a/onnxruntime/core/mlas/lib/kleidiai/kai_ukernel_interface.cpp b/onnxruntime/core/mlas/lib/kleidiai/kai_ukernel_interface.cpp index de7a5d92e59e6..48d44f1b1a3a9 100644 --- a/onnxruntime/core/mlas/lib/kleidiai/kai_ukernel_interface.cpp +++ b/onnxruntime/core/mlas/lib/kleidiai/kai_ukernel_interface.cpp @@ -54,6 +54,8 @@ // FP16 HGEMM kernels #include "kai/ukernels/matmul/matmul_clamp_f16_f16_f16p/kai_matmul_clamp_f16_f16_f16p2vlx2b_1x8vl_sme_mla.h" #include "kai/ukernels/matmul/matmul_clamp_f16_f16_f16p/kai_matmul_clamp_f16_f16_f16p2vlx2b_1x16vl_sme2_dot.h" +// DWCONV +#include "kai/ukernels/dwconv/dwconv_f32_f32_f32p/kai_dwconv_clamp_f32_f32_f32p1vlx1b_3x3_s1_4xc_sme2_mla.h" #if defined(ENABLE_QMX_KERNELS) // QMX kernels (optional) @@ -305,6 +307,9 @@ const KaiF16IMatmulKernel imatmul_f16_conv_sme2 = const KaiBF16SBgemmKernel sbgemm_gemm_sme2 = KAI_WRAP_UKERNEL_RUN_MATMUL_11(matmul_clamp_f32_bf16p2vlx2_bf16p2vlx2_2vlx2vl_sme2_mopa); +const KaiF32DepthwiseConvKernel dwconv_sme2 = + KAI_WRAP_UKERNEL_RUN_DWCONV_PLANAR_4(dwconv_clamp_f32_f32_f32p1vlx1b_3x3_s1_4xc_sme2_mla); + #if defined(ENABLE_QMX_KERNELS) const KaiF32IMatmulKernel imatmul_conv_qmx = KAI_WRAP_UKERNEL_RUN_IMATMUL_PACKED_7(imatmul_clamp_f32_f32p2vlx1_f32p2vlx1b_2vlx2vl_qmx_mopa); @@ -493,3 +498,8 @@ const KaiF16HgemmKernel& GetKleidiAIHgemmUKernel() { return hgemm_sme; } } + +const KaiF32DepthwiseConvKernel& GetKleidiAIDepthwiseConvUKernel() { + // Currently only an SME2 variant exists for FP32 depthwise convolution. + return dwconv_sme2; +} diff --git a/onnxruntime/core/mlas/lib/kleidiai/kai_ukernel_interface.h b/onnxruntime/core/mlas/lib/kleidiai/kai_ukernel_interface.h index f7a2e526f177b..bef87dedf74f5 100644 --- a/onnxruntime/core/mlas/lib/kleidiai/kai_ukernel_interface.h +++ b/onnxruntime/core/mlas/lib/kleidiai/kai_ukernel_interface.h @@ -24,14 +24,18 @@ #include "kai/ukernels/matmul/matmul_clamp_f16_f16_f16p/kai_matmul_clamp_f16_f16_f16p_interface.h" +#include "kai/ukernels/dwconv/dwconv_f32_f32_f32p/kai_dwconv_clamp_f32_f32_f32p_interface.h" + // Wrapper type that carries a stable "name" alongside the KAI ukernel interface. // This avoids needing to infer which underlying microkernel was selected from a function pointer. template -struct KaiMatmulKernel { +struct KaiKernel { const char* name; UkernelFn ukernel; }; +using KaiF32DepthwiseConvKernel = KaiKernel; + enum class KaiQ4RhsPackLayout { SymmetricNxK, SymmetricNxKInterleavedNrx4, @@ -47,10 +51,10 @@ struct KaiQ4MatmulKernel { }; // Wrapper for FP32 GEMM kernels where both LHS and RHS are pre-packed (common SGEMM path). -using KaiF32SgemmKernel = KaiMatmulKernel; +using KaiF32SgemmKernel = KaiKernel; // Wrapper for FP32 kernels used for GEMV-style workloads (typically a single-row/skinny-M use case). -using KaiF32SgemvKernel = KaiMatmulKernel; +using KaiF32SgemvKernel = KaiKernel; // Wrapper for Qnbit GEMM kernels producing FP32 output. using KaiQnbitGemmKernel = KaiQ4MatmulKernel; @@ -59,18 +63,18 @@ using KaiQnbitGemmKernel = KaiQ4MatmulKernel; // Wrapper for dynamic-quantized GEMM kernels producing FP32 output. -using KaiDynamicQGemmKernel = KaiMatmulKernel; +using KaiDynamicQGemmKernel = KaiKernel; // Wrapper for FP32 IMATMUL kernels used by the KleidiAI convolution implementation. -using KaiF32IMatmulKernel = KaiMatmulKernel; +using KaiF32IMatmulKernel = KaiKernel; // Wrapper for FP16 IMATMUL kernels used by the KleidiAI convolution implementation. -using KaiF16IMatmulKernel = KaiMatmulKernel; +using KaiF16IMatmulKernel = KaiKernel; -using KaiBF16SBgemmKernel = KaiMatmulKernel; +using KaiBF16SBgemmKernel = KaiKernel; // Wrapper for FP16 HGEMM kernels producing FP16 output. -using KaiF16HgemmKernel = KaiMatmulKernel; +using KaiF16HgemmKernel = KaiKernel; // Returns the selected Qnbit GEMM ukernel based on runtime CPU capabilities. const KaiQnbitGemmKernel& GetKleidiAIGemmUKernel(); @@ -104,3 +108,5 @@ const KaiBF16SBgemmKernel& GetKleidiAISBGemmUKernel(); // Returns the selected FP16 HGEMM ukernel based on runtime CPU capabilities. const KaiF16HgemmKernel& GetKleidiAIHgemmUKernel(); +// Returns the selected FP32 depthwise convolution ukernel based on runtime CPU capabilities. +const KaiF32DepthwiseConvKernel& GetKleidiAIDepthwiseConvUKernel(); diff --git a/onnxruntime/core/mlas/lib/kleidiai/mlasi_kleidiai.h b/onnxruntime/core/mlas/lib/kleidiai/mlasi_kleidiai.h index 221482c71ec1d..d29bf95744498 100644 --- a/onnxruntime/core/mlas/lib/kleidiai/mlasi_kleidiai.h +++ b/onnxruntime/core/mlas/lib/kleidiai/mlasi_kleidiai.h @@ -68,9 +68,16 @@ inline const bool UseSME2 = MLAS_CPUIDINFO::GetCPUIDInfo().HasArm_SME2(); inline const bool UseSME = MLAS_CPUIDINFO::GetCPUIDInfo().HasArm_SME(); inline const std::string_view vendor_name = MLAS_CPUIDINFO::GetCPUIDInfo().GetCPUVendor(); +bool +MLASCALL +DepthwiseConvKleidiAISupported( + const MLAS_CONV_PARAMETERS* Parameters + ); + // Selects the convolution route for Arm® KleidiAI™ enum class ConvRoute { NoKleidiAi, // decline the conv, caller runs unchanged + Depthwise, // handle the whole conv via SME2 depthwise kernel IGemm, // handle the whole conv via SME IGEMM kernel SGemmFallback, // decline IGEMM, but still route the per-segment SGEMM slices through MlasGemm // so that the Arm® KleidiAI™ SGEMM backend override can pick them up @@ -129,6 +136,12 @@ inline ConvRouteSelection SelectConvRoute(const MLAS_CONV_PARAMETERS* Parameters return ConvRouteSelection{}; } + if (Parameters->InputChannels == 1 && Parameters->FilterCount == 1) { + return DepthwiseConvKleidiAISupported(Parameters) + ? ConvRouteSelection{ConvRoute::Depthwise, Parameters->KernelShape[0], Parameters->KernelShape[1]} + : ConvRouteSelection{}; + } + size_t effective_kernel_h; size_t effective_kernel_w; if (!TryComputeDilatedKernelSize(Parameters->DilationShape[0], Parameters->KernelShape[0], &effective_kernel_h) || @@ -427,6 +440,49 @@ MlasConv( MLAS_THREADPOOL* ThreadPool ); +bool +MLASCALL +DepthwiseConvKleidiAI( + size_t batches, + size_t in_height, + size_t in_width, + size_t channels, + size_t filter_height, + size_t filter_width, + size_t pad_top, + size_t pad_left, + size_t pad_bottom, + size_t pad_right, + bool channels_last, + const float* feature_map, + const float* weights, + const float* bias, + const void* packed_weights, + float* out, + float clamp_min, + float clamp_max, + MLAS_THREADPOOL* thread_pool + ); + +size_t +MLASCALL +MlasDepthwiseConvPackWeightsAndBiasSize( + size_t channels, + size_t filter_height, + size_t filter_width + ); + +void +MLASCALL +MlasDepthwiseConvPackWeightsAndBias( + size_t channels, + size_t filter_height, + size_t filter_width, + const float* weights, + const float* bias, + void* packed_weights + ); + size_t MLASCALL MlasConvSymmetricChannelsLast2DFloatPackWSize( diff --git a/onnxruntime/core/providers/cpu/nn/conv.cc b/onnxruntime/core/providers/cpu/nn/conv.cc index bccf67fd68ab2..e2af3488b92ea 100644 --- a/onnxruntime/core/providers/cpu/nn/conv.cc +++ b/onnxruntime/core/providers/cpu/nn/conv.cc @@ -14,6 +14,7 @@ * limitations under the License. */ /* Modifications Copyright (c) Microsoft. */ +// SPDX-FileCopyrightText: Copyright 2026 Arm Limited and/or its affiliates #include @@ -199,7 +200,8 @@ Status Conv::EnsurePackedChannelsLastFilter(concurrency::ThreadPool* thre size_t filter_count_per_group, size_t input_channels_per_group, const TensorShapeVector& kernel_shape, - const TensorShapeVector& dilations) const { + const TensorShapeVector& dilations, + bool is_depthwise) const { if (!can_cache_packed_filter_) { return Status::OK(); } @@ -214,32 +216,53 @@ Status Conv::EnsurePackedChannelsLastFilter(concurrency::ThreadPool* thre return; } - packed_filter_group_stride_ = - ArmKleidiAI::MlasConvSymmetricChannelsLast2DFloatPackWSize(filter_count_per_group, - input_channels_per_group, - kernel_shape.data(), - dilations.data()); - if (packed_filter_group_stride_ == 0) { + size_t packed_filter_size = 0; + if (is_depthwise) { + packed_filter_size = ArmKleidiAI::MlasDepthwiseConvPackWeightsAndBiasSize( + onnxruntime::narrow(conv_attrs_.group), + onnxruntime::narrow(kernel_shape[0]), + onnxruntime::narrow(kernel_shape[1])); + } else { + packed_filter_group_stride_ = + ArmKleidiAI::MlasConvSymmetricChannelsLast2DFloatPackWSize(filter_count_per_group, + input_channels_per_group, + kernel_shape.data(), + dilations.data()); + packed_filter_size = SafeInt(packed_filter_group_stride_) * + onnxruntime::narrow(conv_attrs_.group); + } + + if (packed_filter_size == 0) { packed_filter_status_ = ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Failed to get KleidiAI packed filter size."); return; } - const size_t packed_filter_size = - packed_filter_group_stride_ * onnxruntime::narrow(conv_attrs_.group); packed_filter_ = IAllocator::MakeUniquePtr(alloc, packed_filter_size, true); memset(packed_filter_.get(), 0, packed_filter_size); - ArmKleidiAI::MlasConvSymmetricChannelsLast2DFloatPackW(filter_count_per_group, - input_channels_per_group, - kernel_shape.data(), - dilations.data(), - onnxruntime::narrow(conv_attrs_.group), - constant_filter_tensor_->Data(), - constant_bias_tensor_ ? constant_bias_tensor_->Data() : nullptr, - packed_filter_.get(), - packed_filter_group_stride_, - thread_pool); + if (is_depthwise) { + ArmKleidiAI::MlasDepthwiseConvPackWeightsAndBias( + onnxruntime::narrow(conv_attrs_.group), + onnxruntime::narrow(kernel_shape[0]), + onnxruntime::narrow(kernel_shape[1]), + constant_filter_tensor_->Data(), + constant_bias_tensor_ ? constant_bias_tensor_->Data() : nullptr, + packed_filter_.get()); + } else { + ArmKleidiAI::MlasConvSymmetricChannelsLast2DFloatPackW(filter_count_per_group, + input_channels_per_group, + kernel_shape.data(), + dilations.data(), + onnxruntime::narrow(conv_attrs_.group), + constant_filter_tensor_->Data(), + constant_bias_tensor_ + ? constant_bias_tensor_->Data() + : nullptr, + packed_filter_.get(), + packed_filter_group_stride_, + thread_pool); + } }); return packed_filter_status_; @@ -324,31 +347,14 @@ Status Conv::Compute(OpKernelContext* context) const { const size_t group_count = narrow(conv_attrs_.group); const size_t input_channels_per_group = narrow(C / conv_attrs_.group); const size_t filter_count_per_group = narrow(M / conv_attrs_.group); - const bool nhwc_fastpath = - wants_channels_last && !sum_present && - (MlasConvSupportsDenseChannelsLast2DFloatKernel( - kernel_rank, - narrow(N), - group_count, - input_shape_size_t.data(), - kernel_shape_size_t.data(), - dilations_size_t.data(), - pads_size_t.data(), - strides_size_t.data(), - filter_count_per_group, - /*Beta*/ 0.0f) || - MlasConvSupportsDepthwiseChannelsLast2DFloatKernel( - kernel_rank, - narrow(N), - group_count, - input_channels_per_group, - input_shape_size_t.data(), - kernel_shape_size_t.data(), - dilations_size_t.data(), - pads_size_t.data(), - strides_size_t.data(), - filter_count_per_group, - /*Beta*/ 0.0f)); + const bool dense_nhwc_fastpath = + wants_channels_last && !sum_present && MlasConvSupportsDenseChannelsLast2DFloatKernel(kernel_rank, narrow(N), group_count, input_shape_size_t.data(), kernel_shape_size_t.data(), dilations_size_t.data(), pads_size_t.data(), strides_size_t.data(), filter_count_per_group, + /*Beta*/ 0.0f); + const bool depthwise_nhwc_fastpath = + wants_channels_last && !sum_present && MlasConvSupportsDepthwiseChannelsLast2DFloatKernel(kernel_rank, narrow(N), group_count, input_channels_per_group, input_shape_size_t.data(), kernel_shape_size_t.data(), dilations_size_t.data(), pads_size_t.data(), strides_size_t.data(), filter_count_per_group, + /*Beta*/ 0.0f); + const bool nhwc_fastpath = mlas_backend_kernel_selector_config_.use_kleidiai && + (dense_nhwc_fastpath || depthwise_nhwc_fastpath); #if defined(USE_KLEIDIAI) && defined(MLAS_TARGET_ARM64) if (nhwc_fastpath && can_cache_packed_filter_) { @@ -356,7 +362,8 @@ Status Conv::Compute(OpKernelContext* context) const { filter_count_per_group, input_channels_per_group, kernel_shape, - dilations)); + dilations, + depthwise_nhwc_fastpath)); } #endif diff --git a/onnxruntime/core/providers/cpu/nn/conv.h b/onnxruntime/core/providers/cpu/nn/conv.h index 9e073df545328..93c81da91d7ad 100644 --- a/onnxruntime/core/providers/cpu/nn/conv.h +++ b/onnxruntime/core/providers/cpu/nn/conv.h @@ -1,4 +1,5 @@ // Copyright (c) Microsoft Corporation. All rights reserved. +// SPDX-FileCopyrightText: Copyright 2026 Arm Limited and/or its affiliates // Licensed under the MIT License. #pragma once @@ -62,7 +63,8 @@ class Conv : public OpKernel { size_t filter_count_per_group, size_t input_channels_per_group, const TensorShapeVector& kernel_shape, - const TensorShapeVector& dilations) const; + const TensorShapeVector& dilations, + bool is_depthwise) const; const Tensor* constant_filter_tensor_{nullptr}; const Tensor* constant_bias_tensor_{nullptr}; diff --git a/onnxruntime/test/contrib_ops/fused_conv_test.cc b/onnxruntime/test/contrib_ops/fused_conv_test.cc index 48e26fa80e4e4..915e949f59a85 100644 --- a/onnxruntime/test/contrib_ops/fused_conv_test.cc +++ b/onnxruntime/test/contrib_ops/fused_conv_test.cc @@ -7,6 +7,7 @@ #include "core/common/narrow.h" #include "core/framework/kernel_registry.h" +#include "core/session/onnxruntime_session_options_config_keys.h" #include "test/common/cuda_op_test_utils.h" #include "test/common/tensor_op_test_utils.h" #include "test/providers/provider_test_utils.h" @@ -562,6 +563,59 @@ TEST(FusedConvTest, Cpu_NhwcDepthwiseConv2D_Relu_NegativePreActivation) { TestNhwcFusedConvFloatOp(attrs, {X, W}, {X_shape, W_shape}, expected_vals, Y_shape, true); #endif } + +TEST(FusedConvTest, Cpu_NhwcDepthwiseConv2D_KleidiAiDisabledFallback) { +#if !defined(MLAS_TARGET_ARM64) + GTEST_SKIP() << "Float NHWC depthwise fast-path requires Arm64."; +#else + constexpr int64_t height = 8; + constexpr int64_t width = 8; + constexpr int64_t channels = 2; + const vector input_shape = {1, height, width, channels}; + const vector weight_shape = {channels, 1, 3, 3}; + const vector pads = {1, 1, 1, 1}; + const vector strides = {1, 1}; + + if (!HasFloatNhwcNoTransposeSupport(input_shape, weight_shape, pads, strides, channels)) { + GTEST_SKIP() << "Float NHWC depthwise fast-path is not available on this configuration."; + } + + vector input; + vector expected_output; + input.reserve(height * width * channels); + expected_output.reserve(height * width * channels); + for (int64_t pixel = 0; pixel < height * width; ++pixel) { + const float channel_zero = static_cast(pixel + 1); + const float channel_one = static_cast(pixel + 101); + input.push_back(channel_zero); + input.push_back(channel_one); + expected_output.push_back(2.0f * channel_zero); + expected_output.push_back(3.0f * channel_one); + } + + vector weights(channels * 3 * 3, 0.0f); + weights[4] = 2.0f; + weights[9 + 4] = 3.0f; + + OpTester test("NhwcFusedConv", 1, onnxruntime::kMSDomain); + test.AddAttribute("group", channels); + test.AddAttribute("kernel_shape", vector{3, 3}); + test.AddAttribute("pads", pads); + test.AddAttribute("strides", strides); + test.AddAttribute("activation", "Relu"); + test.AddInput("X", input_shape, input); + test.AddInput("W", weight_shape, weights, true); + test.AddOutput("Y", input_shape, expected_output); + + SessionOptions session_options; + ASSERT_STATUS_OK(session_options.config_options.AddConfigEntry(kOrtSessionOptionsMlasDisableKleidiAi, "1")); + test.Config(session_options); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCpuExecutionProvider()); + test.ConfigEps(std::move(execution_providers)).RunWithConfig(); +#endif +} #endif TEST(FusedConvTest, Cpu_Conv3D_Batched_Relu) { diff --git a/onnxruntime/test/mlas/unittest/test_dwconv_kleidiai.cpp b/onnxruntime/test/mlas/unittest/test_dwconv_kleidiai.cpp new file mode 100644 index 0000000000000..7899939932a64 --- /dev/null +++ b/onnxruntime/test/mlas/unittest/test_dwconv_kleidiai.cpp @@ -0,0 +1,468 @@ +// +// SPDX-FileCopyrightText: Copyright 2025-2026 Arm Limited and/or its affiliates +// +// SPDX-License-Identifier: MIT +// + +#if defined(USE_KLEIDIAI) && !defined(_MSC_VER) + +#include "test_util.h" +#include "core/mlas/lib/kleidiai/mlasi_kleidiai.h" + +#include +#include +#include +#include +#include +#include + +namespace { + +void ConvertNchwToNhwc(const float* src, + float* dst, + size_t batches, + size_t channels, + size_t height, + size_t width) { + const size_t src_stride_n = channels * height * width; + const size_t src_stride_c = height * width; + const size_t src_stride_h = width; + const size_t dst_stride_n = height * width * channels; + const size_t dst_stride_h = width * channels; + const size_t dst_stride_w = channels; + + for (size_t n = 0; n < batches; ++n) { + for (size_t h = 0; h < height; ++h) { + for (size_t w = 0; w < width; ++w) { + for (size_t c = 0; c < channels; ++c) { + const size_t src_index = n * src_stride_n + c * src_stride_c + h * src_stride_h + w; + const size_t dst_index = n * dst_stride_n + h * dst_stride_h + w * dst_stride_w + c; + dst[dst_index] = src[src_index]; + } + } + } + } +} + +void DepthwiseReferenceNchw(const float* input, + const float* weights, + const float* bias, + size_t batches, + size_t channels, + size_t in_height, + size_t in_width, + size_t filter_height, + size_t filter_width, + size_t pad_top, + size_t pad_left, + size_t pad_bottom, + size_t pad_right, + float clamp_min, + float clamp_max, + float* output) { + const size_t out_height = in_height + pad_top + pad_bottom + 1 - filter_height; + const size_t out_width = in_width + pad_left + pad_right + 1 - filter_width; + + for (size_t b = 0; b < batches; ++b) { + for (size_t c = 0; c < channels; ++c) { + for (size_t oh = 0; oh < out_height; ++oh) { + for (size_t ow = 0; ow < out_width; ++ow) { + float acc = bias != nullptr ? bias[c] : 0.0f; + for (size_t kh = 0; kh < filter_height; ++kh) { + const int in_y = static_cast(oh) + static_cast(kh) - static_cast(pad_top); + if (in_y < 0 || in_y >= static_cast(in_height)) { + continue; + } + for (size_t kw = 0; kw < filter_width; ++kw) { + const int in_x = static_cast(ow) + static_cast(kw) - static_cast(pad_left); + if (in_x < 0 || in_x >= static_cast(in_width)) { + continue; + } + const size_t input_idx = + (((b * channels) + c) * in_height + static_cast(in_y)) * in_width + + static_cast(in_x); + const size_t weight_idx = (c * filter_height + kh) * filter_width + kw; + acc += input[input_idx] * weights[weight_idx]; + } + } + const size_t output_idx = (((b * channels) + c) * out_height + oh) * out_width + ow; + output[output_idx] = std::clamp(acc, clamp_min, clamp_max); + } + } + } + } +} + +void RunDepthwiseConvCase(size_t channels, size_t in_height, size_t in_width, size_t padding, bool channels_last) { + if (!MLAS_CPUIDINFO::GetCPUIDInfo().HasArm_SME2()) { + GTEST_SKIP() << "DepthwiseConvKleidiAI requires ARM64 SME2. Skipping test."; + } + + constexpr size_t batches = 1; + constexpr size_t filter_height = 3; + constexpr size_t filter_width = 3; + + const size_t pad_top = padding; + const size_t pad_left = padding; + const size_t pad_bottom = padding; + const size_t pad_right = padding; + + const size_t input_size = batches * channels * in_height * in_width; + const size_t weights_size = channels * filter_height * filter_width; + const size_t out_height = in_height + pad_top + pad_bottom + 1 - filter_height; + const size_t out_width = in_width + pad_left + pad_right + 1 - filter_width; + const size_t output_size = batches * channels * out_height * out_width; + + std::vector input(input_size); + std::vector weights(weights_size); + std::vector bias(channels); + std::vector input_nhwc(input_size); + std::vector expected(output_size); + std::vector expected_nhwc(output_size); + std::vector output(output_size, std::numeric_limits::quiet_NaN()); + + std::mt19937 rng(static_cast(channels * 131 + in_height * 17 + padding)); + std::uniform_real_distribution dist(-1.0f, 1.0f); + + auto fill_buffer = [&](std::vector& buffer) { + for (float& v : buffer) { + v = dist(rng); + } + }; + + fill_buffer(input); + fill_buffer(weights); + fill_buffer(bias); + + const float clamp_min = -std::numeric_limits::max(); + const float clamp_max = std::numeric_limits::max(); + + DepthwiseReferenceNchw(input.data(), + weights.data(), + bias.data(), + batches, + channels, + in_height, + in_width, + filter_height, + filter_width, + pad_top, + pad_left, + pad_bottom, + pad_right, + clamp_min, + clamp_max, + expected.data()); + + ConvertNchwToNhwc(input.data(), input_nhwc.data(), batches, channels, in_height, in_width); + ConvertNchwToNhwc(expected.data(), expected_nhwc.data(), batches, channels, out_height, out_width); + + const float* input_data = channels_last ? input_nhwc.data() : input.data(); + const float* expected_data = channels_last ? expected_nhwc.data() : expected.data(); + + const bool status = ArmKleidiAI::DepthwiseConvKleidiAI(batches, + in_height, + in_width, + channels, + filter_height, + filter_width, + pad_top, + pad_left, + pad_bottom, + pad_right, + channels_last, + input_data, + weights.data(), + bias.data(), + nullptr, + output.data(), + clamp_min, + clamp_max, + nullptr); + ASSERT_TRUE(status); + + for (size_t i = 0; i < output_size; ++i) { + EXPECT_NEAR(expected_data[i], output[i], 1e-4f) << "Mismatch at element " << i; + } +} + +void RunMlasConvChannelsLastCase(size_t channels, + size_t in_height, + size_t in_width, + size_t padding, + bool use_prepacked_weights = false, + MLAS_THREADPOOL* thread_pool = nullptr) { + if (!MLAS_CPUIDINFO::GetCPUIDInfo().HasArm_SME2()) { + GTEST_SKIP() << "DepthwiseConvKleidiAI requires ARM64 SME2. Skipping test."; + } + + constexpr size_t batches = 1; + constexpr size_t filter_height = 3; + constexpr size_t filter_width = 3; + constexpr size_t filter_count = 1; + + const size_t out_height = in_height + padding + padding + 1 - filter_height; + const size_t out_width = in_width + padding + padding + 1 - filter_width; + const size_t input_size = batches * channels * in_height * in_width; + const size_t output_size = batches * channels * out_height * out_width; + + std::vector input(input_size); + std::vector input_nhwc(input_size); + std::vector weights(channels * filter_height * filter_width); + std::vector bias(channels); + std::vector expected(output_size); + std::vector expected_nhwc(output_size); + std::vector output(output_size, std::numeric_limits::quiet_NaN()); + + std::mt19937 rng(static_cast(channels * 197 + in_width * 23 + padding)); + std::uniform_real_distribution dist(-1.0f, 1.0f); + for (float& v : input) { + v = dist(rng); + } + for (float& v : weights) { + v = dist(rng); + } + for (float& v : bias) { + v = dist(rng); + } + + DepthwiseReferenceNchw(input.data(), + weights.data(), + bias.data(), + batches, + channels, + in_height, + in_width, + filter_height, + filter_width, + padding, + padding, + padding, + padding, + -std::numeric_limits::max(), + std::numeric_limits::max(), + expected.data()); + ConvertNchwToNhwc(input.data(), input_nhwc.data(), batches, channels, in_height, in_width); + ConvertNchwToNhwc(expected.data(), expected_nhwc.data(), batches, channels, out_height, out_width); + + const int64_t input_shape[] = {static_cast(in_height), static_cast(in_width)}; + const int64_t kernel_shape[] = {static_cast(filter_height), static_cast(filter_width)}; + const int64_t dilation_shape[] = {1, 1}; + const int64_t pads[] = {static_cast(padding), static_cast(padding), + static_cast(padding), static_cast(padding)}; + const int64_t strides[] = {1, 1}; + const int64_t output_shape[] = {static_cast(out_height), static_cast(out_width)}; + MLAS_ACTIVATION activation; + activation.ActivationKind = MlasIdentityActivation; + + MLAS_CONV_PARAMETERS parameters{}; + size_t working_buffer_size = 0; + MlasConvPrepare(¶meters, + 2, + batches, + channels, + 1, + input_shape, + kernel_shape, + dilation_shape, + pads, + strides, + output_shape, + filter_count, + &activation, + &working_buffer_size, + true, + 0.0f, + thread_pool); + std::vector working_buffer(working_buffer_size); + std::vector packed_weights; + if (use_prepacked_weights) { + const size_t packed_size = + ArmKleidiAI::MlasDepthwiseConvPackWeightsAndBiasSize(channels, filter_height, filter_width); + ASSERT_NE(packed_size, 0U); + packed_weights.resize(packed_size); + ArmKleidiAI::MlasDepthwiseConvPackWeightsAndBias( + channels, filter_height, filter_width, weights.data(), bias.data(), packed_weights.data()); + parameters.FilterIsPacked = true; + parameters.PackedFilter = packed_weights.data(); + + std::fill(weights.begin(), weights.end(), std::numeric_limits::quiet_NaN()); + std::fill(bias.begin(), bias.end(), std::numeric_limits::quiet_NaN()); + } + + MlasConv(¶meters, + input_nhwc.data(), + weights.data(), + bias.data(), + working_buffer.data(), + output.data(), + thread_pool); + + for (size_t i = 0; i < output_size; ++i) { + EXPECT_NEAR(expected_nhwc[i], output[i], 1e-4f) << "Mismatch at element " << i; + } +} + +} // namespace + +TEST(MlasKleidiDepthwiseTest, ZeroPaddingNchw) { + RunDepthwiseConvCase(/*channels=*/32, /*in_height=*/8, /*in_width=*/8, /*padding=*/0, /*channels_last=*/false); +} + +TEST(MlasKleidiDepthwiseTest, UnitPaddingNchw) { + RunDepthwiseConvCase(/*channels=*/32, /*in_height=*/8, /*in_width=*/8, /*padding=*/1, /*channels_last=*/false); +} + +TEST(MlasKleidiDepthwiseTest, ZeroPaddingNhwc) { + RunDepthwiseConvCase(/*channels=*/32, /*in_height=*/8, /*in_width=*/8, /*padding=*/0, /*channels_last=*/true); +} + +TEST(MlasKleidiDepthwiseTest, UnitPaddingNhwc) { + RunDepthwiseConvCase(/*channels=*/32, /*in_height=*/8, /*in_width=*/8, /*padding=*/1, /*channels_last=*/true); +} + +TEST(MlasKleidiDepthwiseTest, LargeChannelCountNhwc) { + RunDepthwiseConvCase(/*channels=*/128, /*in_height=*/8, /*in_width=*/8, /*padding=*/1, /*channels_last=*/true); +} + +TEST(MlasKleidiDepthwiseTest, ChannelTailNhwc) { + RunDepthwiseConvCase(/*channels=*/129, /*in_height=*/8, /*in_width=*/8, /*padding=*/1, /*channels_last=*/true); +} + +TEST(MlasKleidiDepthwiseTest, ThreadedRowTilesWithTailNhwc) { + ASSERT_NE(GetMlasThreadPool(), nullptr); + RunMlasConvChannelsLastCase(/*channels=*/129, /*in_height=*/17, /*in_width=*/17, /*padding=*/1, + /*use_prepacked_weights=*/true, GetMlasThreadPool()); +} + +TEST(MlasKleidiDepthwiseTest, MlasConvChannelsLast) { + RunMlasConvChannelsLastCase(/*channels=*/32, /*in_height=*/8, /*in_width=*/8, /*padding=*/1); +} + +TEST(MlasKleidiDepthwiseTest, MlasConvChannelsLastPrepackedWeights) { + RunMlasConvChannelsLastCase( + /*channels=*/32, /*in_height=*/8, /*in_width=*/8, /*padding=*/1, /*use_prepacked_weights=*/true); +} + +TEST(MlasKleidiDepthwiseTest, RejectsOutputSmallerThanKernelTile) { + if (!MLAS_CPUIDINFO::GetCPUIDInfo().HasArm_SME2()) { + GTEST_SKIP() << "DepthwiseConvKleidiAI requires ARM64 SME2. Skipping test."; + } + + MLAS_CONV_PARAMETERS parameters{}; + parameters.Dimensions = 2; + parameters.BatchCount = 1; + parameters.GroupCount = 32; + parameters.InputChannels = 1; + parameters.FilterCount = 1; + parameters.InputShape[0] = 5; + parameters.InputShape[1] = 5; + parameters.KernelShape[0] = 3; + parameters.KernelShape[1] = 3; + parameters.DilationShape[0] = 1; + parameters.DilationShape[1] = 1; + parameters.StrideShape[0] = 1; + parameters.StrideShape[1] = 1; + + EXPECT_FALSE(ArmKleidiAI::DepthwiseConvKleidiAISupported(¶meters)); +} + +TEST(MlasKleidiDepthwiseTest, BatchCountTwoFallsBackToGenericMlas) { + if (!MLAS_CPUIDINFO::GetCPUIDInfo().HasArm_SME2()) { + GTEST_SKIP() << "DepthwiseConvKleidiAI requires ARM64 SME2. Skipping test."; + } + + constexpr size_t batches = 2; + constexpr size_t channels = 8; + constexpr size_t in_height = 8; + constexpr size_t in_width = 8; + constexpr size_t filter_height = 3; + constexpr size_t filter_width = 3; + constexpr size_t padding = 1; + constexpr size_t out_height = in_height; + constexpr size_t out_width = in_width; + + std::vector input(batches * channels * in_height * in_width); + std::vector weights(channels * filter_height * filter_width); + std::vector bias(channels); + std::vector expected(batches * channels * out_height * out_width); + std::vector output(expected.size(), std::numeric_limits::quiet_NaN()); + std::mt19937 rng(42); + std::uniform_real_distribution dist(-1.0f, 1.0f); + std::generate(input.begin(), input.end(), [&] { return dist(rng); }); + std::generate(weights.begin(), weights.end(), [&] { return dist(rng); }); + std::generate(bias.begin(), bias.end(), [&] { return dist(rng); }); + + DepthwiseReferenceNchw(input.data(), weights.data(), bias.data(), batches, channels, + in_height, in_width, filter_height, filter_width, + padding, padding, padding, padding, + -std::numeric_limits::max(), std::numeric_limits::max(), expected.data()); + + const int64_t input_shape[] = {in_height, in_width}; + const int64_t kernel_shape[] = {filter_height, filter_width}; + const int64_t dilation_shape[] = {1, 1}; + const int64_t pads[] = {padding, padding, padding, padding}; + const int64_t strides[] = {1, 1}; + const int64_t output_shape[] = {out_height, out_width}; + MLAS_ACTIVATION activation{}; + activation.ActivationKind = MlasIdentityActivation; + MLAS_CONV_PARAMETERS parameters{}; + size_t working_buffer_size = 0; + MlasConvPrepare(¶meters, 2, batches, channels, 1, input_shape, kernel_shape, + dilation_shape, pads, strides, output_shape, 1, &activation, + &working_buffer_size, false, 0.0f, nullptr); + EXPECT_FALSE(ArmKleidiAI::DepthwiseConvKleidiAISupported(¶meters)); + + std::vector working_buffer(working_buffer_size); + MlasConv(¶meters, input.data(), weights.data(), bias.data(), working_buffer.data(), output.data(), nullptr); + for (size_t i = 0; i < output.size(); ++i) { + EXPECT_NEAR(expected[i], output[i], 1e-4f) << "Mismatch at element " << i; + } +} + +TEST(MlasKleidiDepthwiseTest, SelectorDisabledDeclinesPrepare) { + const int64_t shape[] = {8, 8}; + const int64_t kernel_shape[] = {3, 3}; + const int64_t unit_shape[] = {1, 1}; + const int64_t padding[] = {1, 1, 1, 1}; + MLAS_ACTIVATION activation{}; + activation.ActivationKind = MlasIdentityActivation; + MLAS_BACKEND_KERNEL_SELECTOR_CONFIG selector_config{}; + selector_config.use_kleidiai = false; + MLAS_CONV_PARAMETERS parameters{}; + parameters.BackendKernelSelectorConfig = &selector_config; + size_t working_buffer_size = 0; + + EXPECT_FALSE(ArmKleidiAI::MlasConvPrepare( + ¶meters, 2, 1, 32, 1, shape, kernel_shape, unit_shape, padding, unit_shape, + shape, 1, &activation, &working_buffer_size, true, 0.0f, nullptr)); +} + +TEST(MlasKleidiDepthwiseTest, RejectsNon3x3Filters) { + if (!MLAS_CPUIDINFO::GetCPUIDInfo().HasArm_SME2()) { + GTEST_SKIP() << "DepthwiseConvKleidiAI requires ARM64 SME2. Skipping test."; + } + + constexpr size_t batches = 1; + constexpr size_t in_height = 8; + constexpr size_t in_width = 8; + constexpr size_t channels = 32; + + std::vector input(batches * in_height * in_width * channels); + std::vector weights(channels * 3 * 3); + std::vector bias(channels); + std::vector output(batches * in_height * in_width * channels); + + EXPECT_FALSE(ArmKleidiAI::DepthwiseConvKleidiAI( + batches, in_height, in_width, channels, /*filter_height=*/2, /*filter_width=*/3, + /*pad_top=*/0, /*pad_left=*/0, /*pad_bottom=*/0, /*pad_right=*/0, /*channels_last=*/true, + input.data(), weights.data(), bias.data(), nullptr, output.data(), + -std::numeric_limits::max(), std::numeric_limits::max(), nullptr)); + EXPECT_FALSE(ArmKleidiAI::DepthwiseConvKleidiAI( + batches, in_height, in_width, channels, /*filter_height=*/3, /*filter_width=*/2, + /*pad_top=*/0, /*pad_left=*/0, /*pad_bottom=*/0, /*pad_right=*/0, /*channels_last=*/true, + input.data(), weights.data(), bias.data(), nullptr, output.data(), + -std::numeric_limits::max(), std::numeric_limits::max(), nullptr)); +} + +#endif // defined(USE_KLEIDIAI) && !defined(_MSC_VER) diff --git a/onnxruntime/test/optimizer/nhwc_transformer_test.cc b/onnxruntime/test/optimizer/nhwc_transformer_test.cc index b9d8d478fdff2..ee79eae5b17d5 100644 --- a/onnxruntime/test/optimizer/nhwc_transformer_test.cc +++ b/onnxruntime/test/optimizer/nhwc_transformer_test.cc @@ -433,7 +433,7 @@ TEST(NhwcTransformerTests, ConvGlobalAveragePool) { TransformerLevel::Level3); } -TEST(NhwcTransformerTests, ConvDepthwiseFloat_SkipNhwcUntilDepthwiseKernelEnabled) { +TEST(NhwcTransformerTests, ConvDepthwiseFloat_UsesHelperCapability) { auto build_test_case = [&](ModelTestBuilder& builder) { auto* input_arg = builder.MakeInput({1, 8, 7, 7}, -1.0f, 1.0f); auto* weight_arg = builder.MakeInitializer({8, 1, 3, 3}, -1.0f, 1.0f); @@ -445,10 +445,10 @@ TEST(NhwcTransformerTests, ConvDepthwiseFloat_SkipNhwcUntilDepthwiseKernelEnable auto check_nhwc_graph = [&](InferenceSessionWrapper& session) { auto op_to_count = CountOpsInGraph(session.GetGraph()); - EXPECT_FALSE(HasFloatNhwcNoTransposeSupport({1, 8, 7, 7}, {8, 1, 3, 3}, {}, {}, {}, 8)); - EXPECT_EQ(op_to_count["Conv"] + op_to_count["com.microsoft.nchwc.Conv"], 1); - EXPECT_EQ(op_to_count["com.microsoft.NhwcFusedConv"], 0); - EXPECT_EQ(op_to_count["Transpose"], 0); + const bool expect_nhwc = HasFloatNhwcNoTransposeSupport({1, 8, 7, 7}, {8, 1, 3, 3}, {}, {}, {}, 8); + EXPECT_EQ(op_to_count["Conv"] + op_to_count["com.microsoft.nchwc.Conv"], expect_nhwc ? 0 : 1); + EXPECT_EQ(op_to_count["com.microsoft.NhwcFusedConv"], expect_nhwc ? 1 : 0); + EXPECT_EQ(op_to_count["Transpose"], expect_nhwc ? 2 : 0); }; TransformerTester(build_test_case, @@ -584,6 +584,44 @@ TEST(NhwcTransformerTests, ConvFloat_UsesNhwcOnlyWithKleidi) { /*relative_per_sample_tolerance*/ 1e-6); } +TEST(NhwcTransformerTests, ConvFloat_SymbolicChannelsUsesNhwc) { + if (!HasFloatNhwcNoTransposeSupport({1, 8, 7, 7}, {16, 8, 3, 3}, {1, 1, 1, 1})) { + GTEST_SKIP() << "Float NHWC KleidiAI path is not available on this configuration."; + } + + auto build_test_case = [&](ModelTestBuilder& builder) { + auto* input_arg = builder.MakeInput({1, 8, 7, 7}, -1.0f, 1.0f); + auto* weight_arg = builder.MakeInput({16, 8, 3, 3}, -1.0f, 1.0f); + auto* output_arg = builder.MakeOutput(); + + auto input_shape = *input_arg->Shape(); + input_shape.mutable_dim(1)->set_dim_param("channels"); + input_arg->SetShape(input_shape); + + auto weight_shape = *weight_arg->Shape(); + weight_shape.mutable_dim(1)->set_dim_param("channels"); + weight_arg->SetShape(weight_shape); + + Node& conv_node = builder.AddConvNode(input_arg, weight_arg, output_arg); + conv_node.AddAttribute("pads", std::vector{1, 1, 1, 1}); + }; + + auto check_nhwc_graph = [&](InferenceSessionWrapper& session) { + auto op_to_count = CountOpsInGraph(session.GetGraph()); + EXPECT_EQ(op_to_count["Conv"], 0); + EXPECT_EQ(op_to_count["com.microsoft.NhwcFusedConv"], 1); + EXPECT_EQ(op_to_count["Transpose"], 2); + }; + + TransformerTester(build_test_case, + check_nhwc_graph, + TransformerLevel::Level2, + TransformerLevel::Level3, + /*opset_version*/ 12, + /*per_sample_tolerance*/ 1e-6, + /*relative_per_sample_tolerance*/ 1e-6); +} + TEST(NhwcTransformerTests, ConvFloat_RespectsKleidiDisableConfig) { if (!HasFloatNhwcNoTransposeSupport({1, 8, 7, 7}, {16, 8, 3, 3}, {1, 1, 1, 1})) { GTEST_SKIP() << "Float NHWC KleidiAI path is not available on this configuration.";