From ad94a1dd6de37b3c3db8ba9a377ff7f3549520ea Mon Sep 17 00:00:00 2001 From: Scott McKay Date: Sat, 17 Oct 2020 02:39:03 +1000 Subject: [PATCH] Add opset 13 registrations for Identity, IsNaN, NonZero, GatherND and Pad (#5513) --- .../providers/cpu/cpu_execution_provider.cc | 85 +++++++++++++------ .../core/providers/cpu/tensor/gather_nd.cc | 54 +++++++----- .../core/providers/cpu/tensor/identity_op.cc | 9 +- .../core/providers/cpu/tensor/isnan.cc | 23 ++++- .../core/providers/cpu/tensor/nonzero_op.cc | 25 ++++-- onnxruntime/core/providers/cpu/tensor/pad.cc | 60 ++++++++----- 6 files changed, 175 insertions(+), 81 deletions(-) diff --git a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc index 1f75addd6f..873110b3d4 100644 --- a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc +++ b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc @@ -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 KernelCreateInfo BuildKernelCreateInfo() { @@ -777,7 +787,8 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { Gather)>, BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -847,10 +858,10 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -889,16 +900,16 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -1164,7 +1175,8 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { ReduceMin)>, BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -1348,6 +1360,23 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, }; for (auto& function_table_entry : function_table) { diff --git a/onnxruntime/core/providers/cpu/tensor/gather_nd.cc b/onnxruntime/core/providers/cpu/tensor/gather_nd.cc index 40004d874a..f3592357e0 100644 --- a/onnxruntime/core/providers/cpu/tensor/gather_nd.cc +++ b/onnxruntime/core/providers/cpu/tensor/gather_nd.cc @@ -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()), + GatherND); + +// spec added BFloat16 ONNX_CPU_OPERATOR_KERNEL( GatherND, - 12, + 13, KernelDefBuilder() .TypeConstraint("T", DataTypeImpl::AllTensorTypes()) .TypeConstraint("Tind", DataTypeImpl::GetTensorType()), @@ -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(num_slice_dims), - [&lambda](ptrdiff_t first, ptrdiff_t last) { - for (int slice_idx = static_cast(first), end = static_cast(last); slice_idx < end; ++slice_idx) { - lambda(slice_idx); - } - }); + + concurrency::ThreadPool::TryParallelFor( + tp, num_slices, static_cast(num_slice_dims), + [&lambda](ptrdiff_t first, ptrdiff_t last) { + for (int slice_idx = static_cast(first), end = static_cast(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(0); const auto* indices_tensor = context->Input(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(p.bytes_per_slice), - [&lambda](ptrdiff_t first, ptrdiff_t last) { - for (int slice_idx = static_cast(first), end = static_cast(last); slice_idx < end; ++slice_idx) { - lambda(slice_idx); - } - }); + concurrency::ThreadPool::TryParallelFor( + tp, p.slice_offsets.size(), static_cast(p.bytes_per_slice), + [&lambda](ptrdiff_t first, ptrdiff_t last) { + for (int slice_idx = static_cast(first), end = static_cast(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(p.element_count_per_slice), - [&lambda](ptrdiff_t first, ptrdiff_t last) { - for (int slice_idx = static_cast(first), end = static_cast(last); slice_idx < end; ++slice_idx) { - lambda(slice_idx); - } - }); + concurrency::ThreadPool::TryParallelFor( + tp, p.slice_offsets.size(), static_cast(p.element_count_per_slice), + [&lambda](ptrdiff_t first, ptrdiff_t last) { + for (int slice_idx = static_cast(first), end = static_cast(last); slice_idx < end; ++slice_idx) { + lambda(slice_idx); + } + }); return Status::OK(); } diff --git a/onnxruntime/core/providers/cpu/tensor/identity_op.cc b/onnxruntime/core/providers/cpu/tensor/identity_op.cc index 3e00809089..6673793c13 100644 --- a/onnxruntime/core/providers/cpu/tensor/identity_op.cc +++ b/onnxruntime/core/providers/cpu/tensor/identity_op.cc @@ -24,9 +24,16 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL( .TypeConstraint("T1", DataTypeImpl::GetTensorType()), IdentityOp); -ONNX_CPU_OPERATOR_KERNEL( +ONNX_CPU_OPERATOR_VERSIONED_KERNEL( Identity, 1, + 12, + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorTypes()).Alias(0, 0), + IdentityOp); + +ONNX_CPU_OPERATOR_KERNEL( + Identity, + 13, KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorTypes()).Alias(0, 0), IdentityOp); diff --git a/onnxruntime/core/providers/cpu/tensor/isnan.cc b/onnxruntime/core/providers/cpu/tensor/isnan.cc index 19f24df89a..a4ac251e54 100644 --- a/onnxruntime/core/providers/cpu/tensor/isnan.cc +++ b/onnxruntime/core/providers/cpu/tensor/isnan.cc @@ -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()) \ .TypeConstraint("T2", DataTypeImpl::GetTensorType()), \ IsNaN); +#define ADD_TYPED_ISNAN_OP(data_type) \ + ONNX_CPU_OPERATOR_TYPED_KERNEL( \ + IsNaN, \ + 13, \ + data_type, \ + KernelDefBuilder() \ + .TypeConstraint("T1", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("T2", DataTypeImpl::GetTensorType()), \ + IsNaN); + +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::Compute(OpKernelContext* context) const { auto shape_size = dims.Size(); auto& Y = *context->Output(0, dims); - EigenMap(Y) = ConstEigenVectorMap(static_cast(static_cast(X_data)), shape_size).array().isNaN(); + EigenMap(Y) = + ConstEigenVectorMap(static_cast(static_cast(X_data)), shape_size) + .array() + .isNaN(); return Status::OK(); } diff --git a/onnxruntime/core/providers/cpu/tensor/nonzero_op.cc b/onnxruntime/core/providers/cpu/tensor/nonzero_op.cc index c23a897c03..0b09f213b6 100644 --- a/onnxruntime/core/providers/cpu/tensor/nonzero_op.cc +++ b/onnxruntime/core/providers/cpu/tensor/nonzero_op.cc @@ -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()), \ NonZero) -#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()), \ + NonZero) + +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 diff --git a/onnxruntime/core/providers/cpu/tensor/pad.cc b/onnxruntime/core/providers/cpu/tensor/pad.cc index 2e082d5210..51e1ebfe27 100644 --- a/onnxruntime/core/providers/cpu/tensor/pad.cc +++ b/onnxruntime/core/providers/cpu/tensor/pad.cc @@ -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(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType()}), + Pad); + ONNX_CPU_OPERATOR_KERNEL( Pad, - 11, + 13, KernelDefBuilder().TypeConstraint("T", {DataTypeImpl::GetTensorType(), DataTypeImpl::GetTensorType(), DataTypeImpl::GetTensorType(), @@ -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 -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& 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& src_pad, size_t src_dim_coun size_t inner_no_pad_size, std::vector& 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& src_pad, size_t src_dim_coun template static Status PadImpl(OpKernelContext* ctx, - const std::vector& pads, - const std::vector& slices, - const Mode& mode, - T value) { + const std::vector& pads, + const std::vector& slices, + const Mode& mode, + T value) { const auto& input_tensor = *ctx->Input(0); const auto& orig_input_shape = input_tensor.Shape(); std::vector 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 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()) { result.f32 = value; - } - else if (data_type == DataTypeImpl::GetType()) { + } else if (data_type == DataTypeImpl::GetType()) { 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(), "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(); size_t pads_size = static_cast(pads_tensor.Shape().Size());