mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-23 19:32:23 +00:00
Merged PR 5051981: ArgMin/Argmax 12: add select_last_index attribute
Add select_last_index atribute to argmin/argmax. normally the algorithm takes the index of the first occurrence of min/max, but with this attribute enabled it will take the index of the last occurrence. [windowsai pr](https://microsoft.visualstudio.com/WindowsAI/_git/WindowsAI/pullrequest/5051953) Related work items: #27469709
This commit is contained in:
parent
d350d9ac5e
commit
c8155c79cb
8 changed files with 164 additions and 15 deletions
|
|
@ -24,7 +24,7 @@ struct EnumTraits<DML_TENSOR_TYPE>
|
|||
template <>
|
||||
struct EnumTraits<DML_OPERATOR_TYPE>
|
||||
{
|
||||
static constexpr auto ValueCount = 121;
|
||||
static constexpr auto ValueCount = 124;
|
||||
static constexpr size_t ActivationFunctionCount = 20;
|
||||
};
|
||||
|
||||
|
|
@ -62,7 +62,7 @@ struct EnumTraits<DML_CONVOLUTION_DIRECTION>
|
|||
template <>
|
||||
struct EnumTraits<DML_PADDING_MODE>
|
||||
{
|
||||
static constexpr auto ValueCount = 3;
|
||||
static constexpr auto ValueCount = 4;
|
||||
};
|
||||
|
||||
template <>
|
||||
|
|
@ -86,7 +86,7 @@ struct EnumTraits<DML_FEATURE>
|
|||
template <>
|
||||
struct EnumTraits<DML_FEATURE_LEVEL>
|
||||
{
|
||||
static constexpr auto ValueCount = 2;
|
||||
static constexpr auto ValueCount = 4;
|
||||
};
|
||||
|
||||
template <>
|
||||
|
|
@ -113,6 +113,12 @@ struct EnumTraits<DML_ROUNDING_MODE>
|
|||
static constexpr auto ValueCount = 3;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct EnumTraits<DML_RANDOM_GENERATOR_TYPE>
|
||||
{
|
||||
static constexpr auto ValueCount = 1;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
constexpr auto EnumValueCount = EnumTraits<T>::ValueCount;
|
||||
|
||||
|
|
@ -405,6 +411,18 @@ struct OperatorDescTraits<DML_REDUCE_OPERATOR_DESC>
|
|||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_REDUCE;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_ARGMIN_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ARGMIN;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_ARGMAX_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ARGMAX;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_AVERAGE_POOLING_OPERATOR_DESC>
|
||||
{
|
||||
|
|
@ -759,6 +777,12 @@ struct OperatorDescTraits<DML_MEAN_VARIANCE_NORMALIZATION1_OPERATOR_DESC>
|
|||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_MEAN_VARIANCE_NORMALIZATION1;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_RESAMPLE1_OPERATOR_DESC>
|
||||
{
|
||||
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_RESAMPLE1;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorDescTraits<DML_MATRIX_MULTIPLY_INTEGER_OPERATOR_DESC>
|
||||
{
|
||||
|
|
@ -1143,6 +1167,18 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_REDUCE>
|
|||
using DescType = DML_REDUCE_OPERATOR_DESC;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ARGMIN>
|
||||
{
|
||||
using DescType = DML_ARGMIN_OPERATOR_DESC;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ARGMAX>
|
||||
{
|
||||
using DescType = DML_ARGMAX_OPERATOR_DESC;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_AVERAGE_POOLING>
|
||||
{
|
||||
|
|
@ -1497,6 +1533,12 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_MEAN_VARIANCE_NORMALIZ
|
|||
using DescType = DML_MEAN_VARIANCE_NORMALIZATION1_OPERATOR_DESC;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_RESAMPLE1>
|
||||
{
|
||||
using DescType = DML_RESAMPLE1_OPERATOR_DESC;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_MATRIX_MULTIPLY_INTEGER>
|
||||
{
|
||||
|
|
@ -1732,6 +1774,10 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args
|
|||
return std::invoke(std::forward<Visitor>(visitor), DML_GEMM_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_REDUCE:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_REDUCE_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_ARGMIN:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_ARGMIN_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_ARGMAX:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_ARGMAX_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_AVERAGE_POOLING:
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_AVERAGE_POOLING_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
case DML_OPERATOR_LP_POOLING:
|
||||
|
|
@ -1951,6 +1997,8 @@ inline gsl::czstring ToString(DML_OPERATOR_TYPE value)
|
|||
case DML_OPERATOR_CONVOLUTION: return "DML_OPERATOR_CONVOLUTION";
|
||||
case DML_OPERATOR_GEMM: return "DML_OPERATOR_GEMM";
|
||||
case DML_OPERATOR_REDUCE: return "DML_OPERATOR_REDUCE";
|
||||
case DML_OPERATOR_ARGMIN: return "DML_OPERATOR_ARGMIN";
|
||||
case DML_OPERATOR_ARGMAX: return "DML_OPERATOR_ARGMAX";
|
||||
case DML_OPERATOR_AVERAGE_POOLING: return "DML_OPERATOR_AVERAGE_POOLING";
|
||||
case DML_OPERATOR_LP_POOLING: return "DML_OPERATOR_LP_POOLING";
|
||||
case DML_OPERATOR_MAX_POOLING: return "DML_OPERATOR_MAX_POOLING";
|
||||
|
|
|
|||
|
|
@ -627,6 +627,38 @@ constexpr DML_OPERATOR_SCHEMA DML_REDUCE_OPERATOR_SCHEMA {
|
|||
DML_REDUCE_OPERATOR_SCHEMA_FIELDS,
|
||||
};
|
||||
|
||||
constexpr DML_SCHEMA_FIELD DML_ARGMIN_OPERATOR_SCHEMA_FIELDS[5] {
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisCount", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT_ARRAY, "Axes", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisDirection", false },
|
||||
};
|
||||
|
||||
constexpr DML_OPERATOR_SCHEMA DML_ARGMIN_OPERATOR_SCHEMA {
|
||||
"DML_OPERATOR_ARGMIN",
|
||||
DML_OPERATOR_ARGMIN,
|
||||
DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE,
|
||||
5,
|
||||
DML_ARGMIN_OPERATOR_SCHEMA_FIELDS,
|
||||
};
|
||||
|
||||
constexpr DML_SCHEMA_FIELD DML_ARGMAX_OPERATOR_SCHEMA_FIELDS[5] {
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisCount", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT_ARRAY, "Axes", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisDirection", false },
|
||||
};
|
||||
|
||||
constexpr DML_OPERATOR_SCHEMA DML_ARGMAX_OPERATOR_SCHEMA {
|
||||
"DML_OPERATOR_ARGMAX",
|
||||
DML_OPERATOR_ARGMAX,
|
||||
DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE,
|
||||
5,
|
||||
DML_ARGMAX_OPERATOR_SCHEMA_FIELDS,
|
||||
};
|
||||
|
||||
constexpr DML_SCHEMA_FIELD DML_AVERAGE_POOLING_OPERATOR_SCHEMA_FIELDS[8] {
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false },
|
||||
|
|
@ -1656,7 +1688,7 @@ constexpr DML_SCHEMA_FIELD DML_QUANTIZED_LINEAR_CONVOLUTION_OPERATOR_SCHEMA_FIEL
|
|||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "FilterTensor", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "FilterScaleTensor", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "FilterZeroPointTensor", true },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "BiasTensor", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "BiasTensor", true },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputScaleTensor", false },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputZeroPointTensor", true },
|
||||
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false },
|
||||
|
|
|
|||
|
|
@ -344,6 +344,26 @@ inline std::vector<OperatorField> GetFields(const DML_REDUCE_OPERATOR_DESC& desc
|
|||
OperatorField(&DML_REDUCE_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast<const UINT*>(desc.Axes), desc.AxisCount)),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_ARGMIN_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
OperatorField(&DML_ARGMIN_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.InputTensor))),
|
||||
OperatorField(&DML_ARGMIN_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
OperatorField(&DML_ARGMIN_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<UINT>(desc.AxisCount))),
|
||||
OperatorField(&DML_ARGMIN_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast<const UINT*>(desc.Axes), desc.AxisCount)),
|
||||
OperatorField(&DML_ARGMIN_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast<UINT>(desc.AxisDirection))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_ARGMAX_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
OperatorField(&DML_ARGMAX_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.InputTensor))),
|
||||
OperatorField(&DML_ARGMAX_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
|
||||
OperatorField(&DML_ARGMAX_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<UINT>(desc.AxisCount))),
|
||||
OperatorField(&DML_ARGMAX_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast<const UINT*>(desc.Axes), desc.AxisCount)),
|
||||
OperatorField(&DML_ARGMAX_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast<UINT>(desc.AxisDirection))),
|
||||
};
|
||||
}
|
||||
inline std::vector<OperatorField> GetFields(const DML_AVERAGE_POOLING_OPERATOR_DESC& desc)
|
||||
{
|
||||
return {
|
||||
|
|
@ -1211,6 +1231,8 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType)
|
|||
case DML_OPERATOR_CONVOLUTION: return DML_CONVOLUTION_OPERATOR_SCHEMA;
|
||||
case DML_OPERATOR_GEMM: return DML_GEMM_OPERATOR_SCHEMA;
|
||||
case DML_OPERATOR_REDUCE: return DML_REDUCE_OPERATOR_SCHEMA;
|
||||
case DML_OPERATOR_ARGMIN: return DML_ARGMIN_OPERATOR_SCHEMA;
|
||||
case DML_OPERATOR_ARGMAX: return DML_ARGMAX_OPERATOR_SCHEMA;
|
||||
case DML_OPERATOR_AVERAGE_POOLING: return DML_AVERAGE_POOLING_OPERATOR_SCHEMA;
|
||||
case DML_OPERATOR_LP_POOLING: return DML_LP_POOLING_OPERATOR_SCHEMA;
|
||||
case DML_OPERATOR_MAX_POOLING: return DML_MAX_POOLING_OPERATOR_SCHEMA;
|
||||
|
|
@ -1460,6 +1482,14 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc)
|
|||
return AbstractOperatorDesc(
|
||||
&DML_REDUCE_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_REDUCE_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_ARGMIN:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_ARGMIN_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_ARGMIN_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_ARGMAX:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_ARGMAX_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_ARGMAX_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
case DML_OPERATOR_AVERAGE_POOLING:
|
||||
return AbstractOperatorDesc(
|
||||
&DML_AVERAGE_POOLING_OPERATOR_SCHEMA,
|
||||
|
|
|
|||
|
|
@ -54,9 +54,19 @@ public:
|
|||
reducedDims);
|
||||
}
|
||||
|
||||
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> outputDescs = GetDmlOutputDescs();
|
||||
|
||||
// Zero the output tensor's memory for ArgMin & ArgMax, which produce INT64 output.
|
||||
if ((function == DML_REDUCE_FUNCTION_ARGMAX) || (function == DML_REDUCE_FUNCTION_ARGMIN))
|
||||
if (function == DML_REDUCE_FUNCTION_ARGMAX)
|
||||
{
|
||||
DML_ARGMAX_OPERATOR_DESC argmaxDesc;
|
||||
argmaxDesc.AxisDirection = static_cast<DML_AXIS_DIRECTION>(m_selectLastIndex);
|
||||
argmaxDesc.InputTensor = inputDescs.data();
|
||||
argmaxDesc.OutputTensor = outputDescs.data();
|
||||
argmaxDesc.Axes = dmlAxes.data();
|
||||
argmaxDesc.AxisCount = gsl::narrow_cast<uint32_t>(dmlAxes.size());
|
||||
|
||||
// If the 64-bit tensors were remapped to 32-bit, then we need to clear the upper 32-bits
|
||||
// of each element. If the device directly supports 64-bit elements, then no need.
|
||||
DmlOperator::Remap64bitDmlDataTypesTo32bitIfNeeded();
|
||||
|
|
@ -64,20 +74,42 @@ public:
|
|||
{
|
||||
m_zeroOperator = InitializeZeroInt64Tensor(m_outputTensorDescs[0].GetBufferSizeInBytes());
|
||||
}
|
||||
|
||||
DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ARGMAX, &argmaxDesc };
|
||||
SetDmlOperatorDesc(opDesc, kernelInfo);
|
||||
}
|
||||
else if (function == DML_REDUCE_FUNCTION_ARGMIN)
|
||||
{
|
||||
DML_ARGMIN_OPERATOR_DESC argminDesc;
|
||||
argminDesc.AxisDirection = static_cast<DML_AXIS_DIRECTION>(m_selectLastIndex);
|
||||
argminDesc.InputTensor = inputDescs.data();
|
||||
argminDesc.OutputTensor = outputDescs.data();
|
||||
argminDesc.Axes = dmlAxes.data();
|
||||
argminDesc.AxisCount = gsl::narrow_cast<uint32_t>(dmlAxes.size());
|
||||
|
||||
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
|
||||
std::vector<DML_TENSOR_DESC> outputDescs = GetDmlOutputDescs();
|
||||
// If the 64-bit tensors were remapped to 32-bit, then we need to clear the upper 32-bits
|
||||
// of each element. If the device directly supports 64-bit elements, then no need.
|
||||
DmlOperator::Remap64bitDmlDataTypesTo32bitIfNeeded();
|
||||
if (m_outputTensorDescs[0].WasRemapped64bitTo32bit())
|
||||
{
|
||||
m_zeroOperator = InitializeZeroInt64Tensor(m_outputTensorDescs[0].GetBufferSizeInBytes());
|
||||
}
|
||||
|
||||
DML_REDUCE_OPERATOR_DESC reduceDesc = {};
|
||||
reduceDesc.InputTensor = inputDescs.data();
|
||||
reduceDesc.OutputTensor = outputDescs.data();
|
||||
reduceDesc.Function = function;
|
||||
reduceDesc.Axes = dmlAxes.data();
|
||||
reduceDesc.AxisCount = gsl::narrow_cast<uint32_t>(dmlAxes.size());
|
||||
DML_OPERATOR_DESC opDesc = { DML_OPERATOR_ARGMIN, &argminDesc };
|
||||
SetDmlOperatorDesc(opDesc, kernelInfo);
|
||||
}
|
||||
else
|
||||
{
|
||||
DML_REDUCE_OPERATOR_DESC reduceDesc = {};
|
||||
reduceDesc.InputTensor = inputDescs.data();
|
||||
reduceDesc.OutputTensor = outputDescs.data();
|
||||
reduceDesc.Function = function;
|
||||
reduceDesc.Axes = dmlAxes.data();
|
||||
reduceDesc.AxisCount = gsl::narrow_cast<uint32_t>(dmlAxes.size());
|
||||
|
||||
DML_OPERATOR_DESC opDesc = { DML_OPERATOR_REDUCE, &reduceDesc };
|
||||
SetDmlOperatorDesc(opDesc, kernelInfo);
|
||||
DML_OPERATOR_DESC opDesc = { DML_OPERATOR_REDUCE, &reduceDesc };
|
||||
SetDmlOperatorDesc(opDesc, kernelInfo);
|
||||
}
|
||||
}
|
||||
|
||||
void Compute(const MLOperatorKernelContext& kernelContext) override
|
||||
|
|
|
|||
|
|
@ -475,8 +475,10 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation
|
|||
{REG_INFO( 12, ReduceMin, typeNameListDefault, supportedTypeListFloat16to32Int8, DmlGraphSupport::Supported)},
|
||||
{REG_INFO( 7, ArgMax, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)},
|
||||
{REG_INFO( 11, ArgMax, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)},
|
||||
{REG_INFO( 12, ArgMax, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)},
|
||||
{REG_INFO( 7, ArgMin, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)},
|
||||
{REG_INFO( 11, ArgMin, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)},
|
||||
{REG_INFO( 12, ArgMin, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported|DmlGraphSupport::Prefer64BitTensorsDirectly|DmlGraphSupport::SupportedWith64BitTensorsVia32BitStrides)},
|
||||
{REG_INFO( 7, Gemm, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
|
||||
{REG_INFO( 11, Gemm, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
|
||||
{REG_INFO( 9, Gemm, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ namespace AttrName
|
|||
static constexpr const char* InputForget = "input_forget";
|
||||
static constexpr const char* K = "k";
|
||||
static constexpr const char* KeepDims = "keepdims";
|
||||
static constexpr const char* SelectLastIndex = "select_last_index";
|
||||
static constexpr const char* KernelShape = "kernel_shape";
|
||||
static constexpr const char* LinearBeforeReset = "linear_before_reset";
|
||||
static constexpr const char* Lambda = "lambd"; // Deliberate typo to match ONNX spec.
|
||||
|
|
|
|||
|
|
@ -733,6 +733,7 @@ class ReduceHelperBase {
|
|||
template <typename Info_t, typename Shape_t>
|
||||
ReduceHelperBase(const Info_t& info, const Shape_t& shape, bool usingAxes) {
|
||||
m_keepDims = info.GetOptionalAttribute<int>(AttrName::KeepDims, 1);
|
||||
m_selectLastIndex = info.GetOptionalAttribute<int>(AttrName::SelectLastIndex, 0);
|
||||
if (usingAxes) {
|
||||
m_axes = info.GetOptionalAttributeVectorInt32(AttrName::Axes);
|
||||
} else {
|
||||
|
|
@ -751,6 +752,7 @@ class ReduceHelperBase {
|
|||
protected:
|
||||
std::vector<int> m_axes;
|
||||
int m_keepDims = 0;
|
||||
int m_selectLastIndex = 0;
|
||||
};
|
||||
|
||||
class ArgMinArgMaxHelper : public ReduceHelperBase {
|
||||
|
|
|
|||
|
|
@ -256,6 +256,8 @@ namespace OperatorHelper
|
|||
static const int sc_sinceVer_MaxPool = 12;
|
||||
static const int sc_sinceVer_ReduceMax = 12;
|
||||
static const int sc_sinceVer_ReduceMin = 12;
|
||||
static const int sc_sinceVer_ArgMin = 12;
|
||||
static const int sc_sinceVer_ArgMax = 12;
|
||||
} // namespace OnnxOperatorSet12
|
||||
|
||||
namespace MsftOperatorSet1
|
||||
|
|
|
|||
Loading…
Reference in a new issue