mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-27 20:02:15 +00:00
Add opset 13 registrations for Identity, IsNaN, NonZero, GatherND and Pad (#5513)
This commit is contained in:
parent
f207f0bf5e
commit
ad94a1dd6d
6 changed files with 175 additions and 81 deletions
|
|
@ -183,7 +183,7 @@ class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDoma
|
|||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 4, 10, Concat);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Gather);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, Dropout);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Identity);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 12, Identity);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, 10, Pad);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 4, Reshape_1);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 5, 12, Reshape);
|
||||
|
|
@ -225,8 +225,8 @@ class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOn
|
|||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, int32_t, Less);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, int64_t, Less);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, EyeLike);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, IsNaN);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MLFloat16, IsNaN);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, float, IsNaN);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, MLFloat16, IsNaN);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, Sign);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Shrink);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, float, Erf);
|
||||
|
|
@ -250,11 +250,11 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Ata
|
|||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Scan);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Scatter);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, TfIdfVectorizer);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, bool, NonZero);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, NonZero);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, NonZero);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t, NonZero);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, uint8_t, NonZero);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, bool, NonZero);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, float, NonZero);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, int32_t, NonZero);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, int64_t, NonZero);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, uint8_t, NonZero);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, string, Where);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, Where);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, Where);
|
||||
|
|
@ -344,7 +344,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Sc
|
|||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 12, Flatten);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Compress);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 12, Concat);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 12,Gather);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 12, Gather);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 12, Slice);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Split);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Squeeze);
|
||||
|
|
@ -372,7 +372,7 @@ class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDoma
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint8_t, BitShift);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint32_t, BitShift);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint64_t, BitShift);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Pad);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 12, Pad);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 11, GatherND);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Range);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Unique);
|
||||
|
|
@ -411,7 +411,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain,
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, int64_t, ReduceMin);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, int8_t, ReduceMin);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, uint8_t, ReduceMin);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, GatherND);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, 12, GatherND);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, Einsum);
|
||||
|
||||
// REVIEW(codemzs): ConstEigenVectorArrayMap.cast<MLFLoat16) does not seem to be supported.
|
||||
|
|
@ -536,6 +536,16 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, Ga
|
|||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, GatherElements);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, ScatterND);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, ScatterElements);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, Identity);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, float, IsNaN);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, MLFloat16, IsNaN);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, bool, NonZero);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, float, NonZero);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, int32_t, NonZero);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, int64_t, NonZero);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, uint8_t, NonZero);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, GatherND);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, Pad);
|
||||
|
||||
template <>
|
||||
KernelCreateInfo BuildKernelCreateInfo<void>() {
|
||||
|
|
@ -777,7 +787,8 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
Gather)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9,
|
||||
Dropout)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Identity)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 12,
|
||||
Identity)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, 10, Pad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 4,
|
||||
Reshape_1)>,
|
||||
|
|
@ -847,10 +858,10 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12, int64_t,
|
||||
Less)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, EyeLike)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float,
|
||||
IsNaN)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MLFloat16,
|
||||
IsNaN)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12,
|
||||
float, IsNaN)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12,
|
||||
MLFloat16, IsNaN)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12,
|
||||
Sign)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Shrink)>,
|
||||
|
|
@ -889,16 +900,16 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
Scatter)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, TfIdfVectorizer)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, bool,
|
||||
NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float,
|
||||
NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t,
|
||||
NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t,
|
||||
NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, uint8_t,
|
||||
NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12,
|
||||
bool, NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12,
|
||||
float, NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12,
|
||||
int32_t, NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12,
|
||||
int64_t, NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 12,
|
||||
uint8_t, NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, string,
|
||||
Where)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float,
|
||||
|
|
@ -1042,7 +1053,7 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
BitShift)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint64_t,
|
||||
BitShift)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Pad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 12, Pad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 11,
|
||||
GatherND)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Range)>,
|
||||
|
|
@ -1164,7 +1175,8 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
ReduceMin)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, uint8_t,
|
||||
ReduceMin)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, GatherND)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, 12,
|
||||
GatherND)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, Einsum)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, Unsqueeze)>,
|
||||
|
|
@ -1348,6 +1360,23 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, SpaceToDepth)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, ScatterElements)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, ScatterND)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, Identity)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, float,
|
||||
IsNaN)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, MLFloat16,
|
||||
IsNaN)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, bool,
|
||||
NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, float,
|
||||
NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, int32_t,
|
||||
NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, int64_t,
|
||||
NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, uint8_t,
|
||||
NonZero)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, GatherND)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, Pad)>,
|
||||
};
|
||||
|
||||
for (auto& function_table_entry : function_table) {
|
||||
|
|
|
|||
|
|
@ -37,9 +37,18 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
GatherND);
|
||||
|
||||
// opset 12 added batch_dims attribute
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
||||
GatherND,
|
||||
12, 12,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T", DataTypeImpl::AllTensorTypes())
|
||||
.TypeConstraint("Tind", DataTypeImpl::GetTensorType<int64_t>()),
|
||||
GatherND);
|
||||
|
||||
// spec added BFloat16
|
||||
ONNX_CPU_OPERATOR_KERNEL(
|
||||
GatherND,
|
||||
12,
|
||||
13,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T", DataTypeImpl::AllTensorTypes())
|
||||
.TypeConstraint("Tind", DataTypeImpl::GetTensorType<int64_t>()),
|
||||
|
|
@ -93,12 +102,14 @@ Status GatherNDBase::PrepareForCompute(const TensorShape& input_shape, const Ten
|
|||
|
||||
p.slice_offsets[slice_idx] = input_base_offset + relative_slice_offset;
|
||||
};
|
||||
concurrency::ThreadPool::TryParallelFor(tp, num_slices, static_cast<double>(num_slice_dims),
|
||||
[&lambda](ptrdiff_t first, ptrdiff_t last) {
|
||||
for (int slice_idx = static_cast<int>(first), end = static_cast<int>(last); slice_idx < end; ++slice_idx) {
|
||||
lambda(slice_idx);
|
||||
}
|
||||
});
|
||||
|
||||
concurrency::ThreadPool::TryParallelFor(
|
||||
tp, num_slices, static_cast<double>(num_slice_dims),
|
||||
[&lambda](ptrdiff_t first, ptrdiff_t last) {
|
||||
for (int slice_idx = static_cast<int>(first), end = static_cast<int>(last); slice_idx < end; ++slice_idx) {
|
||||
lambda(slice_idx);
|
||||
}
|
||||
});
|
||||
|
||||
return err_index == 0 ? Status::OK()
|
||||
: ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "invalid index found, index = ", err_index);
|
||||
|
|
@ -119,7 +130,8 @@ Status GatherND::Compute(OpKernelContext* context) const {
|
|||
const auto* input_tensor = context->Input<Tensor>(0);
|
||||
const auto* indices_tensor = context->Input<Tensor>(1);
|
||||
|
||||
ORT_ENFORCE(input_tensor != nullptr && indices_tensor != nullptr, "GatherNDBase PrepareForCompute: Input count mismatch");
|
||||
ORT_ENFORCE(input_tensor != nullptr && indices_tensor != nullptr,
|
||||
"GatherNDBase PrepareForCompute: Input count mismatch");
|
||||
|
||||
const auto& input_shape = input_tensor->Shape();
|
||||
const auto& indices_shape = indices_tensor->Shape();
|
||||
|
|
@ -164,12 +176,13 @@ Status GatherND::GatherNumber(const Prepare& p, concurrency::ThreadPool* tp) con
|
|||
memcpy(p.output_base + slice_idx * p.bytes_per_slice, p.input_base + p.slice_offsets[slice_idx] * p.element_bytes,
|
||||
p.bytes_per_slice);
|
||||
};
|
||||
concurrency::ThreadPool::TryParallelFor(tp, p.slice_offsets.size(), static_cast<double>(p.bytes_per_slice),
|
||||
[&lambda](ptrdiff_t first, ptrdiff_t last) {
|
||||
for (int slice_idx = static_cast<int>(first), end = static_cast<int>(last); slice_idx < end; ++slice_idx) {
|
||||
lambda(slice_idx);
|
||||
}
|
||||
});
|
||||
concurrency::ThreadPool::TryParallelFor(
|
||||
tp, p.slice_offsets.size(), static_cast<double>(p.bytes_per_slice),
|
||||
[&lambda](ptrdiff_t first, ptrdiff_t last) {
|
||||
for (int slice_idx = static_cast<int>(first), end = static_cast<int>(last); slice_idx < end; ++slice_idx) {
|
||||
lambda(slice_idx);
|
||||
}
|
||||
});
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
|
|
@ -180,12 +193,13 @@ Status GatherND::GatherString(const Prepare& p, concurrency::ThreadPool* tp) con
|
|||
p.output_str_base[slice_base_offset + j] = p.input_str_base[p.slice_offsets[slice_idx] + j];
|
||||
}
|
||||
};
|
||||
concurrency::ThreadPool::TryParallelFor(tp, p.slice_offsets.size(), static_cast<double>(p.element_count_per_slice),
|
||||
[&lambda](ptrdiff_t first, ptrdiff_t last) {
|
||||
for (int slice_idx = static_cast<int>(first), end = static_cast<int>(last); slice_idx < end; ++slice_idx) {
|
||||
lambda(slice_idx);
|
||||
}
|
||||
});
|
||||
concurrency::ThreadPool::TryParallelFor(
|
||||
tp, p.slice_offsets.size(), static_cast<double>(p.element_count_per_slice),
|
||||
[&lambda](ptrdiff_t first, ptrdiff_t last) {
|
||||
for (int slice_idx = static_cast<int>(first), end = static_cast<int>(last); slice_idx < end; ++slice_idx) {
|
||||
lambda(slice_idx);
|
||||
}
|
||||
});
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -24,9 +24,16 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<bool>()),
|
||||
IdentityOp<true>);
|
||||
|
||||
ONNX_CPU_OPERATOR_KERNEL(
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
||||
Identity,
|
||||
1,
|
||||
12,
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorTypes()).Alias(0, 0),
|
||||
IdentityOp<false>);
|
||||
|
||||
ONNX_CPU_OPERATOR_KERNEL(
|
||||
Identity,
|
||||
13,
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorTypes()).Alias(0, 0),
|
||||
IdentityOp<false>);
|
||||
|
||||
|
|
|
|||
|
|
@ -9,16 +9,28 @@
|
|||
|
||||
namespace onnxruntime {
|
||||
// https://github.com/onnx/onnx/blob/master/docs/Operators.md#IsNaN
|
||||
#define ADD_TYPED_ISNAN_OP(data_type) \
|
||||
ONNX_CPU_OPERATOR_TYPED_KERNEL( \
|
||||
#define ADD_TYPED_ISNAN_OP_9(data_type) \
|
||||
ONNX_CPU_OPERATOR_VERSIONED_TYPED_KERNEL( \
|
||||
IsNaN, \
|
||||
9, \
|
||||
9, 12, \
|
||||
data_type, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<data_type>()) \
|
||||
.TypeConstraint("T2", DataTypeImpl::GetTensorType<bool>()), \
|
||||
IsNaN<data_type>);
|
||||
|
||||
#define ADD_TYPED_ISNAN_OP(data_type) \
|
||||
ONNX_CPU_OPERATOR_TYPED_KERNEL( \
|
||||
IsNaN, \
|
||||
13, \
|
||||
data_type, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<data_type>()) \
|
||||
.TypeConstraint("T2", DataTypeImpl::GetTensorType<bool>()), \
|
||||
IsNaN<data_type>);
|
||||
|
||||
ADD_TYPED_ISNAN_OP_9(float);
|
||||
ADD_TYPED_ISNAN_OP_9(MLFloat16);
|
||||
ADD_TYPED_ISNAN_OP(float);
|
||||
ADD_TYPED_ISNAN_OP(MLFloat16);
|
||||
|
||||
|
|
@ -48,7 +60,10 @@ Status IsNaN<MLFloat16>::Compute(OpKernelContext* context) const {
|
|||
auto shape_size = dims.Size();
|
||||
auto& Y = *context->Output(0, dims);
|
||||
|
||||
EigenMap<bool>(Y) = ConstEigenVectorMap<Eigen::half>(static_cast<const Eigen::half*>(static_cast<const void*>(X_data)), shape_size).array().isNaN();
|
||||
EigenMap<bool>(Y) =
|
||||
ConstEigenVectorMap<Eigen::half>(static_cast<const Eigen::half*>(static_cast<const void*>(X_data)), shape_size)
|
||||
.array()
|
||||
.isNaN();
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -10,16 +10,27 @@
|
|||
|
||||
namespace onnxruntime {
|
||||
// kernel builder functions
|
||||
#define NONZERO_TYPED_KERNEL_WITH_TYPE_NAME(type, type_name) \
|
||||
ONNX_CPU_OPERATOR_TYPED_KERNEL( \
|
||||
#define NONZERO_9_TYPED_KERNEL(type) \
|
||||
ONNX_CPU_OPERATOR_VERSIONED_TYPED_KERNEL( \
|
||||
NonZero, \
|
||||
9, \
|
||||
type_name, \
|
||||
9, 12, \
|
||||
type, \
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<type>()), \
|
||||
NonZero<type>)
|
||||
|
||||
#define NONZERO_TYPED_KERNEL(type) \
|
||||
NONZERO_TYPED_KERNEL_WITH_TYPE_NAME(type, type)
|
||||
#define NONZERO_TYPED_KERNEL(type) \
|
||||
ONNX_CPU_OPERATOR_TYPED_KERNEL( \
|
||||
NonZero, \
|
||||
13, \
|
||||
type, \
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<type>()), \
|
||||
NonZero<type>)
|
||||
|
||||
NONZERO_9_TYPED_KERNEL(bool)
|
||||
NONZERO_9_TYPED_KERNEL(uint8_t)
|
||||
NONZERO_9_TYPED_KERNEL(int32_t)
|
||||
NONZERO_9_TYPED_KERNEL(int64_t)
|
||||
NONZERO_9_TYPED_KERNEL(float)
|
||||
|
||||
// start with a subset of types, enable more as needed...
|
||||
NONZERO_TYPED_KERNEL(bool)
|
||||
|
|
@ -37,7 +48,7 @@ NONZERO_TYPED_KERNEL(float)
|
|||
//NONZERO_TYPED_KERNEL(double)
|
||||
//NONZERO_TYPED_KERNEL_WITH_TYPE_NAME(std::string, string)
|
||||
|
||||
#undef NONZERO_TYPED_KERNEL_WITH_TYPE_NAME
|
||||
#undef NONZERO_9_TYPED_KERNEL
|
||||
#undef NONZERO_TYPED_KERNEL
|
||||
|
||||
template <typename T>
|
||||
|
|
|
|||
|
|
@ -45,9 +45,22 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
// 'pads' and 'value' (attributes previously) became inputs in this version
|
||||
// The core logic remains the same
|
||||
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
||||
Pad,
|
||||
11, 12,
|
||||
KernelDefBuilder().TypeConstraint("T", {DataTypeImpl::GetTensorType<float>(),
|
||||
DataTypeImpl::GetTensorType<double>(),
|
||||
DataTypeImpl::GetTensorType<int32_t>(),
|
||||
DataTypeImpl::GetTensorType<int64_t>(),
|
||||
DataTypeImpl::GetTensorType<uint32_t>(),
|
||||
DataTypeImpl::GetTensorType<uint64_t>(),
|
||||
DataTypeImpl::GetTensorType<int8_t>(),
|
||||
DataTypeImpl::GetTensorType<uint8_t>()}),
|
||||
Pad);
|
||||
|
||||
ONNX_CPU_OPERATOR_KERNEL(
|
||||
Pad,
|
||||
11,
|
||||
13,
|
||||
KernelDefBuilder().TypeConstraint("T", {DataTypeImpl::GetTensorType<float>(),
|
||||
DataTypeImpl::GetTensorType<double>(),
|
||||
DataTypeImpl::GetTensorType<int32_t>(),
|
||||
|
|
@ -60,7 +73,8 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
|
||||
// This is the general padding method to n-dimensionally do edge or reflection padding (based on the inputDelta values)
|
||||
template <typename T>
|
||||
static void PadAxis(T* output, T* input, ptrdiff_t input_delta, ptrdiff_t input_pitch, size_t block_size, size_t block_count) {
|
||||
static void PadAxis(T* output, T* input, ptrdiff_t input_delta, ptrdiff_t input_pitch,
|
||||
size_t block_size, size_t block_count) {
|
||||
for (size_t block_index = 0; block_index < block_count; block_index++) {
|
||||
for (size_t i = 0; i < block_size; i++) {
|
||||
*output++ = *input;
|
||||
|
|
@ -159,7 +173,8 @@ static void FlattenInnerShape(const std::vector<int64_t>& input_dims, const std:
|
|||
break;
|
||||
|
||||
// Break on first Axis that has padding
|
||||
if (!(pads[inner_axis] == 0 && pads[inner_axis + dims_count] == 0 && slices[inner_axis] == 0 && slices[inner_axis + dims_count] == 0))
|
||||
if (!(pads[inner_axis] == 0 && pads[inner_axis + dims_count] == 0 &&
|
||||
slices[inner_axis] == 0 && slices[inner_axis + dims_count] == 0))
|
||||
break;
|
||||
|
||||
} while (inner_axis-- > 0);
|
||||
|
|
@ -175,7 +190,8 @@ static void ReshapePads(const std::vector<int64_t>& src_pad, size_t src_dim_coun
|
|||
size_t inner_no_pad_size, std::vector<int64_t>& reshaped_pad) {
|
||||
size_t inner_axis = new_dim_count - 1;
|
||||
std::copy(src_pad.begin(), src_pad.begin() + inner_axis, reshaped_pad.begin());
|
||||
std::copy(src_pad.begin() + src_dim_count, src_pad.begin() + src_dim_count + inner_axis, reshaped_pad.begin() + new_dim_count);
|
||||
std::copy(src_pad.begin() + src_dim_count, src_pad.begin() + src_dim_count + inner_axis,
|
||||
reshaped_pad.begin() + new_dim_count);
|
||||
|
||||
// Flatten inner axis.
|
||||
reshaped_pad[inner_axis] = src_pad[inner_axis] * inner_no_pad_size;
|
||||
|
|
@ -184,10 +200,10 @@ static void ReshapePads(const std::vector<int64_t>& src_pad, size_t src_dim_coun
|
|||
|
||||
template <typename T>
|
||||
static Status PadImpl(OpKernelContext* ctx,
|
||||
const std::vector<int64_t>& pads,
|
||||
const std::vector<int64_t>& slices,
|
||||
const Mode& mode,
|
||||
T value) {
|
||||
const std::vector<int64_t>& pads,
|
||||
const std::vector<int64_t>& slices,
|
||||
const Mode& mode,
|
||||
T value) {
|
||||
const auto& input_tensor = *ctx->Input<Tensor>(0);
|
||||
const auto& orig_input_shape = input_tensor.Shape();
|
||||
std::vector<int64_t> output_dims(orig_input_shape.GetDims());
|
||||
|
|
@ -204,7 +220,9 @@ static Status PadImpl(OpKernelContext* ctx,
|
|||
// Reshape padding
|
||||
size_t new_dims_count = reshaped_input_dims.size();
|
||||
size_t inner_axis = new_dims_count - 1;
|
||||
size_t inner_no_pad_size = output_dims[inner_axis] > 0 ? reshaped_input_dims[inner_axis] / output_dims[inner_axis] : 0;
|
||||
size_t inner_no_pad_size = output_dims[inner_axis] > 0
|
||||
? reshaped_input_dims[inner_axis] / output_dims[inner_axis]
|
||||
: 0;
|
||||
std::vector<int64_t> reshaped_pad(2 * new_dims_count), reshaped_slice(2 * new_dims_count);
|
||||
ReshapePads(pads, data_rank, new_dims_count, inner_no_pad_size, reshaped_pad);
|
||||
ReshapePads(slices, data_rank, new_dims_count, inner_no_pad_size, reshaped_slice);
|
||||
|
|
@ -219,7 +237,8 @@ static Status PadImpl(OpKernelContext* ctx,
|
|||
for (size_t i = 0; i < new_dims_count; i++) {
|
||||
input_starts.push_back(-1 * reshaped_slice[i]);
|
||||
input_extents.push_back(reshaped_input_dims[i] + reshaped_slice[i] + reshaped_slice[i + new_dims_count]);
|
||||
reshaped_output_dims[i] += reshaped_pad[i] + reshaped_pad[i + new_dims_count] + reshaped_slice[i] + reshaped_slice[i + new_dims_count];
|
||||
reshaped_output_dims[i] += reshaped_pad[i] + reshaped_pad[i + new_dims_count] +
|
||||
reshaped_slice[i] + reshaped_slice[i + new_dims_count];
|
||||
}
|
||||
|
||||
for (size_t i = 0; i < data_rank; i++) {
|
||||
|
|
@ -335,7 +354,8 @@ static Status PadImpl(OpKernelContext* ctx,
|
|||
T* axisStart = output - inner_pitch * input_extents[input_counters.Axis()];
|
||||
int64_t prePad = reshaped_pad[input_counters.Axis()];
|
||||
int64_t postPad = reshaped_pad[input_counters.Axis() + new_dims_count];
|
||||
PadAxis(axisStart - prePad * inner_pitch, axisStart + prePad * inner_pitch, 1, -inner_pitch * 2, inner_pitch, prePad);
|
||||
PadAxis(axisStart - prePad * inner_pitch, axisStart + prePad * inner_pitch, 1, -inner_pitch * 2,
|
||||
inner_pitch, prePad);
|
||||
PadAxis(output, output - 2 * inner_pitch, 1, -inner_pitch * 2, inner_pitch, postPad);
|
||||
output += inner_pitch * postPad;
|
||||
alignSkip += inner_pitch * prePad;
|
||||
|
|
@ -347,24 +367,21 @@ static Status PadImpl(OpKernelContext* ctx,
|
|||
return Status::OK();
|
||||
}
|
||||
|
||||
union PadValue
|
||||
{
|
||||
union PadValue {
|
||||
uint64_t u64;
|
||||
uint32_t u32;
|
||||
uint8_t u8;
|
||||
double f64;
|
||||
float f32;
|
||||
uint8_t u8;
|
||||
double f64;
|
||||
float f32;
|
||||
};
|
||||
|
||||
static PadValue PadValueFromFloat(float value, MLDataType data_type) {
|
||||
PadValue result;
|
||||
if (data_type == DataTypeImpl::GetType<float>()) {
|
||||
result.f32 = value;
|
||||
}
|
||||
else if (data_type == DataTypeImpl::GetType<double>()) {
|
||||
} else if (data_type == DataTypeImpl::GetType<double>()) {
|
||||
result.f64 = value;
|
||||
}
|
||||
else {
|
||||
} else {
|
||||
ORT_THROW("Unsupported input data type of ", data_type);
|
||||
}
|
||||
return result;
|
||||
|
|
@ -390,7 +407,8 @@ Status Pad::Compute(OpKernelContext* ctx) const {
|
|||
ORT_ENFORCE(pads_tensor.IsDataType<int64_t>(),
|
||||
"Pads tensor should be an INT64 tensor");
|
||||
ORT_ENFORCE(pads_tensor_dims.size() == 1 || (pads_tensor_dims.size() == 2 && pads_tensor_dims[0] == 1),
|
||||
"Pads tensor should be a 1D tensor of shape [2 * input_rank] or a 2D tensor of shape [1, 2 * input_rank]");
|
||||
"Pads tensor should be a 1D tensor of shape [2 * input_rank] "
|
||||
"or a 2D tensor of shape [1, 2 * input_rank]");
|
||||
|
||||
const int64_t* pads_tensor_raw_data = pads_tensor.template Data<int64_t>();
|
||||
size_t pads_size = static_cast<size_t>(pads_tensor.Shape().Size());
|
||||
|
|
|
|||
Loading…
Reference in a new issue