Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions cmake/onnxruntime_mlas.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
47 changes: 32 additions & 15 deletions onnxruntime/core/mlas/lib/convolve.cpp
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
/*++

Copyright (c) Microsoft Corporation. All rights reserved.
SPDX-FileCopyrightText: Copyright 2026 Arm Limited and/or its affiliates <open-source-office@arm.com>

Licensed under the MIT License.

Expand All @@ -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
Expand Down Expand Up @@ -1351,6 +1355,7 @@ static constexpr size_t ComputeChannelsLastConvOutSize(size_t input, size_t kern

return 0;
}

#endif

} // namespace
Expand Down Expand Up @@ -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(&parameters);

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

done

#endif
}

Expand Down
130 changes: 85 additions & 45 deletions onnxruntime/core/mlas/lib/kleidiai/convolve_kleidiai.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand All @@ -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
Expand Down Expand Up @@ -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;
}

Expand Down Expand Up @@ -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<float>::max(),
std::numeric_limits<float>::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<const std::byte*>(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<const std::byte*>(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;
}
Loading