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
This commit is contained in:
Dwayne Robinson 2020-08-26 02:22:21 +00:00
parent 83b7c1151a
commit cb5e199a79
11 changed files with 119 additions and 12 deletions

View file

@ -24,7 +24,7 @@ struct EnumTraits<DML_TENSOR_TYPE>
template <>
struct EnumTraits<DML_OPERATOR_TYPE>
{
static constexpr auto ValueCount = 124;
static constexpr auto ValueCount = 141;
static constexpr size_t ActivationFunctionCount = 20;
};
@ -891,6 +891,12 @@ struct OperatorDescTraits<DML_ROI_ALIGN_OPERATOR_DESC>
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ROI_ALIGN;
};
template <>
struct OperatorDescTraits<DML_GATHER_ND1_OPERATOR_DESC>
{
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_GATHER_ND1;
};
template <>
struct OperatorDescTraits<DML_ACTIVATION_ELU_OPERATOR_DESC>
{
@ -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>(visitor), DML_ADAM_OPTIMIZER_OPERATOR_DESC{}, std::forward<Ts>(args)...);
case DML_OPERATOR_ROI_ALIGN:
return std::invoke(std::forward<Visitor>(visitor), DML_ROI_ALIGN_OPERATOR_DESC{}, std::forward<Ts>(args)...);
case DML_OPERATOR_GATHER_ND1:
return std::invoke(std::forward<Visitor>(visitor), DML_GATHER_ND1_OPERATOR_DESC{}, std::forward<Ts>(args)...);
case DML_OPERATOR_ACTIVATION_ELU:
return std::invoke(std::forward<Visitor>(visitor), DML_ACTIVATION_ELU_OPERATOR_DESC{}, std::forward<Ts>(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 "<unknown>";

View file

@ -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 },

View file

@ -1169,6 +1169,17 @@ inline std::vector<OperatorField> GetFields(const DML_ROI_ALIGN_OPERATOR_DESC& d
OperatorField(&DML_ROI_ALIGN_OPERATOR_SCHEMA.Fields[10], ToOperatorFieldType(static_cast<UINT>(desc.MaximumSamplesPerOutput))),
};
}
inline std::vector<OperatorField> GetFields(const DML_GATHER_ND1_OPERATOR_DESC& desc)
{
return {
OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.InputTensor))),
OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.IndicesTensor))),
OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast<UINT>(desc.InputDimensionCount))),
OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast<UINT>(desc.IndicesDimensionCount))),
OperatorField(&DML_GATHER_ND1_OPERATOR_SCHEMA.Fields[5], ToOperatorFieldType(static_cast<UINT>(desc.BatchDimensionCount))),
};
}
inline std::vector<OperatorField> 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<const DML_ROI_ALIGN_OPERATOR_DESC*>(opDesc.Desc)));
case DML_OPERATOR_GATHER_ND1:
return AbstractOperatorDesc(
&DML_GATHER_ND1_OPERATOR_SCHEMA,
GetFields(*static_cast<const DML_GATHER_ND1_OPERATOR_DESC*>(opDesc.Desc)));
case DML_OPERATOR_ACTIVATION_ELU:
return AbstractOperatorDesc(
&DML_ACTIVATION_ELU_OPERATOR_SCHEMA,

View file

@ -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<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
std::vector<DML_TENSOR_DESC> outputDescs = GetDmlOutputDescs();
assert(inputDescs.size() == 2);
@ -97,17 +109,18 @@ public:
auto outputTensorShapeDescription = kernelCreationContext.GetTensorShapeDescription();
std::vector<DimensionType> dataDimensions = outputTensorShapeDescription.GetInputTensorShape(0);
std::vector<DimensionType> 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<uint32_t>(dataDimensions.size());
operatorDesc.IndicesDimensionCount = static_cast<uint32_t>(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);
}
};

View file

@ -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)},

View file

@ -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<int32_t>(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;
}

View file

@ -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<const uint32_t> GetSizes() const { return { m_sizes, m_sizes + m_bufferTensorDesc.DimensionCount }; }
gsl::span<const uint32_t> GetStrides() const;

View file

@ -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";

View file

@ -737,21 +737,27 @@ namespace OperatorHelper
{
std::vector<DimensionType> inputDimensions = shapeInfo.GetInputTensorShape(0);
std::vector<DimensionType> 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<uint32_t>(inputDimensions.size()) - numberOfCoordinatesPerIndex;
const uint32_t numberOfOutputDimensionsFromIndices = static_cast<uint32_t>(indicesDimensions.size()) - 1; // Strip off last dimension.
uint32_t outputDimensionCount = gsl::narrow_cast<uint32_t>(numberOfOutputDimensionsFromIndices + numberOfOutputDimensionsFromInput);
ML_CHECK_VALID_ARGUMENT(inputDimensions.size() >= batchCount + numberOfCoordinatesPerIndex);
const uint32_t numberOfOutputDimensionsFromInput = static_cast<uint32_t>(inputDimensions.size()) - batchCount - numberOfCoordinatesPerIndex;
const uint32_t numberOfOutputDimensionsFromIndices = static_cast<uint32_t>(indicesDimensions.size()) - batchCount - 1; // Strip off last dimension.
uint32_t outputDimensionCount = gsl::narrow_cast<uint32_t>(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<DimensionType> 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)) };

View file

@ -977,9 +977,13 @@ class GatherNdHelper {
// Shape_t is used to obtain input shape which will be used for adjusting attribute value.
template <typename Info_t, typename Shape_t>
GatherNdHelper(const Info_t& info, const Shape_t& shape) {
m_batchCount = info.GetOptionalAttribute<int32_t>(AttrName::BatchDimensions, 0);
}
std::vector<EdgeShapes> GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const;
protected:
int32_t m_batchCount;
};
class PoolingHelperBase {

View file

@ -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;