mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
Add kernels, including stubs.
This commit is contained in:
parent
be0c0da9ef
commit
d489288e3c
23 changed files with 1393 additions and 67 deletions
|
|
@ -13,7 +13,7 @@ struct EnumTraits
|
|||
template <>
|
||||
struct EnumTraits<DML_TENSOR_DATA_TYPE>
|
||||
{
|
||||
static constexpr auto ValueCount = 9;
|
||||
static constexpr auto ValueCount = 12;
|
||||
};
|
||||
|
||||
template <>
|
||||
|
|
@ -25,7 +25,7 @@ struct EnumTraits<DML_TENSOR_TYPE>
|
|||
template <>
|
||||
struct EnumTraits<DML_OPERATOR_TYPE>
|
||||
{
|
||||
static constexpr auto ValueCount = 97;
|
||||
static constexpr auto ValueCount = 107;
|
||||
static constexpr size_t ActivationFunctionCount = 19;
|
||||
};
|
||||
|
||||
|
|
@ -90,6 +90,24 @@ struct EnumTraits<DML_FEATURE_LEVEL>
|
|||
static constexpr auto ValueCount = 2;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct EnumTraits<DML_IS_INFINITY_MODE>
|
||||
{
|
||||
static constexpr auto ValueCount = 3;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct EnumTraits<DML_AXIS_DIRECTION>
|
||||
{
|
||||
static constexpr auto ValueCount = 2;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct EnumTraits<DML_ROUNDING_MODE>
|
||||
{
|
||||
static constexpr auto ValueCount = 3;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
constexpr auto EnumValueCount = EnumTraits<T>::ValueCount;
|
||||
|
||||
|
|
@ -610,6 +628,66 @@ struct OperatorDescTraits<DML_RESAMPLE_OPERATOR_DESC>
|
|||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_RESAMPLE;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_ELEMENT_WISE_ROUND_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_ROUND;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_IS_INFINITY;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_FILL_VALUE_CONSTANT_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_FILL_VALUE_CONSTANT;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_FILL_VALUE_SEQUENCE_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_FILL_VALUE_SEQUENCE;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_CUMULATIVE_SUMMATION_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_CUMULATIVE_SUMMATION;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_REVERSE_SUBSEQUENCES_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_REVERSE_SUBSEQUENCES;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_ACTIVATION_ELU_OPERATOR_DESC>
|
||||
{
|
||||
|
|
@ -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>(visitor), DML_ONE_HOT_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_RESAMPLE:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_RESAMPLE_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_ELEMENT_WISE_ROUND:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_ELEMENT_WISE_ROUND_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_ELEMENT_WISE_IS_INFINITY:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_FILL_VALUE_CONSTANT:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_FILL_VALUE_CONSTANT_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_FILL_VALUE_SEQUENCE:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_FILL_VALUE_SEQUENCE_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_CUMULATIVE_SUMMATION:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_CUMULATIVE_SUMMATION_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_REVERSE_SUBSEQUENCES:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_REVERSE_SUBSEQUENCES_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_ACTIVATION_ELU:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_ACTIVATION_ELU_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_ACTIVATION_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 "<unknown>";
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -738,6 +738,90 @@ inline std::vector<OperatorField> GetFields(const DML_RESAMPLE_OPERATOR_DESC& de
|
|||
OperatorField(&DML_RESAMPLE_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast<const FLOAT*>(desc.Scales), desc.ScaleCount)),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> 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<const DML_TENSOR_DESC*>(desc.ATensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.BTensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> 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<const DML_TENSOR_DESC*>(desc.ATensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.BTensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_ELEMENT_WISE_ROUND_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
OperatorField(&DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.InputTensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<UINT>(desc.RoundingMode))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
OperatorField(&DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.InputTensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<UINT>(desc.InfinityMode))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
OperatorField(&DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.ATensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.BTensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
OperatorField(&DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.ATensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.BTensor))),
|
||||
OperatorField(&DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_FILL_VALUE_CONSTANT_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
OperatorField(&DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
OperatorField(&DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<UINT>(desc.ValueDataType))),
|
||||
OperatorField(&DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<DML_SCALAR_UNION>(desc.Value))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_FILL_VALUE_SEQUENCE_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
OperatorField(&DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
OperatorField(&DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<UINT>(desc.ValueDataType))),
|
||||
OperatorField(&DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<DML_SCALAR_UNION>(desc.ValueStart))),
|
||||
OperatorField(&DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast<DML_SCALAR_UNION>(desc.ValueDelta))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_CUMULATIVE_SUMMATION_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
OperatorField(&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.InputTensor))),
|
||||
OperatorField(&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
OperatorField(&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<UINT>(desc.Axis))),
|
||||
OperatorField(&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast<UINT>(desc.HasExclusiveSum))),
|
||||
OperatorField(&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast<UINT>(desc.AxisDirection))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_REVERSE_SUBSEQUENCES_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
OperatorField(&DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.InputTensor))),
|
||||
OperatorField(&DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.SequenceLengthsTensor))),
|
||||
OperatorField(&DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
OperatorField(&DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast<UINT>(desc.Axis))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> 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<const DML_RESAMPLE_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_LEFT:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_ELEMENT_WISE_BIT_SHIFT_LEFT_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_ELEMENT_WISE_BIT_SHIFT_RIGHT:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_ELEMENT_WISE_BIT_SHIFT_RIGHT_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_ELEMENT_WISE_ROUND:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_ELEMENT_WISE_ROUND_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_ELEMENT_WISE_ROUND_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_ELEMENT_WISE_IS_INFINITY:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_ELEMENT_WISE_MODULUS_TRUNCATE:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_ELEMENT_WISE_MODULUS_TRUNCATE_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_ELEMENT_WISE_MODULUS_FLOOR:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_ELEMENT_WISE_MODULUS_FLOOR_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_FILL_VALUE_CONSTANT:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_FILL_VALUE_CONSTANT_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_FILL_VALUE_CONSTANT_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_FILL_VALUE_SEQUENCE:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_FILL_VALUE_SEQUENCE_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_FILL_VALUE_SEQUENCE_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_CUMULATIVE_SUMMATION:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_CUMULATIVE_SUMMATION_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_CUMULATIVE_SUMMATION_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_REVERSE_SUBSEQUENCES:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_REVERSE_SUBSEQUENCES_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_REVERSE_SUBSEQUENCES_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_ACTIVATION_ELU:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_ACTIVATION_ELU_OPERATOR_SCHEMA,
|
||||
|
|
|
|||
|
|
@ -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<AbstractOperatorDesc>; // DML_SCHEMA_FIELD_TYPE_OPERATOR_DESC
|
||||
using OperatorDescArray = std::optional<std::vector<AbstractOperatorDesc>>; // 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<std::vector<uint32_t>>; // DML_SCHEMA_FIELD_TYPE_UINT_ARRAY
|
||||
using IntArray = std::optional<std::vector<int32_t>>; // DML_SCHEMA_FIELD_TYPE_INT_ARRAY
|
||||
using FloatArray = std::optional<std::vector<float>>; // DML_SCHEMA_FIELD_TYPE_FLOAT_ARRAY
|
||||
using ScaleBias = std::optional<DML_SCALE_BIAS>; // 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<OperatorFieldTypes::UInt>(m_data); }
|
||||
OperatorFieldTypes::UInt& AsUInt() { return std::get<OperatorFieldTypes::UInt>(m_data); }
|
||||
|
||||
const OperatorFieldTypes::UInt64& AsUInt64() const { return std::get<OperatorFieldTypes::UInt64>(m_data); }
|
||||
OperatorFieldTypes::UInt64& AsUInt64() { return std::get<OperatorFieldTypes::UInt64>(m_data); }
|
||||
|
||||
const OperatorFieldTypes::Int& AsInt() const { return std::get<OperatorFieldTypes::Int>(m_data); }
|
||||
OperatorFieldTypes::Int& AsInt() { return std::get<OperatorFieldTypes::Int>(m_data); }
|
||||
|
||||
|
|
@ -89,6 +101,9 @@ public:
|
|||
const OperatorFieldTypes::UIntArray& AsUIntArray() const { return std::get<OperatorFieldTypes::UIntArray>(m_data); }
|
||||
OperatorFieldTypes::UIntArray& AsUIntArray() { return std::get<OperatorFieldTypes::UIntArray>(m_data); }
|
||||
|
||||
const OperatorFieldTypes::IntArray& AsIntArray() const { return std::get<OperatorFieldTypes::IntArray>(m_data); }
|
||||
OperatorFieldTypes::IntArray& AsIntArray() { return std::get<OperatorFieldTypes::IntArray>(m_data); }
|
||||
|
||||
const OperatorFieldTypes::FloatArray& AsFloatArray() const { return std::get<OperatorFieldTypes::FloatArray>(m_data); }
|
||||
OperatorFieldTypes::FloatArray& AsFloatArray() { return std::get<OperatorFieldTypes::FloatArray>(m_data); }
|
||||
|
||||
|
|
@ -98,6 +113,9 @@ public:
|
|||
const OperatorFieldTypes::Size2D& AsSize2D() const { return std::get<OperatorFieldTypes::Size2D>(m_data); }
|
||||
OperatorFieldTypes::Size2D& AsSize2D() { return std::get<OperatorFieldTypes::Size2D>(m_data); }
|
||||
|
||||
const OperatorFieldTypes::ScalarUnion& AsScalarUnion() const { return std::get<OperatorFieldTypes::ScalarUnion>(m_data); }
|
||||
OperatorFieldTypes::ScalarUnion& AsScalarUnion() { return std::get<OperatorFieldTypes::ScalarUnion>(m_data); }
|
||||
|
||||
private:
|
||||
const DML_SCHEMA_FIELD* m_schema;
|
||||
OperatorFieldVariant m_data;
|
||||
|
|
|
|||
|
|
@ -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<int32_t>(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);
|
||||
|
|
|
|||
|
|
@ -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<std::optional<uint32_t>> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'.
|
||||
std::vector<std::optional<uint32_t>> 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<uint32_t> 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<uint32_t>(indicesDimensions.size()),
|
||||
m_inputTensorDescs.front().GetDimensionCount()
|
||||
);
|
||||
|
||||
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> 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
|
||||
|
|
@ -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<int32_t>(AttrName::Exclusive, 0);
|
||||
int32_t isReversed = kernelCreationContext.GetOptionalAttribute<int32_t>(AttrName::Reverse, 0);
|
||||
int32_t onnxAxis = kernelCreationContext.GetOptionalAttribute<int32_t>(AttrName::Axis, -1);
|
||||
uint32_t dmlAxis = GetDmlAdjustedAxis(onnxAxis, kernelCreationContext, m_inputTensorDescs.front().GetDimensionCount());
|
||||
|
||||
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> 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
|
||||
|
|
@ -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<std::optional<uint32_t>> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'.
|
||||
std::vector<std::optional<uint32_t>> 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<uint32_t> 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<uint32_t>(indicesDimensions.size()),
|
||||
m_inputTensorDescs.front().GetDimensionCount()
|
||||
);
|
||||
|
||||
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> 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
|
||||
|
|
@ -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<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> outputDescs = GetDmlOutputDescs();
|
||||
|
||||
auto fmod = kernelInfo.GetOptionalAttribute<int>(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<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> 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<std::string>(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<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> outputDescs = GetDmlOutputDescs();
|
||||
|
||||
auto detectPositive = kernelInfo.GetOptionalAttribute<int>(AttrName::DetectPositive, 1);
|
||||
auto detectNegative = kernelInfo.GetOptionalAttribute<int>(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<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> 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_ELEMENT_WISE_SQRT_OPERATOR_DESC>);
|
||||
DML_OP_DEFINE_CREATION_FUNCTION(Reciprocal, DmlOperatorElementwiseUnary<DML_ELEMENT_WISE_RECIP_OPERATOR_DESC>);
|
||||
|
|
@ -582,6 +686,10 @@ DML_OP_DEFINE_CREATION_FUNCTION(Pow, DmlOperatorElementwisePow);
|
|||
DML_OP_DEFINE_CREATION_FUNCTION(QuantizeLinear, DmlOperatorElementwiseQLinear<DML_ELEMENT_WISE_QUANTIZE_LINEAR_OPERATOR_DESC>);
|
||||
DML_OP_DEFINE_CREATION_FUNCTION(DequantizeLinear, DmlOperatorElementwiseQLinear<DML_ELEMENT_WISE_DEQUANTIZE_LINEAR_OPERATOR_DESC>);
|
||||
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<DML_ELEMENT_WISE_ADD1_OPERATOR_DESC>);
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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<std::optional<uint32_t>> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'.
|
||||
std::vector<std::optional<uint32_t>> 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<uint32_t> 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<uint32_t>(indicesDimensions.size()),
|
||||
m_inputTensorDescs.front().GetDimensionCount()
|
||||
);
|
||||
|
||||
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> 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
|
||||
|
|
@ -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<std::optional<uint32_t>> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'.
|
||||
std::vector<std::optional<uint32_t>> 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<uint32_t> 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<uint32_t>(indicesDimensions.size()),
|
||||
m_inputTensorDescs.front().GetDimensionCount()
|
||||
);
|
||||
|
||||
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> 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
|
||||
|
|
@ -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<std::optional<uint32_t>> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'.
|
||||
std::vector<std::optional<uint32_t>> 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<uint32_t> 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<uint32_t>(indicesDimensions.size()),
|
||||
m_inputTensorDescs.front().GetDimensionCount()
|
||||
);
|
||||
|
||||
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> 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
|
||||
|
|
@ -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<std::optional<uint32_t>> inputIndices = { 0, 2 }; // The second tensor ('depth') is not bound, just 'indices' and 'values'.
|
||||
std::vector<std::optional<uint32_t>> 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<uint32_t> 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<uint32_t>(indicesDimensions.size()),
|
||||
m_inputTensorDescs.front().GetDimensionCount()
|
||||
);
|
||||
|
||||
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> 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
|
||||
|
|
@ -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<uint32_t> inputDimensions = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0);
|
||||
std::vector<uint32_t> sequenceLengthDimensions = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(1);
|
||||
|
||||
// Read axis.
|
||||
int32_t onnxAxis = kernelCreationContext.GetOptionalAttribute<int32_t>(AttrName::TimeAxis, 0);
|
||||
onnxAxis = HandleNegativeAxis(onnxAxis, static_cast<uint32_t>(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<uint32_t> 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<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -277,10 +277,14 @@ DML_TENSOR_DESC TensorDesc::GetDmlDesc()
|
|||
// requires coercion by the caller.
|
||||
void TensorDesc::ForceUnsignedDataType()
|
||||
{
|
||||
static_assert(ApiTraits::EnumValueCount<DML_TENSOR_DATA_TYPE> == 9, "New tensor data type. Update cases.");
|
||||
static_assert(ApiTraits::EnumValueCount<DML_TENSOR_DATA_TYPE> == 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;
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -528,6 +528,7 @@ public:
|
|||
{
|
||||
ends.push_back(gsl::narrow_cast<int32_t>(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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Reference in a new issue