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 2870946480..666dc9c0f1 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h @@ -13,7 +13,7 @@ struct EnumTraits template <> struct EnumTraits { - static constexpr auto ValueCount = 9; + static constexpr auto ValueCount = 12; }; template <> @@ -25,7 +25,7 @@ struct EnumTraits template <> struct EnumTraits { - static constexpr auto ValueCount = 97; + static constexpr auto ValueCount = 107; static constexpr size_t ActivationFunctionCount = 19; }; @@ -90,6 +90,24 @@ struct EnumTraits static constexpr auto ValueCount = 2; }; +template <> +struct EnumTraits +{ + static constexpr auto ValueCount = 3; +}; + +template <> +struct EnumTraits +{ + static constexpr auto ValueCount = 2; +}; + +template <> +struct EnumTraits +{ + static constexpr auto ValueCount = 3; +}; + template constexpr auto EnumValueCount = EnumTraits::ValueCount; @@ -610,6 +628,66 @@ struct OperatorDescTraits static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_RESAMPLE; }; +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_ROUND; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_IS_INFINITY; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_FILL_VALUE_CONSTANT; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_FILL_VALUE_SEQUENCE; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_CUMULATIVE_SUMMATION; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_REVERSE_SUBSEQUENCES; +}; + template <> struct OperatorDescTraits { @@ -1192,6 +1270,66 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_RESAMPLE> using DescType = DML_RESAMPLE_OPERATOR_DESC; }; +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT> +{ + using DescType = DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT> +{ + using DescType = DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ELEMENT_WISE_ROUND> +{ + using DescType = DML_ELEMENT_WISE_ROUND_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ELEMENT_WISE_IS_INFINITY> +{ + using DescType = DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE> +{ + using DescType = DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR> +{ + using DescType = DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_FILL_VALUE_CONSTANT> +{ + using DescType = DML_FILL_VALUE_CONSTANT_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_FILL_VALUE_SEQUENCE> +{ + using DescType = DML_FILL_VALUE_SEQUENCE_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_CUMULATIVE_SUMMATION> +{ + using DescType = DML_CUMULATIVE_SUMMATION_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_REVERSE_SUBSEQUENCES> +{ + using DescType = DML_REVERSE_SUBSEQUENCES_OPERATOR_DESC; +}; + template <> struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_ELU> { @@ -1306,7 +1444,6 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_SHRINK> using DescType = DML_ACTIVATION_SHRINK_OPERATOR_DESC; }; - // Calls a visitor functor, supplying an empty operator desc corresponding to the given DML_OPERATOR_TYPE as // the first argument. // @@ -1474,6 +1611,26 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args return std::invoke(std::forward(visitor), DML_ONE_HOT_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_RESAMPLE: return std::invoke(std::forward(visitor), DML_RESAMPLE_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT: + return std::invoke(std::forward(visitor), DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT: + return std::invoke(std::forward(visitor), DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ELEMENT_WISE_ROUND: + return std::invoke(std::forward(visitor), DML_ELEMENT_WISE_ROUND_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ELEMENT_WISE_IS_INFINITY: + return std::invoke(std::forward(visitor), DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE: + return std::invoke(std::forward(visitor), DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR: + return std::invoke(std::forward(visitor), DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_FILL_VALUE_CONSTANT: + return std::invoke(std::forward(visitor), DML_FILL_VALUE_CONSTANT_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_FILL_VALUE_SEQUENCE: + return std::invoke(std::forward(visitor), DML_FILL_VALUE_SEQUENCE_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_CUMULATIVE_SUMMATION: + return std::invoke(std::forward(visitor), DML_CUMULATIVE_SUMMATION_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_REVERSE_SUBSEQUENCES: + return std::invoke(std::forward(visitor), DML_REVERSE_SUBSEQUENCES_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_HARDMAX: @@ -1601,6 +1758,16 @@ inline gsl::czstring ToString(DML_OPERATOR_TYPE value) case DML_OPERATOR_SCATTER: return "DML_OPERATOR_SCATTER"; case DML_OPERATOR_ONE_HOT: return "DML_OPERATOR_ONE_HOT"; case DML_OPERATOR_RESAMPLE: return "DML_OPERATOR_RESAMPLE"; + case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT: return "DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT"; + case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT: return "DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT"; + case DML_OPERATOR_ELEMENT_WISE_ROUND: return "DML_OPERATOR_ELEMENT_WISE_ROUND"; + case DML_OPERATOR_ELEMENT_WISE_IS_INFINITY: return "DML_OPERATOR_ELEMENT_WISE_IS_INFINITY"; + case DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE: return "DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE"; + case DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR: return "DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR"; + case DML_OPERATOR_FILL_VALUE_CONSTANT: return "DML_OPERATOR_FILL_VALUE_CONSTANT"; + case DML_OPERATOR_FILL_VALUE_SEQUENCE: return "DML_OPERATOR_FILL_VALUE_SEQUENCE"; + case DML_OPERATOR_CUMULATIVE_SUMMATION: return "DML_OPERATOR_CUMULATIVE_SUMMATION"; + case DML_OPERATOR_REVERSE_SUBSEQUENCES: return "DML_OPERATOR_REVERSE_SUBSEQUENCES"; 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 7c46a8a6a2..b95fbfc95c 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h @@ -19,12 +19,15 @@ enum DML_SCHEMA_FIELD_TYPE DML_SCHEMA_FIELD_TYPE_OPERATOR_DESC, DML_SCHEMA_FIELD_TYPE_OPERATOR_DESC_ARRAY, DML_SCHEMA_FIELD_TYPE_UINT, + DML_SCHEMA_FIELD_TYPE_UINT64, DML_SCHEMA_FIELD_TYPE_INT, DML_SCHEMA_FIELD_TYPE_FLOAT, DML_SCHEMA_FIELD_TYPE_UINT_ARRAY, + DML_SCHEMA_FIELD_TYPE_INT_ARRAY, DML_SCHEMA_FIELD_TYPE_FLOAT_ARRAY, DML_SCHEMA_FIELD_TYPE_SCALE_BIAS, DML_SCHEMA_FIELD_TYPE_SIZE_2D, + DML_SCHEMA_FIELD_TYPE_SCALAR_UNION, }; enum DML_SCHEMA_OPERATOR_SUPPORT_FLAGS @@ -1246,6 +1249,150 @@ constexpr DML_OPERATOR_SCHEMA DML_RESAMPLE_OPERATOR_SCHEMA { DML_RESAMPLE_OPERATOR_SCHEMA_FIELDS, }; +constexpr DML_SCHEMA_FIELD DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA_FIELDS[3] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "ATensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "BTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA { + "DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT", + DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 3, + DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA_FIELDS[3] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "ATensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "BTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA { + "DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT", + DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 3, + DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_ELEMENT_WISE_ROUND_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 }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "RoundingMode", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA { + "DML_OPERATOR_ELEMENT_WISE_ROUND", + DML_OPERATOR_ELEMENT_WISE_ROUND, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 3, + DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_ELEMENT_WISE_IS_INFINITY_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 }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "InfinityMode", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA { + "DML_OPERATOR_ELEMENT_WISE_IS_INFINITY", + DML_OPERATOR_ELEMENT_WISE_IS_INFINITY, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 3, + DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA_FIELDS[3] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "ATensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "BTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA { + "DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE", + DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 3, + DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA_FIELDS[3] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "ATensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "BTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA { + "DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR", + DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 3, + DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA_FIELDS[3] { + 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, "ValueDataType", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_SCALAR_UNION, "Value", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA { + "DML_OPERATOR_FILL_VALUE_CONSTANT", + DML_OPERATOR_FILL_VALUE_CONSTANT, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 3, + DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA_FIELDS[4] { + 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, "ValueDataType", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_SCALAR_UNION, "ValueStart", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_SCALAR_UNION, "ValueDelta", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA { + "DML_OPERATOR_FILL_VALUE_SEQUENCE", + DML_OPERATOR_FILL_VALUE_SEQUENCE, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 4, + DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_CUMULATIVE_SUMMATION_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, "Axis", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "HasExclusiveSum", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisDirection", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA { + "DML_OPERATOR_CUMULATIVE_SUMMATION", + DML_OPERATOR_CUMULATIVE_SUMMATION, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 5, + DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA_FIELDS[4] { + 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, "SequenceLengthsTensor", 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, "Axis", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA { + "DML_OPERATOR_REVERSE_SUBSEQUENCES", + DML_OPERATOR_REVERSE_SUBSEQUENCES, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 4, + DML_REVERSE_SUBSEQUENCES_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 }, @@ -1511,4 +1658,10 @@ constexpr DML_OPERATOR_SCHEMA DML_ACTIVATION_SHRINK_OPERATOR_SCHEMA { DML_ACTIVATION_SHRINK_OPERATOR_SCHEMA_FIELDS, }; +constexpr DML_SCHEMA_FIELD DML_RNN_ZERO_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_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "SequenceLengthsTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, +}; + } // extern "C" 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 b8285c77a7..d565590ece 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h @@ -738,6 +738,90 @@ inline std::vector GetFields(const DML_RESAMPLE_OPERATOR_DESC& de OperatorField(&DML_RESAMPLE_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast(desc.Scales), desc.ScaleCount)), }; } +inline std::vector GetFields(const DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.ATensor))), + OperatorField(&DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.BTensor))), + OperatorField(&DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.OutputTensor))), + }; +} +inline std::vector GetFields(const DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.ATensor))), + OperatorField(&DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.BTensor))), + OperatorField(&DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.OutputTensor))), + }; +} +inline std::vector GetFields(const DML_ELEMENT_WISE_ROUND_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.RoundingMode))), + }; +} +inline std::vector GetFields(const DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.InfinityMode))), + }; +} +inline std::vector GetFields(const DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.ATensor))), + OperatorField(&DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.BTensor))), + OperatorField(&DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.OutputTensor))), + }; +} +inline std::vector GetFields(const DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.ATensor))), + OperatorField(&DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.BTensor))), + OperatorField(&DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.OutputTensor))), + }; +} +inline std::vector GetFields(const DML_FILL_VALUE_CONSTANT_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.ValueDataType))), + OperatorField(&DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.Value))), + }; +} +inline std::vector GetFields(const DML_FILL_VALUE_SEQUENCE_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.ValueDataType))), + OperatorField(&DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.ValueStart))), + OperatorField(&DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.ValueDelta))), + }; +} +inline std::vector GetFields(const DML_CUMULATIVE_SUMMATION_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.Axis))), + OperatorField(&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.HasExclusiveSum))), + OperatorField(&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast(desc.AxisDirection))), + }; +} +inline std::vector GetFields(const DML_REVERSE_SUBSEQUENCES_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.SequenceLengthsTensor))), + OperatorField(&DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.Axis))), + }; +} inline std::vector GetFields(const DML_ACTIVATION_ELU_OPERATOR_DESC& desc) { return { @@ -970,6 +1054,16 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType) case DML_OPERATOR_SCATTER: return DML_SCATTER_OPERATOR_SCHEMA; case DML_OPERATOR_ONE_HOT: return DML_ONE_HOT_OPERATOR_SCHEMA; case DML_OPERATOR_RESAMPLE: return DML_RESAMPLE_OPERATOR_SCHEMA; + case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT: return DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA; + case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT: return DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA; + case DML_OPERATOR_ELEMENT_WISE_ROUND: return DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA; + case DML_OPERATOR_ELEMENT_WISE_IS_INFINITY: return DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA; + case DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE: return DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA; + case DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR: return DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA; + case DML_OPERATOR_FILL_VALUE_CONSTANT: return DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA; + case DML_OPERATOR_FILL_VALUE_SEQUENCE: return DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA; + case DML_OPERATOR_CUMULATIVE_SUMMATION: return DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA; + case DML_OPERATOR_REVERSE_SUBSEQUENCES: return DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_ELU: return DML_ACTIVATION_ELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_HARDMAX: return DML_ACTIVATION_HARDMAX_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_HARD_SIGMOID: return DML_ACTIVATION_HARD_SIGMOID_OPERATOR_SCHEMA; @@ -989,6 +1083,7 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType) case DML_OPERATOR_ACTIVATION_TANH: return DML_ACTIVATION_TANH_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_THRESHOLDED_RELU: return DML_ACTIVATION_THRESHOLDED_RELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_SHRINK: return DML_ACTIVATION_SHRINK_OPERATOR_SCHEMA; + default: THROW_HR(E_INVALIDARG); } } @@ -1305,6 +1400,46 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc) return AbstractOperatorDesc( &DML_RESAMPLE_OPERATOR_SCHEMA, GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT: + return AbstractOperatorDesc( + &DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT: + return AbstractOperatorDesc( + &DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ELEMENT_WISE_ROUND: + return AbstractOperatorDesc( + &DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ELEMENT_WISE_IS_INFINITY: + return AbstractOperatorDesc( + &DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE: + return AbstractOperatorDesc( + &DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR: + return AbstractOperatorDesc( + &DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_FILL_VALUE_CONSTANT: + return AbstractOperatorDesc( + &DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_FILL_VALUE_SEQUENCE: + return AbstractOperatorDesc( + &DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_CUMULATIVE_SUMMATION: + return AbstractOperatorDesc( + &DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_REVERSE_SUBSEQUENCES: + return AbstractOperatorDesc( + &DML_REVERSE_SUBSEQUENCES_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/External/DirectMLHelpers/GeneratedSchemaTypes.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaTypes.h index 57c8ec8ce0..21c8c46af7 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaTypes.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaTypes.h @@ -7,13 +7,16 @@ using ApiAttributeVariant = std::variant< const DML_TENSOR_DESC*, const DML_OPERATOR_DESC*, UINT, + UINT64, INT, FLOAT, const UINT*, + const INT*, const FLOAT*, const DML_SCALE_BIAS*, - DML_SIZE_2D - >; + DML_SIZE_2D, + DML_SCALAR_UNION +>; namespace OperatorFieldTypes { @@ -22,12 +25,15 @@ namespace OperatorFieldTypes using OperatorDesc = std::optional; // DML_SCHEMA_FIELD_TYPE_OPERATOR_DESC using OperatorDescArray = std::optional>; // DML_SCHEMA_FIELD_TYPE_OPERATOR_DESC_ARRAY using UInt = uint32_t; // DML_SCHEMA_FIELD_TYPE_UINT + using UInt64 = uint64_t; // DML_SCHEMA_FIELD_TYPE_UINT64 using Int = int32_t; // DML_SCHEMA_FIELD_TYPE_INT using Float = float; // DML_SCHEMA_FIELD_TYPE_FLOAT using UIntArray = std::optional>; // DML_SCHEMA_FIELD_TYPE_UINT_ARRAY + using IntArray = std::optional>; // DML_SCHEMA_FIELD_TYPE_INT_ARRAY using FloatArray = std::optional>; // DML_SCHEMA_FIELD_TYPE_FLOAT_ARRAY using ScaleBias = std::optional; // DML_SCHEMA_FIELD_TYPE_SCALE_BIAS using Size2D = DML_SIZE_2D; // DML_SCHEMA_FIELD_TYPE_SIZE_2D + using ScalarUnion = DML_SCALAR_UNION; // DML_SCHEMA_FIELD_TYPE_SCALAR_UNION } using OperatorFieldVariant = std::variant< @@ -36,13 +42,16 @@ using OperatorFieldVariant = std::variant< OperatorFieldTypes::OperatorDesc, OperatorFieldTypes::OperatorDescArray, OperatorFieldTypes::UInt, + OperatorFieldTypes::UInt64, OperatorFieldTypes::Int, OperatorFieldTypes::Float, OperatorFieldTypes::UIntArray, + OperatorFieldTypes::IntArray, OperatorFieldTypes::FloatArray, OperatorFieldTypes::ScaleBias, - OperatorFieldTypes::Size2D - >; + OperatorFieldTypes::Size2D, + OperatorFieldTypes::ScalarUnion +>; class OperatorField { @@ -80,6 +89,9 @@ public: const OperatorFieldTypes::UInt& AsUInt() const { return std::get(m_data); } OperatorFieldTypes::UInt& AsUInt() { return std::get(m_data); } + const OperatorFieldTypes::UInt64& AsUInt64() const { return std::get(m_data); } + OperatorFieldTypes::UInt64& AsUInt64() { return std::get(m_data); } + const OperatorFieldTypes::Int& AsInt() const { return std::get(m_data); } OperatorFieldTypes::Int& AsInt() { return std::get(m_data); } @@ -89,6 +101,9 @@ public: const OperatorFieldTypes::UIntArray& AsUIntArray() const { return std::get(m_data); } OperatorFieldTypes::UIntArray& AsUIntArray() { return std::get(m_data); } + const OperatorFieldTypes::IntArray& AsIntArray() const { return std::get(m_data); } + OperatorFieldTypes::IntArray& AsIntArray() { return std::get(m_data); } + const OperatorFieldTypes::FloatArray& AsFloatArray() const { return std::get(m_data); } OperatorFieldTypes::FloatArray& AsFloatArray() { return std::get(m_data); } @@ -98,6 +113,9 @@ public: const OperatorFieldTypes::Size2D& AsSize2D() const { return std::get(m_data); } OperatorFieldTypes::Size2D& AsSize2D() { return std::get(m_data); } + const OperatorFieldTypes::ScalarUnion& AsScalarUnion() const { return std::get(m_data); } + OperatorFieldTypes::ScalarUnion& AsScalarUnion() { return std::get(m_data); } + private: const DML_SCHEMA_FIELD* m_schema; OperatorFieldVariant m_data; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/SchemaHelpers.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/SchemaHelpers.h index fba6503dd6..09f1b1cdc4 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/SchemaHelpers.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/SchemaHelpers.h @@ -50,6 +50,11 @@ namespace SchemaHelpers return value; } + inline OperatorFieldTypes::UInt64 ToOperatorFieldType(uint64_t value) + { + return value; + } + inline OperatorFieldTypes::Int ToOperatorFieldType(int32_t value) { return value; @@ -71,6 +76,17 @@ namespace SchemaHelpers return field; } + inline OperatorFieldTypes::IntArray ToOperatorFieldType(const int32_t* values, uint32_t count) + { + OperatorFieldTypes::IntArray field; + if (values && count != 0) + { + field.emplace(count); + std::copy_n(values, count, field->begin()); + } + return field; + } + inline OperatorFieldTypes::FloatArray ToOperatorFieldType(const float* values, uint32_t count) { OperatorFieldTypes::FloatArray field; @@ -92,6 +108,10 @@ namespace SchemaHelpers return value; } + inline OperatorFieldTypes::ScalarUnion ToOperatorFieldType(DML_SCALAR_UNION value) + { + return value; + } class StructFieldWriter { @@ -250,6 +270,12 @@ namespace SchemaHelpers dst->Write(value); } break; + case DML_SCHEMA_FIELD_TYPE_UINT64: + { + uint64_t value = field.AsUInt64(); + dst->Write(value); + } break; + case DML_SCHEMA_FIELD_TYPE_INT: { int32_t value = field.AsInt(); @@ -276,6 +302,20 @@ namespace SchemaHelpers dst->Write(arrayPtr); } break; + case DML_SCHEMA_FIELD_TYPE_INT_ARRAY: + { + int32_t* arrayPtr = nullptr; + + const auto& values = field.AsIntArray(); + if (values) + { + arrayPtr = allocator->Allocate(values->size()); + std::copy(values->begin(), values->end(), arrayPtr); + } + + dst->Write(arrayPtr); + } break; + case DML_SCHEMA_FIELD_TYPE_FLOAT_ARRAY: { float* arrayPtr = nullptr; @@ -310,6 +350,12 @@ namespace SchemaHelpers dst->Write(value); } break; + case DML_SCHEMA_FIELD_TYPE_SCALAR_UNION: + { + uint64_t value = field.AsScalarUnion().UInt64; + dst->Write(value); + } break; + default: assert(false); THROW_HR(E_UNEXPECTED); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorConvInteger.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorConvInteger.cpp new file mode 100644 index 0000000000..45f4204abd --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorConvInteger.cpp @@ -0,0 +1,81 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +//TODO::: + +#include "precomp.h" + +namespace Dml +{ + +class DmlOperatorConvInteger : public DmlOperator, OneHotHelper// TODO::: +{ +public: + using Self = DmlOperatorConvInteger; + + DmlOperatorConvInteger(const MLOperatorKernelCreationContext& kernelCreationContext) + : DmlOperator(kernelCreationContext), + OneHotHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) + { + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() == 3); + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() == 1); + std::vector> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'. + std::vector> outputIndices = { 0 }; + DmlOperator::Initialize(kernelCreationContext, inputIndices, outputIndices); + + // Unsqueeze the indices tensor by inserting a flat dimension of size 1, + // and compute the output tensor by expanding along the active axis. + // This way they are both size-compatible and directly consumable by DirectML. + std::vector indicesDimensions; + indicesDimensions = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0); + indicesDimensions.insert(indicesDimensions.begin() + m_absoluteAxis, 1u); + + // Update the tensor descriptions with new sizes. + m_inputTensorDescs[0] = + TensorDesc( + m_inputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(indicesDimensions), + gsl::make_span(indicesDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + m_outputTensorDescs[0] = + TensorDesc( + m_outputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(m_outputDimensions), + gsl::make_span(m_outputDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + // Adjust the axis so it's in DML's terms rather than the original ONNX indexing. + uint32_t dmlAxis = GetDmlAdjustedAxis( + m_absoluteAxis, + gsl::narrow_cast(indicesDimensions.size()), + m_inputTensorDescs.front().GetDimensionCount() + ); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + DML_ONE_HOT_OPERATOR_DESC operatorDesc = {}; + operatorDesc.IndicesTensor = &inputDescs[0]; + operatorDesc.ValuesTensor = &inputDescs[1]; + operatorDesc.OutputTensor = outputDescs.data(); + operatorDesc.Axis = dmlAxis; + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ONE_HOT, &operatorDesc }; + SetDmlOperatorDesc(opDesc, kernelCreationContext); + } +}; + +DML_OP_DEFINE_CREATION_FUNCTION(ConvInteger, DmlOperatorConvInteger); + +} // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCumSum.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCumSum.cpp new file mode 100644 index 0000000000..4a8e438eb1 --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCumSum.cpp @@ -0,0 +1,45 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "precomp.h" + +namespace Dml +{ + +class DmlOperatorCumSum : public DmlOperator, OneHotHelper +{ +public: + using Self = DmlOperatorCumSum; + + DmlOperatorCumSum(const MLOperatorKernelCreationContext& kernelCreationContext) + : DmlOperator(kernelCreationContext), + OneHotHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) + { + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() == 1); + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() == 1); + DmlOperator::Initialize(kernelCreationContext); + + // Adjust the axis so it's in DML's terms rather than the original ONNX indexing. + int32_t hasExclusiveSum = kernelCreationContext.GetOptionalAttribute(AttrName::Exclusive, 0); + int32_t isReversed = kernelCreationContext.GetOptionalAttribute(AttrName::Reverse, 0); + int32_t onnxAxis = kernelCreationContext.GetOptionalAttribute(AttrName::Axis, -1); + uint32_t dmlAxis = GetDmlAdjustedAxis(onnxAxis, kernelCreationContext, m_inputTensorDescs.front().GetDimensionCount()); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + DML_CUMULATIVE_SUMMATION_OPERATOR_DESC operatorDesc = {}; + operatorDesc.InputTensor = inputDescs.data(); + operatorDesc.OutputTensor = outputDescs.data(); + operatorDesc.HasExclusiveSum = hasExclusiveSum; + operatorDesc.Axis = dmlAxis; + operatorDesc.AxisDirection = isReversed ? DML_AXIS_DIRECTION_DECREASING : DML_AXIS_DIRECTION_INCREASING; + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_CUMULATIVE_SUMMATION, &operatorDesc }; + SetDmlOperatorDesc(opDesc, kernelCreationContext); + } +}; + +DML_OP_DEFINE_CREATION_FUNCTION(CumSum, DmlOperatorCumSum); + +} // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorDynamicQuantizeLinear.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorDynamicQuantizeLinear.cpp new file mode 100644 index 0000000000..0f4487de52 --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorDynamicQuantizeLinear.cpp @@ -0,0 +1,79 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "precomp.h" +//TODO::: +namespace Dml +{ + +class DmlOperatorDynamicQuantizeLinear : public DmlOperator, OneHotHelper +{ +public: + using Self = DmlOperatorDynamicQuantizeLinear; + + DmlOperatorDynamicQuantizeLinear(const MLOperatorKernelCreationContext& kernelCreationContext) + : DmlOperator(kernelCreationContext), + OneHotHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) + { + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() == 3); + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() == 1); + std::vector> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'. + std::vector> outputIndices = { 0 }; + DmlOperator::Initialize(kernelCreationContext, inputIndices, outputIndices); + + // Unsqueeze the indices tensor by inserting a flat dimension of size 1, + // and compute the output tensor by expanding along the active axis. + // This way they are both size-compatible and directly consumable by DirectML. + std::vector indicesDimensions; + indicesDimensions = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0); + indicesDimensions.insert(indicesDimensions.begin() + m_absoluteAxis, 1u); + + // Update the tensor descriptions with new sizes. + m_inputTensorDescs[0] = + TensorDesc( + m_inputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(indicesDimensions), + gsl::make_span(indicesDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + m_outputTensorDescs[0] = + TensorDesc( + m_outputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(m_outputDimensions), + gsl::make_span(m_outputDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + // Adjust the axis so it's in DML's terms rather than the original ONNX indexing. + uint32_t dmlAxis = GetDmlAdjustedAxis( + m_absoluteAxis, + gsl::narrow_cast(indicesDimensions.size()), + m_inputTensorDescs.front().GetDimensionCount() + ); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + DML_ONE_HOT_OPERATOR_DESC operatorDesc = {}; + operatorDesc.IndicesTensor = &inputDescs[0]; + operatorDesc.ValuesTensor = &inputDescs[1]; + operatorDesc.OutputTensor = outputDescs.data(); + operatorDesc.Axis = dmlAxis; + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ONE_HOT, &operatorDesc }; + SetDmlOperatorDesc(opDesc, kernelCreationContext); + } +}; + +DML_OP_DEFINE_CREATION_FUNCTION(DynamicQuantizeLinear, DmlOperatorDynamicQuantizeLinear); + +} // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorElementWise.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorElementWise.cpp index 3e7341c625..9226816503 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorElementWise.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorElementWise.cpp @@ -70,7 +70,7 @@ public: } else { - // Dml doesn't support UINT datatypes redirect to Identity because abs doesn't do anything to UINT + // DML doesn't support UINT datatypes. So redirect to Identity because Abs doesn't do anything to UINT. DML_ELEMENT_WISE_IDENTITY_OPERATOR_DESC opDesc = {}; opDesc.InputTensor = inputDescs.data(); opDesc.OutputTensor = outputDescs.data(); @@ -534,6 +534,110 @@ public: } }; +class DmlOperatorElementwiseMod : public DmlOperator +{ +public: + DmlOperatorElementwiseMod(const MLOperatorKernelCreationContext& kernelInfo) : DmlOperator(kernelInfo) + { + ML_CHECK_VALID_ARGUMENT(kernelInfo.GetInputCount() == 2); + ML_CHECK_VALID_ARGUMENT(kernelInfo.GetOutputCount() == 1); + + Initialize(kernelInfo, std::nullopt, std::nullopt, kernelInfo.GetTensorShapeDescription().GetOutputTensorShape(0)); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + auto fmod = kernelInfo.GetOptionalAttribute(AttrName::Fmod, 0); + + // Note TRUNCATE and FLOOR modulus operator descriptions are identical. + static_assert(sizeof(DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_DESC) == sizeof(DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_DESC)); + DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_DESC opDesc = {}; + opDesc.ATensor = &inputDescs[0]; + opDesc.BTensor = &inputDescs[1]; + opDesc.OutputTensor = &outputDescs[0]; + + DML_OPERATOR_TYPE type = fmod ? DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE : DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR; + SetDmlOperatorDesc({ type, &opDesc}, kernelInfo); + } +}; + +class DmlOperatorElementwiseBitShift : public DmlOperator +{ +public: + DmlOperatorElementwiseBitShift(const MLOperatorKernelCreationContext& kernelInfo) : DmlOperator(kernelInfo) + { + ML_CHECK_VALID_ARGUMENT(kernelInfo.GetInputCount() == 2); + ML_CHECK_VALID_ARGUMENT(kernelInfo.GetOutputCount() == 1); + + Initialize(kernelInfo, std::nullopt, std::nullopt, kernelInfo.GetTensorShapeDescription().GetOutputTensorShape(0)); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + // Note LEFT and RIGHT shift operator descriptions are identical. + static_assert(sizeof(DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_DESC) == sizeof(DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_DESC)); + DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_DESC opDesc = {}; + opDesc.ATensor = &inputDescs[0]; + opDesc.BTensor = &inputDescs[1]; + opDesc.OutputTensor = &outputDescs[0]; + + std::string mode = kernelInfo.GetOptionalAttribute(AttrName::Direction, ""); + ML_CHECK_VALID_ARGUMENT(mode == "LEFT" || mode == "RIGHT"); + + DML_OPERATOR_TYPE type = (mode == "LEFT") ? DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT : DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT; + SetDmlOperatorDesc({ type, &opDesc}, kernelInfo); + } +}; + +class DmlOperatorElementwiseIsInf : public DmlOperator +{ +public: + DmlOperatorElementwiseIsInf(const MLOperatorKernelCreationContext& kernelInfo) : DmlOperator(kernelInfo) + { + ML_CHECK_VALID_ARGUMENT(kernelInfo.GetInputCount() == 1); + ML_CHECK_VALID_ARGUMENT(kernelInfo.GetOutputCount() == 1); + + Initialize(kernelInfo, std::nullopt, std::nullopt, kernelInfo.GetTensorShapeDescription().GetOutputTensorShape(0)); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + auto detectPositive = kernelInfo.GetOptionalAttribute(AttrName::DetectPositive, 1); + auto detectNegative = kernelInfo.GetOptionalAttribute(AttrName::DetectNegative, 1); + + DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_DESC opDesc = {}; + opDesc.InputTensor = inputDescs.data(); + opDesc.OutputTensor = outputDescs.data(); + opDesc.InfinityMode = (detectPositive == detectNegative) ? DML_IS_INFINITY_MODE_EITHER + : detectPositive ? DML_IS_INFINITY_MODE_POSITIVE + : DML_IS_INFINITY_MODE_NEGATIVE; + + SetDmlOperatorDesc({ DML_OPERATOR_ELEMENT_WISE_CLIP, &opDesc}, kernelInfo); + } +}; + +class DmlOperatorElementwiseRound : public DmlOperator +{ +public: + DmlOperatorElementwiseRound(const MLOperatorKernelCreationContext& kernelInfo) : DmlOperator(kernelInfo) + { + ML_CHECK_VALID_ARGUMENT(kernelInfo.GetInputCount() == 1); + ML_CHECK_VALID_ARGUMENT(kernelInfo.GetOutputCount() == 1); + + Initialize(kernelInfo, std::nullopt, std::nullopt, kernelInfo.GetTensorShapeDescription().GetOutputTensorShape(0)); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + DML_ELEMENT_WISE_ROUND_OPERATOR_DESC opDesc = {}; + opDesc.InputTensor = inputDescs.data(); + opDesc.OutputTensor = outputDescs.data(); + opDesc.RoundingMode = DML_ROUNDING_MODE_HALVES_TO_NEAREST_EVEN; + + SetDmlOperatorDesc({ DML_OPERATOR_ELEMENT_WISE_ROUND, &opDesc}, kernelInfo); + } +}; + // Unary operators: DML_OP_DEFINE_CREATION_FUNCTION(Sqrt, DmlOperatorElementwiseUnary); DML_OP_DEFINE_CREATION_FUNCTION(Reciprocal, DmlOperatorElementwiseUnary); @@ -582,6 +686,10 @@ DML_OP_DEFINE_CREATION_FUNCTION(Pow, DmlOperatorElementwisePow); DML_OP_DEFINE_CREATION_FUNCTION(QuantizeLinear, DmlOperatorElementwiseQLinear); DML_OP_DEFINE_CREATION_FUNCTION(DequantizeLinear, DmlOperatorElementwiseQLinear); DML_OP_DEFINE_CREATION_FUNCTION(Where, DmlOperatorElementwiseIf); +DML_OP_DEFINE_CREATION_FUNCTION(Mod, DmlOperatorElementwiseMod); +DML_OP_DEFINE_CREATION_FUNCTION(BitShift, DmlOperatorElementwiseBitShift); +DML_OP_DEFINE_CREATION_FUNCTION(IsInf, DmlOperatorElementwiseIsInf); +DML_OP_DEFINE_CREATION_FUNCTION(Round, DmlOperatorElementwiseRound); // Fused operators: DML_OP_DEFINE_CREATION_FUNCTION(FusedAdd, DmlOperatorElementwiseBinary); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGather.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGather.cpp index b06e1c3afd..01dd7379d3 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGather.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGather.cpp @@ -44,5 +44,8 @@ public: }; DML_OP_DEFINE_CREATION_FUNCTION(Gather, DmlOperatorGather); +// TODO::: +DML_OP_DEFINE_CREATION_FUNCTION(GatherElements, DmlOperatorGather); +DML_OP_DEFINE_CREATION_FUNCTION(GatherND, DmlOperatorGather); } // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGemm.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGemm.cpp index e5c1ff555f..5fcc8a89fe 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGemm.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorGemm.cpp @@ -36,7 +36,7 @@ public: DML_GEMM_OPERATOR_DESC gemmDesc = {}; gemmDesc.ATensor = &inputDescs[0]; gemmDesc.BTensor = &inputDescs[1]; - gemmDesc.CTensor = &inputDescs[2]; + gemmDesc.CTensor = kernelInfo.IsInputValid(2) ? &inputDescs[2] : nullptr; gemmDesc.OutputTensor = &outputDescs[0]; gemmDesc.TransA = (m_transA ? DML_MATRIX_TRANSFORM_TRANSPOSE : DML_MATRIX_TRANSFORM_NONE); gemmDesc.TransB = (m_transB ? DML_MATRIX_TRANSFORM_TRANSPOSE : DML_MATRIX_TRANSFORM_NONE); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorMatMulInteger.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorMatMulInteger.cpp new file mode 100644 index 0000000000..bf9e62d3bb --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorMatMulInteger.cpp @@ -0,0 +1,81 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + + +//TODO::: +#include "precomp.h" + +namespace Dml +{ + +class DmlOperatorMatMulInteger : public DmlOperator, OneHotHelper +{ +public: + using Self = DmlOperatorMatMulInteger; + + DmlOperatorMatMulInteger(const MLOperatorKernelCreationContext& kernelCreationContext) + : DmlOperator(kernelCreationContext), + OneHotHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) + { + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() == 3); + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() == 1); + std::vector> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'. + std::vector> outputIndices = { 0 }; + DmlOperator::Initialize(kernelCreationContext, inputIndices, outputIndices); + + // Unsqueeze the indices tensor by inserting a flat dimension of size 1, + // and compute the output tensor by expanding along the active axis. + // This way they are both size-compatible and directly consumable by DirectML. + std::vector indicesDimensions; + indicesDimensions = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0); + indicesDimensions.insert(indicesDimensions.begin() + m_absoluteAxis, 1u); + + // Update the tensor descriptions with new sizes. + m_inputTensorDescs[0] = + TensorDesc( + m_inputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(indicesDimensions), + gsl::make_span(indicesDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + m_outputTensorDescs[0] = + TensorDesc( + m_outputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(m_outputDimensions), + gsl::make_span(m_outputDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + // Adjust the axis so it's in DML's terms rather than the original ONNX indexing. + uint32_t dmlAxis = GetDmlAdjustedAxis( + m_absoluteAxis, + gsl::narrow_cast(indicesDimensions.size()), + m_inputTensorDescs.front().GetDimensionCount() + ); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + DML_ONE_HOT_OPERATOR_DESC operatorDesc = {}; + operatorDesc.IndicesTensor = &inputDescs[0]; + operatorDesc.ValuesTensor = &inputDescs[1]; + operatorDesc.OutputTensor = outputDescs.data(); + operatorDesc.Axis = dmlAxis; + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ONE_HOT, &operatorDesc }; + SetDmlOperatorDesc(opDesc, kernelCreationContext); + } +}; + +DML_OP_DEFINE_CREATION_FUNCTION(MatMulInteger, DmlOperatorMatMulInteger); + +} // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorQLinearConv.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorQLinearConv.cpp new file mode 100644 index 0000000000..211d34f051 --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorQLinearConv.cpp @@ -0,0 +1,81 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +// TODO::: + +#include "precomp.h" + +namespace Dml +{ + +class DmlOperatorQLinearConv : public DmlOperator, OneHotHelper +{ +public: + using Self = DmlOperatorQLinearConv; + + DmlOperatorQLinearConv(const MLOperatorKernelCreationContext& kernelCreationContext) + : DmlOperator(kernelCreationContext), + OneHotHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) + { + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() == 3); + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() == 1); + std::vector> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'. + std::vector> outputIndices = { 0 }; + DmlOperator::Initialize(kernelCreationContext, inputIndices, outputIndices); + + // Unsqueeze the indices tensor by inserting a flat dimension of size 1, + // and compute the output tensor by expanding along the active axis. + // This way they are both size-compatible and directly consumable by DirectML. + std::vector indicesDimensions; + indicesDimensions = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0); + indicesDimensions.insert(indicesDimensions.begin() + m_absoluteAxis, 1u); + + // Update the tensor descriptions with new sizes. + m_inputTensorDescs[0] = + TensorDesc( + m_inputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(indicesDimensions), + gsl::make_span(indicesDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + m_outputTensorDescs[0] = + TensorDesc( + m_outputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(m_outputDimensions), + gsl::make_span(m_outputDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + // Adjust the axis so it's in DML's terms rather than the original ONNX indexing. + uint32_t dmlAxis = GetDmlAdjustedAxis( + m_absoluteAxis, + gsl::narrow_cast(indicesDimensions.size()), + m_inputTensorDescs.front().GetDimensionCount() + ); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + DML_ONE_HOT_OPERATOR_DESC operatorDesc = {}; + operatorDesc.IndicesTensor = &inputDescs[0]; + operatorDesc.ValuesTensor = &inputDescs[1]; + operatorDesc.OutputTensor = outputDescs.data(); + operatorDesc.Axis = dmlAxis; + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ONE_HOT, &operatorDesc }; + SetDmlOperatorDesc(opDesc, kernelCreationContext); + } +}; + +DML_OP_DEFINE_CREATION_FUNCTION(QLinearConv, DmlOperatorQLinearConv); + +} // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorQLinearMatMul.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorQLinearMatMul.cpp new file mode 100644 index 0000000000..2d46a4cb92 --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorQLinearMatMul.cpp @@ -0,0 +1,81 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +//TODO::: + +#include "precomp.h" + +namespace Dml +{ + +class DmlOperatorQLinearMatMul : public DmlOperator, OneHotHelper +{ +public: + using Self = DmlOperatorQLinearMatMul; + + DmlOperatorQLinearMatMul(const MLOperatorKernelCreationContext& kernelCreationContext) + : DmlOperator(kernelCreationContext), + OneHotHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) + { + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() == 3); + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() == 1); + std::vector> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'. + std::vector> outputIndices = { 0 }; + DmlOperator::Initialize(kernelCreationContext, inputIndices, outputIndices); + + // Unsqueeze the indices tensor by inserting a flat dimension of size 1, + // and compute the output tensor by expanding along the active axis. + // This way they are both size-compatible and directly consumable by DirectML. + std::vector indicesDimensions; + indicesDimensions = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0); + indicesDimensions.insert(indicesDimensions.begin() + m_absoluteAxis, 1u); + + // Update the tensor descriptions with new sizes. + m_inputTensorDescs[0] = + TensorDesc( + m_inputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(indicesDimensions), + gsl::make_span(indicesDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + m_outputTensorDescs[0] = + TensorDesc( + m_outputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(m_outputDimensions), + gsl::make_span(m_outputDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + // Adjust the axis so it's in DML's terms rather than the original ONNX indexing. + uint32_t dmlAxis = GetDmlAdjustedAxis( + m_absoluteAxis, + gsl::narrow_cast(indicesDimensions.size()), + m_inputTensorDescs.front().GetDimensionCount() + ); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + DML_ONE_HOT_OPERATOR_DESC operatorDesc = {}; + operatorDesc.IndicesTensor = &inputDescs[0]; + operatorDesc.ValuesTensor = &inputDescs[1]; + operatorDesc.OutputTensor = outputDescs.data(); + operatorDesc.Axis = dmlAxis; + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ONE_HOT, &operatorDesc }; + SetDmlOperatorDesc(opDesc, kernelCreationContext); + } +}; + +DML_OP_DEFINE_CREATION_FUNCTION(QLinearMatMul, DmlOperatorQLinearMatMul); + +} // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRange.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRange.cpp new file mode 100644 index 0000000000..e9048056c3 --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRange.cpp @@ -0,0 +1,82 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + + +// TODO::: + +#include "precomp.h" + +namespace Dml +{ + +class DmlOperatorRange : public DmlOperator, OneHotHelper +{ +public: + using Self = DmlOperatorRange; + + DmlOperatorRange(const MLOperatorKernelCreationContext& kernelCreationContext) + : DmlOperator(kernelCreationContext), + OneHotHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) + { + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() == 3); + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() == 1); + std::vector> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'. + std::vector> outputIndices = { 0 }; + DmlOperator::Initialize(kernelCreationContext, inputIndices, outputIndices); + + // Unsqueeze the indices tensor by inserting a flat dimension of size 1, + // and compute the output tensor by expanding along the active axis. + // This way they are both size-compatible and directly consumable by DirectML. + std::vector indicesDimensions; + indicesDimensions = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0); + indicesDimensions.insert(indicesDimensions.begin() + m_absoluteAxis, 1u); + + // Update the tensor descriptions with new sizes. + m_inputTensorDescs[0] = + TensorDesc( + m_inputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(indicesDimensions), + gsl::make_span(indicesDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + m_outputTensorDescs[0] = + TensorDesc( + m_outputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(m_outputDimensions), + gsl::make_span(m_outputDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + // Adjust the axis so it's in DML's terms rather than the original ONNX indexing. + uint32_t dmlAxis = GetDmlAdjustedAxis( + m_absoluteAxis, + gsl::narrow_cast(indicesDimensions.size()), + m_inputTensorDescs.front().GetDimensionCount() + ); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + DML_ONE_HOT_OPERATOR_DESC operatorDesc = {}; + operatorDesc.IndicesTensor = &inputDescs[0]; + operatorDesc.ValuesTensor = &inputDescs[1]; + operatorDesc.OutputTensor = outputDescs.data(); + operatorDesc.Axis = dmlAxis; + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ONE_HOT, &operatorDesc }; + SetDmlOperatorDesc(opDesc, kernelCreationContext); + } +}; + +DML_OP_DEFINE_CREATION_FUNCTION(Range, DmlOperatorRange); + +} // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorReverseSequence.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorReverseSequence.cpp new file mode 100644 index 0000000000..842ad8f4cc --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorReverseSequence.cpp @@ -0,0 +1,67 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + + +// TODO::: + +#include "precomp.h" + +namespace Dml +{ + +class DmlOperatorReverseSequence : public DmlOperator, OneHotHelper +{ +public: + using Self = DmlOperatorReverseSequence; + + DmlOperatorReverseSequence(const MLOperatorKernelCreationContext& kernelCreationContext) + : DmlOperator(kernelCreationContext), + OneHotHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) + { + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() == 2); + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() == 1); + DmlOperator::Initialize(kernelCreationContext); + + std::vector inputDimensions = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0); + std::vector sequenceLengthDimensions = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(1); + + // Read axis. + int32_t onnxAxis = kernelCreationContext.GetOptionalAttribute(AttrName::TimeAxis, 0); + onnxAxis = HandleNegativeAxis(onnxAxis, static_cast(inputDimensions.size())); + const uint32_t dmlAxis = GetDmlAdjustedAxis(onnxAxis, onnxAxis, m_inputTensorDescs.front().GetDimensionCount()); + + // Fix up the sequence lengths tensor (originally 1D) to be rank compatible with input, + // with all dimensions being the same as input except the active reversal axis. + std::vector adjustedSequenceLengthDimensions = inputDimensions; + adjustedSequenceLengthDimensions[onnxAxis] = 1; + ML_CHECK_VALID_ARGUMENT(ComputeElementCountFromDimensions(adjustedSequenceLengthDimensions), ComputeElementCountFromDimensions(sequenceLengthDimensions)); + + m_inputTensorDescs[1] = + TensorDesc( + m_inputTensorDescs[0].GetMlOperatorDataType(), + gsl::make_span(adjustedSequenceLengthDimensions), + gsl::make_span(adjustedSequenceLengthDimensions), + TensorAxis::DoNotCoerce, + TensorAxis::W, + TensorAxis::RightAligned, + NchwDimensionCount, // minDimensionCount + 0 + ); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + DML_REVERSE_SUBSEQUENCES_OPERATOR_DESC operatorDesc = {}; + operatorDesc.InputTensor = &inputDescs[0]; + operatorDesc.SequenceLengthsTensor = &inputDescs[1]; + operatorDesc.OutputTensor = outputDescs.data(); + operatorDesc.Axis = dmlAxis; + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_REVERSE_SUBSEQUENCES, &operatorDesc }; + SetDmlOperatorDesc(opDesc, kernelCreationContext); + } +}; + +DML_OP_DEFINE_CREATION_FUNCTION(ReverseSequence, DmlOperatorReverseSequence); + +} // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorScatter.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorScatter.cpp index ee742c4ccb..97b881d542 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorScatter.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorScatter.cpp @@ -38,13 +38,10 @@ public: assert(inputDescs.size() == 1); assert(outputDescs.size() == 1); - DML_SCALE_BIAS scaleBias = {}; - scaleBias.Scale = 1.0f; - DML_ELEMENT_WISE_IDENTITY_OPERATOR_DESC operatorDesc = {}; operatorDesc.InputTensor = &inputDescs[0]; operatorDesc.OutputTensor = outputDescs.data(); - operatorDesc.ScaleBias = &scaleBias; + operatorDesc.ScaleBias = nullptr; DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ELEMENT_WISE_IDENTITY, &operatorDesc }; SetDmlOperatorDesc(opDesc, kernelCreationContext); @@ -78,5 +75,8 @@ public: }; DML_OP_DEFINE_CREATION_FUNCTION(Scatter, DmlOperatorScatter); +// TODO::: +DML_OP_DEFINE_CREATION_FUNCTION(ScatterElements, DmlOperatorScatter); +DML_OP_DEFINE_CREATION_FUNCTION(ScatterND, DmlOperatorScatter); } // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorSlice.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorSlice.cpp index 0e9d0feb5a..bbeb9d9a87 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorSlice.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorSlice.cpp @@ -71,11 +71,12 @@ public: } }; -void CALLBACK QuerySlice(IMLOperatorSupportQueryContextPrivate* context, bool *isSupported) +void CALLBACK QuerySlice(IMLOperatorSupportQueryContextPrivate* context, bool* isSupported) { *isSupported = (context->GetInputCount() <= 4); } DML_OP_DEFINE_CREATION_FUNCTION(Slice7, DmlOperatorSliceTemplate<7>); DML_OP_DEFINE_CREATION_FUNCTION(Slice10, DmlOperatorSliceTemplate<10>); +DML_OP_DEFINE_CREATION_FUNCTION(Slice11, DmlOperatorSliceTemplate<11>); } // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index f19bf8d6a9..077cc38c3d 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -101,6 +101,7 @@ DML_OP_EXTERN_CREATION_FUNCTION(Tile); DML_OP_EXTERN_CREATION_FUNCTION(Concat); DML_OP_EXTERN_CREATION_FUNCTION(Slice7); DML_OP_EXTERN_CREATION_FUNCTION(Slice10); +DML_OP_EXTERN_CREATION_FUNCTION(Slice11); DML_OP_EXTERN_CREATION_FUNCTION(Pad); DML_OP_EXTERN_CREATION_FUNCTION(SpaceToDepth); DML_OP_EXTERN_CREATION_FUNCTION(DepthToSpace); @@ -203,6 +204,22 @@ DML_OP_EXTERN_CREATION_FUNCTION(MaxUnpool); DML_OP_EXTERN_CREATION_FUNCTION(Scatter); DML_OP_EXTERN_CREATION_FUNCTION(Resize); DML_OP_EXTERN_CREATION_FUNCTION(ConstantOfShape); +DML_OP_EXTERN_CREATION_FUNCTION(IsInf); +DML_OP_EXTERN_CREATION_FUNCTION(Mod); +DML_OP_EXTERN_CREATION_FUNCTION(BitShift); +DML_OP_EXTERN_CREATION_FUNCTION(CumSum); +DML_OP_EXTERN_CREATION_FUNCTION(GatherElements); +DML_OP_EXTERN_CREATION_FUNCTION(GatherND); +DML_OP_EXTERN_CREATION_FUNCTION(Range); +DML_OP_EXTERN_CREATION_FUNCTION(ReverseSequence); +DML_OP_EXTERN_CREATION_FUNCTION(Round); +DML_OP_EXTERN_CREATION_FUNCTION(ScatterElements); +DML_OP_EXTERN_CREATION_FUNCTION(ScatterND); +DML_OP_EXTERN_CREATION_FUNCTION(QLinearConv); +DML_OP_EXTERN_CREATION_FUNCTION(QLinearMatMul); +DML_OP_EXTERN_CREATION_FUNCTION(DynamicQuantizeLinear); +DML_OP_EXTERN_CREATION_FUNCTION(MatMulInteger); +DML_OP_EXTERN_CREATION_FUNCTION(ConvInteger); DML_OP_EXTERN_QUERY_FUNCTION(MaxPool); DML_OP_EXTERN_QUERY_FUNCTION(Slice); @@ -210,16 +227,17 @@ DML_OP_EXTERN_QUERY_FUNCTION(Slice); const static char* const typeNameListDefault[1] = {"T"}; const static char* const typeNameListTopK[2] = { "T", "I" }; const static char* const typeNameListLogicalComparison[2] = { "T", "T1" }; -const static char* const typeNameListCast[2] = { "T1", "T2" }; -const static char* const typeNameListIsNan[2] = { "T1", "T2" }; +const static char* const typeNameListT1T2[2] = { "T1", "T2" }; const static char* const typeNameListConstantOfShape[2] = { "T1", "T2" }; const static char* const typeNameListScatterGather[2] = { "T", "Tind" }; +const static char* const typeNameListScatterGatherND[2] = { "T" }; // Tind is curiously missing, only allowing 64-bit. const static char* const typeNameListQuantize[2] = { "T1", "T2" }; const static char* const typeNameListWhere[2] = { "B", "T" }; const static char* const typeNameListOneHot[3] = { "T1", "T2", "T3" }; const static char* const typeNameListEyeLike[1] = { "T2" }; const static SupportedTensorDataTypes supportedTypeListAll[1] = {SupportedTensorDataTypes::All}; const static SupportedTensorDataTypes supportedTypeListFloat16to32[1] = {SupportedTensorDataTypes::Float16to32}; +const static SupportedTensorDataTypes supportedTypeListInt8to32[1] = {SupportedTensorDataTypes::Int8to32}; const static SupportedTensorDataTypes supportedTypeListInt32to64AndFloat16to32[1] = {SupportedTensorDataTypes::Int32to64|SupportedTensorDataTypes::Float16to32}; const static SupportedTensorDataTypes supportedTypeListNumericDefault[1] = { SupportedTensorDataTypes::NumericDefault }; const static SupportedTensorDataTypes supportedTypeListAllScalars[1] = { SupportedTensorDataTypes::AllScalars }; @@ -228,10 +246,12 @@ const static SupportedTensorDataTypes supportedTypeListTopK[2] = {SupportedTenso const static SupportedTensorDataTypes supportedTypeListIndices[1] = { SupportedTensorDataTypes::Int32|SupportedTensorDataTypes::Int64 }; const static SupportedTensorDataTypes supportedTypeListCast[2] = { SupportedTensorDataTypes::AllScalars, SupportedTensorDataTypes::Scalars8to32 }; const static SupportedTensorDataTypes supportedTypeListScatterGather[2] = { SupportedTensorDataTypes::NumericDefault, SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::Int64 }; +const static SupportedTensorDataTypes supportedTypeListScatterGatherND[2] = { SupportedTensorDataTypes::NumericDefault }; const static SupportedTensorDataTypes supportedTypeListQuantizeLinear[2] = { SupportedTensorDataTypes::Float32 | SupportedTensorDataTypes::Int32, SupportedTensorDataTypes::UInt8 | SupportedTensorDataTypes::Int8 }; const static SupportedTensorDataTypes supportedTypeListDequantizeLinear[2] = { SupportedTensorDataTypes::Float32, SupportedTensorDataTypes::UInt8 | SupportedTensorDataTypes::Int8 | SupportedTensorDataTypes::Int32 }; const static SupportedTensorDataTypes supportedTypeListQuantize[2] = { SupportedTensorDataTypes::Float32, SupportedTensorDataTypes::UInt8 }; const static SupportedTensorDataTypes supportedTypeListIsNan[2] = { SupportedTensorDataTypes::Float16to32, SupportedTensorDataTypes::UInt8 }; +const static SupportedTensorDataTypes supportedTypeListIsInf[2] = { SupportedTensorDataTypes::Float16to32, SupportedTensorDataTypes::UInt8 }; const static SupportedTensorDataTypes supportedTypeListConstantOfShape[2] = { SupportedTensorDataTypes::Int32|SupportedTensorDataTypes::Int64, SupportedTensorDataTypes::Float16to32 }; const static SupportedTensorDataTypes supportedTypeListWhere[2] = { SupportedTensorDataTypes::UInt8, SupportedTensorDataTypes::Float16to32 }; const static SupportedTensorDataTypes supportedTypeListOneHot[3] = /* indices, depth, values */ { SupportedTensorDataTypes::Int32to64, SupportedTensorDataTypes::AllScalars, SupportedTensorDataTypes::Float16to32 }; @@ -272,11 +292,20 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO( 7, Conv, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ConvTranspose, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, AveragePool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, +#if 0 + // TODO:DwayneR add ceil mode https://microsoft.visualstudio.com/OS/_workitems/edit/24674310 + {REG_INFO( 10, AveragePool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, AveragePool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, +#endif {REG_INFO( 7, GlobalAveragePool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 8, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported, {}, std::nullopt, QueryMaxPool)}, {REG_INFO( 10, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported, {}, std::nullopt, QueryMaxPool)}, - +#if 0 + // TODO:DwayneR add ceil mode https://microsoft.visualstudio.com/OS/_workitems/edit/24674310 + {REG_INFO( 11, MaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported, {}, std::nullopt, QueryMaxPool)}, +#endif + {REG_INFO( 7, GlobalMaxPool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, LpPool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, GlobalLpPool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, @@ -294,13 +323,23 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl // Data Reorganization Layers {REG_INFO( 7, Split, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported)}, + {REG_INFO( 11, Split, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported)}, // Adds negative axis. {REG_INFO( 7, Transpose, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Concat, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported)}, + {REG_INFO( 11, Concat, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported)}, // Adds negative axis. {REG_INFO_VER( 7, Slice, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported)}, - {REG_INFO_VER( 10, Slice, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported, {1, 2, 3}, std::nullopt, QuerySlice)}, + {REG_INFO_VER( 10, Slice, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported, {1, 2, 3}, std::nullopt, QuerySlice)}, + {REG_INFO_VER( 11, Slice, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported, {1, 2, 3}, std::nullopt, QuerySlice)}, // Adds negative axes. {REG_INFO( 7, Pad, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, +#if 0 // TODO:DwayneR Pads and Value are inputs. https://github.com/onnx/onnx/blob/master/docs/Changelog.md#Pad-11 + {REG_INFO( 11, Pad, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, +#endif {REG_INFO( 7, SpaceToDepth, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, DepthToSpace, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, +#if 0 + // TODO:Dwayner https://microsoft.visualstudio.com/OS/_workitems/edit/24672169 + {REG_INFO( 11, DepthToSpace, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, +#endif {REG_INFO( 7, Tile, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported, {1})}, {REG_INFO( 8, Expand, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported, {1})}, {REG_INFO( 9, ConstantOfShape, typeNameListConstantOfShape, supportedTypeListConstantOfShape, DmGraphSupport::NotSupported, {0})}, @@ -312,8 +351,13 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO_ID( 7, Identity, typeNameListDefault, supportedTypeListAllScalars, DmGraphSupport::Supported)}, {REG_INFO_ID( 7, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmGraphSupport::Supported)}, {REG_INFO_ID( 9, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmGraphSupport::Supported)}, + //!!!TODO:::DwayneR check remaining 11's for other work besides negative axes. + //Also verify that negative axes are handled. + {REG_INFO_ID( 11, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmGraphSupport::Supported)}, {REG_INFO_ID( 7, Squeeze, typeNameListDefault, supportedTypeListAllScalars, DmGraphSupport::Supported)}, + {REG_INFO_ID( 11, Squeeze, typeNameListDefault, supportedTypeListAllScalars, DmGraphSupport::Supported)}, {REG_INFO_ID( 7, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmGraphSupport::Supported)}, + {REG_INFO_ID( 11, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmGraphSupport::Supported)}, {REG_INFO_ID( 7, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmGraphSupport::Supported, {1})}, // Elementwise @@ -326,6 +370,9 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO( 7, Ceil, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Floor, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Clip, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, +#if 0 // TODO:DwayneR https://microsoft.visualstudio.com/OS/_workitems/edit/24674103 + {REG_INFO( 11, Clip, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, +#endif {REG_INFO( 7, Add, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Sub, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Mul, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, @@ -350,7 +397,7 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO_MS( 1, QuantizeLinear, typeNameListQuantize, supportedTypeListQuantize, DmGraphSupport::Supported)}, {REG_INFO_MS( 1, DequantizeLinear, typeNameListQuantize, supportedTypeListQuantize, DmGraphSupport::Supported)}, {REG_INFO( 9, Sign, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported)}, - {REG_INFO( 9, IsNan, typeNameListIsNan, supportedTypeListIsNan, DmGraphSupport::Supported)}, + {REG_INFO( 9, IsNan, typeNameListT1T2, supportedTypeListIsNan, DmGraphSupport::Supported)}, {REG_INFO( 9, Sinh, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 9, Cosh, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 9, Asinh, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, @@ -359,18 +406,31 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO( 9, Erf, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 9, Where, typeNameListWhere, supportedTypeListWhere, DmGraphSupport::Supported)}, {REG_INFO( 7, ReduceSum, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReduceSum, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ReduceMean, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReduceMean, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ReduceProd, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReduceProd, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ReduceLogSum, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReduceLogSum, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ReduceLogSumExp, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReduceLogSumExp, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ReduceSumSquare, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReduceSumSquare, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ReduceL1, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReduceL1, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ReduceL2, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReduceL2, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ReduceMax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReduceMax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ReduceMin, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReduceMin, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ArgMax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ArgMax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ArgMin, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ArgMin, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Gemm, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, Gemm, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 9, Gemm, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Neg, typeNameListDefault, supportedTypeListSigned, DmGraphSupport::Supported)}, {REG_INFO( 7, Greater, typeNameListLogicalComparison, supportedTypeListLogicalComparison7,DmGraphSupport::Supported)}, @@ -405,8 +465,11 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO( 7, Elu, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Selu, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Softsign, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, Softplus, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 7, ParametricSoftplus, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, @@ -416,12 +479,19 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl // Uncategorized {REG_INFO( 7, MatMul, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO( 9, MatMul, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 7, Cast, typeNameListCast, supportedTypeListCast, DmGraphSupport::Supported)}, - {REG_INFO( 9, Cast, typeNameListCast, supportedTypeListCast, DmGraphSupport::Supported)}, + {REG_INFO( 7, Cast, typeNameListT1T2, supportedTypeListCast, DmGraphSupport::Supported)}, + {REG_INFO( 9, Cast, typeNameListT1T2, supportedTypeListCast, DmGraphSupport::Supported)}, {REG_INFO( 7, MemcpyFromHost, typeNameListDefault, supportedTypeListAll)}, {REG_INFO( 7, MemcpyToHost, typeNameListDefault, supportedTypeListAll)}, {REG_INFO( 7, TopK, typeNameListTopK, supportedTypeListTopK, DmGraphSupport::Supported)}, +#if 0 + // TODO:Dwayner https://microsoft.visualstudio.com/OS/_workitems/edit/24674287 + {REG_INFO( 10, TopK, typeNameListTopK, supportedTypeListTopK, DmGraphSupport::Supported)}, + // TODO:Dwayner https://microsoft.visualstudio.com/OS/_workitems/edit/24671996 + {REG_INFO( 11, TopK, typeNameListTopK, supportedTypeListTopK, DmGraphSupport::Supported)}, +#endif {REG_INFO( 9, OneHot, typeNameListOneHot, supportedTypeListOneHot, DmGraphSupport::Supported, {1})}, + {REG_INFO( 11, OneHot, typeNameListOneHot, supportedTypeListOneHot, DmGraphSupport::Supported, {1})}, // Fused operators {REG_INFO_MSDML(1, FusedConv, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, @@ -434,51 +504,58 @@ const static OperatorRegistrationInformation operatorRegistrationInformationTabl {REG_INFO_MSDML(1, FusedAdd, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, {REG_INFO_MSDML(1, FusedSum, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported, {}, 2)}, + {REG_INFO( 10, IsInf, typeNameListT1T2, supportedTypeListIsInf, DmGraphSupport::Supported)}, + {REG_INFO( 10, Mod, typeNameListDefault, supportedTypeListNumericDefault, DmGraphSupport::Supported)}, + {REG_INFO( 11, BitShift, typeNameListDefault, supportedTypeListInt8to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, Round, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, ReverseSequence, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, // TODO::: data types + {REG_INFO( 11, CumSum, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, +// {REG_INFO( 11, Range, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported), {0,1,2}}, #if 0 {REG_INFO( 9, MaxUnpool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported, {2})}, - {REG_INFO( 10, IsInf, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 10, Mod, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Argmax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Argmin, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, AveragePool, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, BitShift, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Clip, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Compress, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Concat, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, CumSum, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, DepthToSpace, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Flatten, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Gather, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, GatherElements, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, GatherND, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Gemm, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, OneHot, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Pad, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Range, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReduceL1, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReduceL2, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReduceLogSum, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReduceLogSumExp, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReduceMax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReduceMean, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReduceMin, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReduceProd, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReduceSum, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReduceSumSquare, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Resize, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ReverseSequence, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Round, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Scan, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ScatterElements, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, ScatterND, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Slice, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Split, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Squeeze, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, TopK, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, - {REG_INFO( 11, Unsqueeze, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + {REG_INFO( 11, Gather, typeNameListScatterGather, supportedTypeListScatterGather, DmGraphSupport::Supported)}, + -{REG_INFO( 11, GatherElements, typeNameListScatterGather, supportedTypeListScatterGather, DmGraphSupport::Supported)}, + -{REG_INFO( 11, GatherND, typeNameListScatterGatherND, supportedTypeListScatterGatherND, DmGraphSupport::Supported)}, + {REG_INFO( 11, Resize, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported, {1})}, + -{REG_INFO( 11, ScatterElements, typeNameListScatterGather, supportedTypeListScatterGather, DmGraphSupport::Supported)}, + -{REG_INFO( 11, ScatterND, typeNameListScatterGatherND, supportedTypeListScatterGatherND, DmGraphSupport::Supported)}, + + {REG_INFO( 11, QLinearConv, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + // T1 : tensor(int8), tensor(uint8) + // Constrain input type to 8-bit integer tensor. + // T2 : tensor(int8), tensor(uint8) + // Constrain filter type to 8-bit integer tensor. + // T3 : tensor(int8), tensor(uint8) + // Constrain output type to 8-bit integer tensor. + // T4 : tensor(int32) + // Constrain bias type to 32-bit integer tensor. + {REG_INFO( 11, QLinearMatMul, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + // T1 : tensor(int8), tensor(uint8) + // Constrain input a and its zero point data type to 8-bit integer tensor. + // T2 : tensor(int8), tensor(uint8) + // Constrain input b and its zero point data type to 8-bit integer tensor. + // T3 : tensor(int8), tensor(uint8) + // Constrain output y and its zero point data type to 8-bit integer tensor. + {REG_INFO( 11, DynamicQuantizeLinear, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + //T1 : tensor(float) + //Constrain 'x' to float tensor. + //T2 : tensor(uint8) + + {REG_INFO( 11, MatMulInteger, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + MatMulInteger + //T1 : tensor(int8), tensor(uint8) + //Constrain input A data type to 8-bit integer tensor. + //T2 : tensor(int8), tensor(uint8) + //Constrain input B data type to 8-bit integer tensor. + //T3 : tensor(int32) + //Constrain output Y data type as 32-bit integer tensor. + {REG_INFO( 11, ConvInteger, typeNameListDefault, supportedTypeListFloat16to32, DmGraphSupport::Supported)}, + // T1 : tensor(int8), tensor(uint8) + // Constrain input x and its zero point data type to 8-bit integer tensor. + // T2 : tensor(int8), tensor(uint8) + // Constrain input w and its zero point data type to 8-bit integer tensor. + // T3 : tensor(int32) + // Constrain output y data type to 32-bit integer tensor. #endif }; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp index 1afb7d83e6..4bf97d889c 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp @@ -277,10 +277,14 @@ DML_TENSOR_DESC TensorDesc::GetDmlDesc() // requires coercion by the caller. void TensorDesc::ForceUnsignedDataType() { - static_assert(ApiTraits::EnumValueCount == 9, "New tensor data type. Update cases."); + static_assert(ApiTraits::EnumValueCount == 12, "New tensor data type. Update cases."); switch (m_bufferTensorDesc.DataType) { + case DML_TENSOR_DATA_TYPE_INT64: + m_bufferTensorDesc.DataType = DML_TENSOR_DATA_TYPE_UINT64; + break; + case DML_TENSOR_DATA_TYPE_INT32: m_bufferTensorDesc.DataType = DML_TENSOR_DATA_TYPE_UINT32; break; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h index dc7def086d..02caf8f28b 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h @@ -22,14 +22,18 @@ namespace AttrName static constexpr const char* CeilMode = "ceil_mode"; static constexpr const char* Clip = "clip"; static constexpr const char* CountIncludePad = "count_include_pad"; + static constexpr const char* DetectPositive = "detect_negative"; + static constexpr const char* DetectNegative = "detect_negative "; static constexpr const char* Dilations = "dilations"; static constexpr const char* Direction = "direction"; static constexpr const char* Dtype = "dtype"; static constexpr const char* Ends = "ends"; static constexpr const char* Epsilon = "epsilon"; static constexpr const char* Exponent = "exponent"; + static constexpr const char* Fmod = "fmod"; static constexpr const char* Gamma = "gamma"; static constexpr const char* Group = "group"; + static constexpr const char* Exclusive = "exclusive"; static constexpr const char* HeightScale = "height_scale"; static constexpr const char* HiddenSize = "hidden_size"; static constexpr const char* High = "high"; @@ -50,6 +54,7 @@ namespace AttrName static constexpr const char* OutputPadding = "output_padding"; static constexpr const char* Pads = "pads"; static constexpr const char* PooledShape = "pooled_shape"; + static constexpr const char* Reverse = "reverse"; static constexpr const char* SampleSize = "sample_size"; static constexpr const char* Scale = "scale"; static constexpr const char* Scales = "scales"; @@ -64,6 +69,7 @@ namespace AttrName static constexpr const char* StorageOrder = "storage_order"; static constexpr const char* Strides = "strides"; static constexpr const char* Tiles = "tiles"; + static constexpr const char* TimeAxis = "time_axis"; static constexpr const char* To = "to"; static constexpr const char* TransA = "transA"; static constexpr const char* TransB = "transB"; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h index aa8486117f..f61f63b69a 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h @@ -528,6 +528,7 @@ public: { ends.push_back(gsl::narrow_cast(endsData[i])); } + uint32_t inputCount = operatorInfo.GetInputCount(); if (inputCount > 3) { @@ -1193,6 +1194,7 @@ using ShapeInferenceHelper_Transpose = TransposeHelper; using ShapeInferenceHelper_Concat = ConcatHelper; using ShapeInferenceHelper_Slice7 = SliceHelper; using ShapeInferenceHelper_Slice10 = Slice10Helper; +using ShapeInferenceHelper_Slice11 = Slice10Helper; // 11 and 10 are identical. using ShapeInferenceHelper_Pad = PaddingHelper; using ShapeInferenceHelper_SpaceToDepth = SpaceToDepthHelper; using ShapeInferenceHelper_DepthToSpace = DepthToSpaceHelper; @@ -1250,6 +1252,10 @@ using ShapeInferenceHelper_Asinh = GetBroadcastedOutputShapeHelper; using ShapeInferenceHelper_Acosh = GetBroadcastedOutputShapeHelper; using ShapeInferenceHelper_Atanh = GetBroadcastedOutputShapeHelper; using ShapeInferenceHelper_Where = GetBroadcastedOutputShapeHelper; +using ShapeInferenceHelper_IsInf = GetBroadcastedOutputShapeHelper; +using ShapeInferenceHelper_Mod = GetBroadcastedOutputShapeHelper; +using ShapeInferenceHelper_BitShift= GetBroadcastedOutputShapeHelper; +using ShapeInferenceHelper_Round = GetBroadcastedOutputShapeHelper; using ShapeInferenceHelper_ReduceSum = ReduceHelper; using ShapeInferenceHelper_ReduceMean = ReduceHelper; @@ -1302,6 +1308,10 @@ using ShapeInferenceHelper_RandomNormal = RandomNormalHelper; using ShapeInferenceHelper_RandomNormalLike = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Multinomial = MultinomialHelper; +using ShapeInferenceHelper_ReverseSequence = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_CumSum = GetOutputShapeAsInputShapeHelper; +// TODO::: using ShapeInferenceHelper_ShapeInferenceHelper_Range = ... + using ShapeInferenceHelper_FusedConv = ConvHelper; using ShapeInferenceHelper_FusedConvTranspose = ConvTransposeHelper; using ShapeInferenceHelper_FusedInstanceNormalization = GetOutputShapeAsInputShapeHelper; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h index 7fff38af20..8e55c78497 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorRegistration.h @@ -189,8 +189,8 @@ namespace OperatorHelper namespace OnnxOperatorSet11 { - static const int sc_sinceVer_Argmax = 11; - static const int sc_sinceVer_Argmin = 11; + static const int sc_sinceVer_ArgMax = 11; + static const int sc_sinceVer_ArgMin = 11; static const int sc_sinceVer_AveragePool = 11; static const int sc_sinceVer_BitShift = 11; static const int sc_sinceVer_Clip = 11; @@ -206,6 +206,7 @@ namespace OperatorHelper static const int sc_sinceVer_Gemm = 11; static const int sc_sinceVer_Hardmax = 11; static const int sc_sinceVer_LogSoftmax = 11; + static const int sc_sinceVer_MaxPool = 11; static const int sc_sinceVer_OneHot = 11; static const int sc_sinceVer_Pad = 11; static const int sc_sinceVer_Range = 11;