From c8155c79cbd31bc4c65426b5d963bf6d0d34912c Mon Sep 17 00:00:00 2001 From: Jake Mathern Date: Wed, 19 Aug 2020 18:47:53 +0000 Subject: [PATCH] Merged PR 5051981: ArgMin/Argmax 12: add select_last_index attribute Add select_last_index atribute to argmin/argmax. normally the algorithm takes the index of the first occurrence of min/max, but with this attribute enabled it will take the index of the last occurrence. [windowsai pr](https://microsoft.visualstudio.com/WindowsAI/_git/WindowsAI/pullrequest/5051953) Related work items: #27469709 --- .../src/External/DirectMLHelpers/ApiTraits.h | 54 +++++++++++++++++-- .../External/DirectMLHelpers/DirectMLSchema.h | 34 +++++++++++- .../DirectMLHelpers/GeneratedSchemaHelpers.h | 30 +++++++++++ .../src/Operators/DmlOperatorReduce.cpp | 54 +++++++++++++++---- .../src/Operators/OperatorRegistration.cpp | 2 + .../dml/OperatorAuthorHelper/Attributes.h | 1 + .../dml/OperatorAuthorHelper/OperatorHelper.h | 2 + .../OperatorRegistration.h | 2 + 8 files changed, 164 insertions(+), 15 deletions(-) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h index 8563c39052..ac5b63820f 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h @@ -24,7 +24,7 @@ struct EnumTraits template <> struct EnumTraits { - static constexpr auto ValueCount = 121; + static constexpr auto ValueCount = 124; static constexpr size_t ActivationFunctionCount = 20; }; @@ -62,7 +62,7 @@ struct EnumTraits template <> struct EnumTraits { - static constexpr auto ValueCount = 3; + static constexpr auto ValueCount = 4; }; template <> @@ -86,7 +86,7 @@ struct EnumTraits template <> struct EnumTraits { - static constexpr auto ValueCount = 2; + static constexpr auto ValueCount = 4; }; template <> @@ -113,6 +113,12 @@ struct EnumTraits static constexpr auto ValueCount = 3; }; +template <> +struct EnumTraits +{ + static constexpr auto ValueCount = 1; +}; + template constexpr auto EnumValueCount = EnumTraits::ValueCount; @@ -405,6 +411,18 @@ struct OperatorDescTraits static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_REDUCE; }; +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ARGMIN; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ARGMAX; +}; + template <> struct OperatorDescTraits { @@ -759,6 +777,12 @@ struct OperatorDescTraits static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_MEAN_VARIANCE_NORMALIZATION1; }; +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_RESAMPLE1; +}; + template <> struct OperatorDescTraits { @@ -1143,6 +1167,18 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_REDUCE> using DescType = DML_REDUCE_OPERATOR_DESC; }; +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ARGMIN> +{ + using DescType = DML_ARGMIN_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ARGMAX> +{ + using DescType = DML_ARGMAX_OPERATOR_DESC; +}; + template <> struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_AVERAGE_POOLING> { @@ -1497,6 +1533,12 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_MEAN_VARIANCE_NORMALIZ using DescType = DML_MEAN_VARIANCE_NORMALIZATION1_OPERATOR_DESC; }; +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_RESAMPLE1> +{ + using DescType = DML_RESAMPLE1_OPERATOR_DESC; +}; + template <> struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_MATRIX_MULTIPLY_INTEGER> { @@ -1732,6 +1774,10 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args return std::invoke(std::forward(visitor), DML_GEMM_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_REDUCE: return std::invoke(std::forward(visitor), DML_REDUCE_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ARGMIN: + return std::invoke(std::forward(visitor), DML_ARGMIN_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ARGMAX: + return std::invoke(std::forward(visitor), DML_ARGMAX_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_AVERAGE_POOLING: return std::invoke(std::forward(visitor), DML_AVERAGE_POOLING_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_LP_POOLING: @@ -1951,6 +1997,8 @@ inline gsl::czstring ToString(DML_OPERATOR_TYPE value) case DML_OPERATOR_CONVOLUTION: return "DML_OPERATOR_CONVOLUTION"; case DML_OPERATOR_GEMM: return "DML_OPERATOR_GEMM"; case DML_OPERATOR_REDUCE: return "DML_OPERATOR_REDUCE"; + case DML_OPERATOR_ARGMIN: return "DML_OPERATOR_ARGMIN"; + case DML_OPERATOR_ARGMAX: return "DML_OPERATOR_ARGMAX"; case DML_OPERATOR_AVERAGE_POOLING: return "DML_OPERATOR_AVERAGE_POOLING"; case DML_OPERATOR_LP_POOLING: return "DML_OPERATOR_LP_POOLING"; case DML_OPERATOR_MAX_POOLING: return "DML_OPERATOR_MAX_POOLING"; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h index 746e39d601..35e28bf3b3 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h @@ -627,6 +627,38 @@ constexpr DML_OPERATOR_SCHEMA DML_REDUCE_OPERATOR_SCHEMA { DML_REDUCE_OPERATOR_SCHEMA_FIELDS, }; +constexpr DML_SCHEMA_FIELD DML_ARGMIN_OPERATOR_SCHEMA_FIELDS[5] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisCount", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT_ARRAY, "Axes", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisDirection", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ARGMIN_OPERATOR_SCHEMA { + "DML_OPERATOR_ARGMIN", + DML_OPERATOR_ARGMIN, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE, + 5, + DML_ARGMIN_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_ARGMAX_OPERATOR_SCHEMA_FIELDS[5] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisCount", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT_ARRAY, "Axes", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisDirection", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ARGMAX_OPERATOR_SCHEMA { + "DML_OPERATOR_ARGMAX", + DML_OPERATOR_ARGMAX, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE, + 5, + DML_ARGMAX_OPERATOR_SCHEMA_FIELDS, +}; + constexpr DML_SCHEMA_FIELD DML_AVERAGE_POOLING_OPERATOR_SCHEMA_FIELDS[8] { DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, @@ -1656,7 +1688,7 @@ constexpr DML_SCHEMA_FIELD DML_QUANTIZED_LINEAR_CONVOLUTION_OPERATOR_SCHEMA_FIEL DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "FilterTensor", false }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "FilterScaleTensor", false }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "FilterZeroPointTensor", true }, - DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "BiasTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "BiasTensor", true }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputScaleTensor", false }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputZeroPointTensor", true }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h index 78a42e97a3..15277f5e9e 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h @@ -344,6 +344,26 @@ inline std::vector GetFields(const DML_REDUCE_OPERATOR_DESC& desc OperatorField(&DML_REDUCE_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast(desc.Axes), desc.AxisCount)), }; } +inline std::vector GetFields(const DML_ARGMIN_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ARGMIN_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_ARGMIN_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_ARGMIN_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.AxisCount))), + OperatorField(&DML_ARGMIN_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.Axes), desc.AxisCount)), + OperatorField(&DML_ARGMIN_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast(desc.AxisDirection))), + }; +} +inline std::vector GetFields(const DML_ARGMAX_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ARGMAX_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_ARGMAX_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_ARGMAX_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.AxisCount))), + OperatorField(&DML_ARGMAX_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.Axes), desc.AxisCount)), + OperatorField(&DML_ARGMAX_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast(desc.AxisDirection))), + }; +} inline std::vector GetFields(const DML_AVERAGE_POOLING_OPERATOR_DESC& desc) { return { @@ -1211,6 +1231,8 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType) case DML_OPERATOR_CONVOLUTION: return DML_CONVOLUTION_OPERATOR_SCHEMA; case DML_OPERATOR_GEMM: return DML_GEMM_OPERATOR_SCHEMA; case DML_OPERATOR_REDUCE: return DML_REDUCE_OPERATOR_SCHEMA; + case DML_OPERATOR_ARGMIN: return DML_ARGMIN_OPERATOR_SCHEMA; + case DML_OPERATOR_ARGMAX: return DML_ARGMAX_OPERATOR_SCHEMA; case DML_OPERATOR_AVERAGE_POOLING: return DML_AVERAGE_POOLING_OPERATOR_SCHEMA; case DML_OPERATOR_LP_POOLING: return DML_LP_POOLING_OPERATOR_SCHEMA; case DML_OPERATOR_MAX_POOLING: return DML_MAX_POOLING_OPERATOR_SCHEMA; @@ -1460,6 +1482,14 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc) return AbstractOperatorDesc( &DML_REDUCE_OPERATOR_SCHEMA, GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ARGMIN: + return AbstractOperatorDesc( + &DML_ARGMIN_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ARGMAX: + return AbstractOperatorDesc( + &DML_ARGMAX_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); case DML_OPERATOR_AVERAGE_POOLING: return AbstractOperatorDesc( &DML_AVERAGE_POOLING_OPERATOR_SCHEMA, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorReduce.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorReduce.cpp index 128d4308fe..72758d743b 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorReduce.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorReduce.cpp @@ -54,9 +54,19 @@ public: reducedDims); } + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + // Zero the output tensor's memory for ArgMin & ArgMax, which produce INT64 output. - if ((function == DML_REDUCE_FUNCTION_ARGMAX) || (function == DML_REDUCE_FUNCTION_ARGMIN)) + if (function == DML_REDUCE_FUNCTION_ARGMAX) { + DML_ARGMAX_OPERATOR_DESC argmaxDesc; + argmaxDesc.AxisDirection = static_cast(m_selectLastIndex); + argmaxDesc.InputTensor = inputDescs.data(); + argmaxDesc.OutputTensor = outputDescs.data(); + argmaxDesc.Axes = dmlAxes.data(); + argmaxDesc.AxisCount = gsl::narrow_cast(dmlAxes.size()); + // If the 64-bit tensors were remapped to 32-bit, then we need to clear the upper 32-bits // of each element. If the device directly supports 64-bit elements, then no need. DmlOperator::Remap64bitDmlDataTypesTo32bitIfNeeded(); @@ -64,20 +74,42 @@ public: { m_zeroOperator = InitializeZeroInt64Tensor(m_outputTensorDescs[0].GetBufferSizeInBytes()); } + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ARGMAX, &argmaxDesc }; + SetDmlOperatorDesc(opDesc, kernelInfo); } + else if (function == DML_REDUCE_FUNCTION_ARGMIN) + { + DML_ARGMIN_OPERATOR_DESC argminDesc; + argminDesc.AxisDirection = static_cast(m_selectLastIndex); + argminDesc.InputTensor = inputDescs.data(); + argminDesc.OutputTensor = outputDescs.data(); + argminDesc.Axes = dmlAxes.data(); + argminDesc.AxisCount = gsl::narrow_cast(dmlAxes.size()); - std::vector inputDescs = GetDmlInputDescs(); - std::vector outputDescs = GetDmlOutputDescs(); + // If the 64-bit tensors were remapped to 32-bit, then we need to clear the upper 32-bits + // of each element. If the device directly supports 64-bit elements, then no need. + DmlOperator::Remap64bitDmlDataTypesTo32bitIfNeeded(); + if (m_outputTensorDescs[0].WasRemapped64bitTo32bit()) + { + m_zeroOperator = InitializeZeroInt64Tensor(m_outputTensorDescs[0].GetBufferSizeInBytes()); + } - DML_REDUCE_OPERATOR_DESC reduceDesc = {}; - reduceDesc.InputTensor = inputDescs.data(); - reduceDesc.OutputTensor = outputDescs.data(); - reduceDesc.Function = function; - reduceDesc.Axes = dmlAxes.data(); - reduceDesc.AxisCount = gsl::narrow_cast(dmlAxes.size()); + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ARGMIN, &argminDesc }; + SetDmlOperatorDesc(opDesc, kernelInfo); + } + else + { + DML_REDUCE_OPERATOR_DESC reduceDesc = {}; + reduceDesc.InputTensor = inputDescs.data(); + reduceDesc.OutputTensor = outputDescs.data(); + reduceDesc.Function = function; + reduceDesc.Axes = dmlAxes.data(); + reduceDesc.AxisCount = gsl::narrow_cast(dmlAxes.size()); - DML_OPERATOR_DESC opDesc = { DML_OPERATOR_REDUCE, &reduceDesc }; - SetDmlOperatorDesc(opDesc, kernelInfo); + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_REDUCE, &reduceDesc }; + SetDmlOperatorDesc(opDesc, kernelInfo); + } } void Compute(const MLOperatorKernelContext& kernelContext) override diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index 23fc1fc991..0eb72a4f29 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -475,8 +475,10 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 12, ReduceMin, typeNameListDefault, supportedTypeListFloat16to32Int8, DmlGraphSupport::Supported)}, {REG_INFO( 7, ArgMax, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, {REG_INFO( 11, ArgMax, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, + {REG_INFO( 12, ArgMax, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, {REG_INFO( 7, ArgMin, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, {REG_INFO( 11, ArgMin, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, + {REG_INFO( 12, ArgMin, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)}, {REG_INFO( 7, Gemm, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, Gemm, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 9, Gemm, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h index e6f0f995b8..541c0fe78b 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h @@ -45,6 +45,7 @@ namespace AttrName static constexpr const char* InputForget = "input_forget"; static constexpr const char* K = "k"; static constexpr const char* KeepDims = "keepdims"; + static constexpr const char* SelectLastIndex = "select_last_index"; static constexpr const char* KernelShape = "kernel_shape"; static constexpr const char* LinearBeforeReset = "linear_before_reset"; static constexpr const char* Lambda = "lambd"; // Deliberate typo to match ONNX spec. diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h index 0e78a5afeb..0bc87d52f9 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h @@ -733,6 +733,7 @@ class ReduceHelperBase { template ReduceHelperBase(const Info_t& info, const Shape_t& shape, bool usingAxes) { m_keepDims = info.GetOptionalAttribute(AttrName::KeepDims, 1); + m_selectLastIndex = info.GetOptionalAttribute(AttrName::SelectLastIndex, 0); if (usingAxes) { m_axes = info.GetOptionalAttributeVectorInt32(AttrName::Axes); } else { @@ -751,6 +752,7 @@ class ReduceHelperBase { protected: std::vector m_axes; int m_keepDims = 0; + int m_selectLastIndex = 0; }; class ArgMinArgMaxHelper : public ReduceHelperBase { diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h index cc69d0cbcd..b11da1ef42 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h @@ -256,6 +256,8 @@ namespace OperatorHelper static const int sc_sinceVer_MaxPool = 12; static const int sc_sinceVer_ReduceMax = 12; static const int sc_sinceVer_ReduceMin = 12; + static const int sc_sinceVer_ArgMin = 12; + static const int sc_sinceVer_ArgMax = 12; } // namespace OnnxOperatorSet12 namespace MsftOperatorSet1