From cb5e199a790239b200cbb601f1493a602d6e0ba8 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Wed, 26 Aug 2020 02:22:21 +0000 Subject: [PATCH] Merged PR 5093868: GatherND1 ORT DML EP Add batchDimensionCount. https://github.com/onnx/onnx/pull/2585 - add batch_dim parameter. DML PR: https://microsoft.visualstudio.com/WindowsAI/_git/WindowsAI/pullrequest/5089850 --- .../src/External/DirectMLHelpers/ApiTraits.h | 17 +++++++++- .../External/DirectMLHelpers/DirectMLSchema.h | 17 ++++++++++ .../DirectMLHelpers/GeneratedSchemaHelpers.h | 16 ++++++++++ .../src/Operators/DmlOperatorGather.cpp | 21 +++++++++--- .../src/Operators/OperatorRegistration.cpp | 1 + .../DmlExecutionProvider/src/TensorDesc.cpp | 32 +++++++++++++++++++ .../dml/DmlExecutionProvider/src/TensorDesc.h | 1 + .../dml/OperatorAuthorHelper/Attributes.h | 1 + .../OperatorAuthorHelper/OperatorHelper.cpp | 20 ++++++++---- .../dml/OperatorAuthorHelper/OperatorHelper.h | 4 +++ .../OperatorRegistration.h | 1 + 11 files changed, 119 insertions(+), 12 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 31cadafbae..2b177faeda 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 = 124; + static constexpr auto ValueCount = 141; static constexpr size_t ActivationFunctionCount = 20; }; @@ -891,6 +891,12 @@ struct OperatorDescTraits static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ROI_ALIGN; }; +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_GATHER_ND1; +}; + template <> struct OperatorDescTraits { @@ -1731,6 +1737,12 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ROI_ALIGN> using DescType = DML_ROI_ALIGN_OPERATOR_DESC; }; +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_GATHER_ND1> +{ + using DescType = DML_GATHER_ND1_OPERATOR_DESC; +}; + template <> struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_ELU> { @@ -2102,6 +2114,8 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args return std::invoke(std::forward(visitor), DML_ADAM_OPTIMIZER_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ROI_ALIGN: return std::invoke(std::forward(visitor), DML_ROI_ALIGN_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_GATHER_ND1: + return std::invoke(std::forward(visitor), DML_GATHER_ND1_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_ELU: return std::invoke(std::forward(visitor), DML_ACTIVATION_ELU_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_CELU: @@ -2273,6 +2287,7 @@ inline gsl::czstring ToString(DML_OPERATOR_TYPE value) case DML_OPERATOR_SLICE_GRAD: return "DML_OPERATOR_SLICE_GRAD"; case DML_OPERATOR_ADAM_OPTIMIZER: return "DML_OPERATOR_ADAM_OPTIMIZER"; case DML_OPERATOR_ROI_ALIGN: return "DML_OPERATOR_ROI_ALIGN"; + case DML_OPERATOR_GATHER_ND1: return "DML_OPERATOR_GATHER_ND1"; default: assert(false); return ""; 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 31d0d07406..ec3d24070b 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h @@ -1932,6 +1932,23 @@ constexpr DML_OPERATOR_SCHEMA DML_ROI_ALIGN_OPERATOR_SCHEMA { DML_ROI_ALIGN_OPERATOR_SCHEMA_FIELDS, }; +constexpr DML_SCHEMA_FIELD DML_GATHER_ND1_OPERATOR_SCHEMA_FIELDS[6] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "IndicesTensor", 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, "InputDimensionCount", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "IndicesDimensionCount", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "BatchDimensionCount", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_GATHER_ND1_OPERATOR_SCHEMA { + "DML_OPERATOR_GATHER_ND1", + DML_OPERATOR_GATHER_ND1, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE, + 6, + DML_GATHER_ND1_OPERATOR_SCHEMA_FIELDS, +}; + constexpr DML_SCHEMA_FIELD DML_ACTIVATION_ELU_OPERATOR_SCHEMA_FIELDS[3] { 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 }, 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 66a62f5042..7630e98489 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h @@ -1169,6 +1169,17 @@ inline std::vector GetFields(const DML_ROI_ALIGN_OPERATOR_DESC& d OperatorField(&DML_ROI_ALIGN_OPERATOR_SCHEMA.Fields[10], ToOperatorFieldType(static_cast(desc.MaximumSamplesPerOutput))), }; } +inline std::vector GetFields(const DML_GATHER_ND1_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.IndicesTensor))), + OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.InputDimensionCount))), + OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast(desc.IndicesDimensionCount))), + OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[5], ToOperatorFieldType(static_cast(desc.BatchDimensionCount))), + }; +} inline std::vector GetFields(const DML_ACTIVATION_ELU_OPERATOR_DESC& desc) { return { @@ -1451,6 +1462,7 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType) case DML_OPERATOR_SLICE_GRAD: return DML_SLICE_GRAD_OPERATOR_SCHEMA; case DML_OPERATOR_ADAM_OPTIMIZER: return DML_ADAM_OPTIMIZER_OPERATOR_SCHEMA; case DML_OPERATOR_ROI_ALIGN: return DML_ROI_ALIGN_OPERATOR_SCHEMA; + case DML_OPERATOR_GATHER_ND1: return DML_GATHER_ND1_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_ELU: return DML_ACTIVATION_ELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_CELU: return DML_ACTIVATION_CELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_HARDMAX: return DML_ACTIVATION_HARDMAX_OPERATOR_SCHEMA; @@ -1956,6 +1968,10 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc) return AbstractOperatorDesc( &DML_ROI_ALIGN_OPERATOR_SCHEMA, GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_GATHER_ND1: + return AbstractOperatorDesc( + &DML_GATHER_ND1_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); case DML_OPERATOR_ACTIVATION_ELU: return AbstractOperatorDesc( &DML_ACTIVATION_ELU_OPERATOR_SCHEMA, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGather.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGather.cpp index ec70544657..dea317686f 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGather.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGather.cpp @@ -89,6 +89,18 @@ public: DmlOperator::Initialize(kernelCreationContext); DmlOperator::Remap64bitDmlDataTypesTo32bitIfNeeded(); + uint32_t maxDimensionCount = std::max({ + m_inputTensorDescs[0].GetDimensionCount(), + m_inputTensorDescs[1].GetDimensionCount(), + m_outputTensorDescs[0].GetDimensionCount() + }); + + // DML expects all tensors to have the same dimension count. + // Update the tensor descriptions with new sizes. + m_inputTensorDescs[0].SetDimensionCount(maxDimensionCount, TensorAxis::RightAligned); + m_inputTensorDescs[1].SetDimensionCount(maxDimensionCount, TensorAxis::RightAligned); + m_outputTensorDescs[0].SetDimensionCount(maxDimensionCount, TensorAxis::RightAligned); + std::vector inputDescs = GetDmlInputDescs(); std::vector outputDescs = GetDmlOutputDescs(); assert(inputDescs.size() == 2); @@ -97,17 +109,18 @@ public: auto outputTensorShapeDescription = kernelCreationContext.GetTensorShapeDescription(); std::vector dataDimensions = outputTensorShapeDescription.GetInputTensorShape(0); std::vector indicesDimensions = outputTensorShapeDescription.GetInputTensorShape(1); - ML_CHECK_VALID_ARGUMENT(dataDimensions.size() <= OperatorHelper::NchwDimensionCount); - ML_CHECK_VALID_ARGUMENT(indicesDimensions.size() <= OperatorHelper::NchwDimensionCount); + ML_CHECK_VALID_ARGUMENT(dataDimensions.size() > m_batchCount); + ML_CHECK_VALID_ARGUMENT(indicesDimensions.size() > m_batchCount); - DML_GATHER_ND_OPERATOR_DESC operatorDesc = {}; + DML_GATHER_ND1_OPERATOR_DESC operatorDesc = {}; operatorDesc.InputTensor = &inputDescs[0]; operatorDesc.IndicesTensor = &inputDescs[1]; operatorDesc.OutputTensor = outputDescs.data(); operatorDesc.InputDimensionCount = static_cast(dataDimensions.size()); operatorDesc.IndicesDimensionCount = static_cast(indicesDimensions.size()); + operatorDesc.BatchDimensionCount = m_batchCount; - DML_OPERATOR_DESC opDesc = { DML_OPERATOR_GATHER_ND, &operatorDesc }; + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_GATHER_ND1, &operatorDesc }; SetDmlOperatorDesc(opDesc, kernelCreationContext); } }; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index 7eb53ba54c..442ecf4e89 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -393,6 +393,7 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 11, Gather, typeNameListScatterGather, supportedTypeListScatterGather, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, {REG_INFO( 11, GatherElements, typeNameListScatterGather, supportedTypeListScatterGather, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, {REG_INFO( 11, GatherND, typeNameListScatterGatherND, supportedTypeListScatterGatherND, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, + {REG_INFO( 12, GatherND, typeNameListScatterGatherND, supportedTypeListScatterGatherND, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, {REG_INFO_VER( 9, Scatter, typeNameListScatterGather, supportedTypeListScatterGather, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, {REG_INFO_VER( 11, Scatter, typeNameListScatterGather, supportedTypeListScatterGather, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, {REG_INFO( 11, ScatterElements, typeNameListScatterGather, supportedTypeListScatterGather, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStridesFromAnyEp)}, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp index a470378e06..e4fb01dfa9 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp @@ -325,3 +325,35 @@ void TensorDesc::ForceUnsignedDataType() ML_INVALID_ARGUMENT("Can't coerce unknown or non-integral data type"); } } + +void TensorDesc::SetDimensionCount(uint32_t newDimensionCount, TensorAxis alignment) +{ + ML_CHECK_VALID_ARGUMENT(newDimensionCount <= MaximumDimensionCount); + ML_CHECK_VALID_ARGUMENT(alignment == TensorAxis::RightAligned || alignment == TensorAxis::LeftAligned); + + const uint32_t oldDimensionCount = m_bufferTensorDesc.DimensionCount; + const int32_t difference = static_cast(newDimensionCount - oldDimensionCount); + if (difference == 0) + { + return; + } + + int32_t fillOffset = oldDimensionCount; + int32_t fillCount = std::max(0, difference); + + // alignment == TensorAxis::LeftAligned is the easy case. + // Right alignment needs more work, shifting values over. + if (alignment == TensorAxis::RightAligned) + { + fillOffset = 0; // Fill leading dimensions with 1's starting at the front. + uint32_t moveCount = std::min(newDimensionCount, oldDimensionCount); + memmove(&m_sizes[fillCount], &m_sizes[oldDimensionCount - moveCount], sizeof(m_sizes[0]) * moveCount); + memmove(&m_strides[fillCount], &m_strides[oldDimensionCount - moveCount], sizeof(m_strides[0]) * moveCount); + } + if (fillCount > 0) + { + std::fill(&m_sizes[fillOffset], &m_sizes[fillOffset] + fillCount, 1u); + std::fill(&m_strides[fillOffset], &m_strides[fillOffset] + fillCount, 0u); + } + m_bufferTensorDesc.DimensionCount = newDimensionCount; +} diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.h index 39f4b0e56c..e9b7d97c48 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.h @@ -42,6 +42,7 @@ namespace Dml inline bool IsValid() const { return m_tensorType != DML_TENSOR_TYPE_INVALID; } inline uint32_t GetDimensionCount() const { return m_bufferTensorDesc.DimensionCount; } + void SetDimensionCount(uint32_t newDimensionCount, TensorAxis alignment); gsl::span GetSizes() const { return { m_sizes, m_sizes + m_bufferTensorDesc.DimensionCount }; } gsl::span GetStrides() const; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h index 5b286f2b23..deb194ff7e 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h @@ -15,6 +15,7 @@ namespace AttrName static constexpr const char* Axis = "axis"; static constexpr const char* AxisW = "axis_w"; static constexpr const char* BatchAxis = "batch_axis"; + static constexpr const char* BatchDimensions = "batch_dims"; static constexpr const char* Beta = "beta"; static constexpr const char* Bias = "bias"; static constexpr const char* BlockSize = "blocksize"; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp index 23d604be7e..c635407ee2 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp @@ -737,21 +737,27 @@ namespace OperatorHelper { std::vector inputDimensions = shapeInfo.GetInputTensorShape(0); std::vector indicesDimensions = shapeInfo.GetInputTensorShape(1); + int32_t batchCount = m_batchCount; // Determine the number of output dimensions. ML_CHECK_VALID_ARGUMENT(inputDimensions.size() >= 1); ML_CHECK_VALID_ARGUMENT(indicesDimensions.size() >= 1); + ML_CHECK_VALID_ARGUMENT(inputDimensions.size() > batchCount); + ML_CHECK_VALID_ARGUMENT(indicesDimensions.size() > batchCount); const uint32_t numberOfCoordinatesPerIndex = indicesDimensions.back(); - ML_CHECK_VALID_ARGUMENT(inputDimensions.size() >= numberOfCoordinatesPerIndex); - const uint32_t numberOfOutputDimensionsFromInput = static_cast(inputDimensions.size()) - numberOfCoordinatesPerIndex; - const uint32_t numberOfOutputDimensionsFromIndices = static_cast(indicesDimensions.size()) - 1; // Strip off last dimension. - uint32_t outputDimensionCount = gsl::narrow_cast(numberOfOutputDimensionsFromIndices + numberOfOutputDimensionsFromInput); + ML_CHECK_VALID_ARGUMENT(inputDimensions.size() >= batchCount + numberOfCoordinatesPerIndex); + const uint32_t numberOfOutputDimensionsFromInput = static_cast(inputDimensions.size()) - batchCount - numberOfCoordinatesPerIndex; + const uint32_t numberOfOutputDimensionsFromIndices = static_cast(indicesDimensions.size()) - batchCount - 1; // Strip off last dimension. + uint32_t outputDimensionCount = gsl::narrow_cast(batchCount + numberOfOutputDimensionsFromIndices + numberOfOutputDimensionsFromInput); ML_CHECK_VALID_ARGUMENT(outputDimensionCount > 0); - // Form the full expected size by concatenating the prefix part of the indices tensor shape - // with the suffix of the input tensor shape. + // Form the full expected size by concatenating fragments: + // 1 - batch count + // 2 - prefix part of the indices tensor shape + // 3 - suffix of the input tensor shape. std::vector outputDimensions; - outputDimensions.assign(indicesDimensions.begin(), indicesDimensions.end() - 1); + outputDimensions.assign(inputDimensions.begin(), inputDimensions.begin() + batchCount); + outputDimensions.insert(outputDimensions.end(), indicesDimensions.begin() + batchCount, indicesDimensions.end() - 1); outputDimensions.insert(outputDimensions.end(), inputDimensions.end() - numberOfOutputDimensionsFromInput, inputDimensions.end()); return { EdgeShapes(std::move(outputDimensions)) }; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h index 983db6b237..a32da47dbc 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h @@ -977,9 +977,13 @@ class GatherNdHelper { // Shape_t is used to obtain input shape which will be used for adjusting attribute value. template GatherNdHelper(const Info_t& info, const Shape_t& shape) { + m_batchCount = info.GetOptionalAttribute(AttrName::BatchDimensions, 0); } std::vector GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const; + +protected: + int32_t m_batchCount; }; class PoolingHelperBase { diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h index 25b9573a6e..c8784a79d8 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h @@ -251,6 +251,7 @@ namespace OperatorHelper static const int sc_sinceVer_LessOrEqual = 12; static const int sc_sinceVer_Celu = 12; static const int sc_sinceVer_Clip = 12; + static const int sc_sinceVer_GatherND = 12; static const int sc_sinceVer_Min = 12; static const int sc_sinceVer_Max = 12; static const int sc_sinceVer_Pow = 12;