diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md index 20e0fcada0..cc78a81872 100644 --- a/docs/OperatorKernels.md +++ b/docs/OperatorKernels.md @@ -303,9 +303,9 @@ Do not modify directly.* |Softsign|*in* input:**T**
*out* output:**T**|1+|**T** = tensor(float)| |SpaceToDepth|*in* input:**T**
*out* output:**T**|13+|**T** = tensor(double), tensor(float)| |||[1, 12]|**T** = tensor(double), tensor(float)| -|Split|*in* input:**T**
*in* split:**T**
*out* outputs...:**T**

or

*in* input:**T**
*in* split:**tensor(int64)**
*out* outputs:**T**

or

*in* input:**T**
*out* outputs:**T**|13+|**T** = tensor(float), tensor(int32), tensor(int64), tensor(string), tensor(uint8)| -|||[11, 12]|**T** = tensor(float), tensor(int32), tensor(int64), tensor(string), tensor(uint8)| -|||[2, 10]|**T** = tensor(float), tensor(int32), tensor(int64), tensor(string), tensor(uint8)| +|Split|*in* input:**T**
*in* split:**T**
*out* outputs...:**T**

or

*in* input:**T**
*in* split:**tensor(int64)**
*out* outputs:**T**

or

*in* input:**T**
*out* outputs:**T**|13+|**T** = tensor(float), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint8)| +|||[11, 12]|**T** = tensor(float), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint8)| +|||[2, 10]|**T** = tensor(float), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint8)| |SplitToSequence|*in* input:**T**
*in* split:**I**
*out* output_sequence:**S**|11+|**I** = tensor(int32), tensor(int64)
**S** = seq(tensor(bfloat16)), seq(tensor(bool)), seq(tensor(double)), seq(tensor(float)), seq(tensor(float16)), seq(tensor(int16)), seq(tensor(int32)), seq(tensor(int64)), seq(tensor(int8)), seq(tensor(string)), seq(tensor(uint16)), seq(tensor(uint32)), seq(tensor(uint64)), seq(tensor(uint8))
**T** = tensor(double), tensor(float), tensor(int32), tensor(int64), tensor(string)| |Sqrt|*in* X:**T**
*out* Y:**T**|13+|**T** = tensor(double), tensor(float)| |||[6, 12]|**T** = tensor(double), tensor(float)| diff --git a/onnxruntime/core/providers/cpu/tensor/split.cc b/onnxruntime/core/providers/cpu/tensor/split.cc index 08df9dd012..1666efd6a1 100644 --- a/onnxruntime/core/providers/cpu/tensor/split.cc +++ b/onnxruntime/core/providers/cpu/tensor/split.cc @@ -16,7 +16,7 @@ namespace onnxruntime { namespace op_kernel_type_control { ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPES_ALL_OPSETS( kCpuExecutionProvider, kOnnxDomain, Split, Input, 0, - float, int32_t, int64_t, uint8_t, std::string); + float, int8_t, int32_t, int64_t, uint8_t, std::string); ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPES_ALL_OPSETS( kCpuExecutionProvider, kOnnxDomain, Split, Input, 0, int32_t, int64_t); @@ -27,13 +27,18 @@ using SplitDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS( using EnabledSplitDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS( kCpuExecutionProvider, kOnnxDomain, Split, Input, 0); +using OldSplitDataTypes = onnxruntime::TypeList; + ONNX_CPU_OPERATOR_VERSIONED_KERNEL( Split, 2, 10, KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraintsFromTypeList(), - BuildKernelDefConstraintsFromTypeList()), + BuildKernelDefConstraintsFromTypeList()) + .FixedTypeConstraintForHash( + "T", + BuildKernelDefConstraintsFromTypeList()), Split); // Opset 11 starts to support Neg Axis. @@ -43,7 +48,10 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL( 12, KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraintsFromTypeList(), - BuildKernelDefConstraintsFromTypeList()), + BuildKernelDefConstraintsFromTypeList()) + .FixedTypeConstraintForHash( + "T", + BuildKernelDefConstraintsFromTypeList()), Split); // Opset 13 starts to supports 'split' as optional input. @@ -52,7 +60,10 @@ ONNX_CPU_OPERATOR_KERNEL( 13, KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraintsFromTypeList(), - BuildKernelDefConstraintsFromTypeList()), + BuildKernelDefConstraintsFromTypeList()) + .FixedTypeConstraintForHash( + "T", + BuildKernelDefConstraintsFromTypeList()), Split); Status SplitBase::PrepareForCompute(const TensorShape& input_shape, int num_outputs, int64_t& axis, int& before_dims, @@ -108,6 +119,8 @@ Status Split::Compute(OpKernelContext* context) const { status = ComputeImpl(*context, input); else if (input.IsDataType()) status = ComputeImpl(*context, input); + else if (input.IsDataType()) + status = ComputeImpl(*context, input); else if (input.IsDataTypeString()) status = ComputeImpl(*context, input); else diff --git a/onnxruntime/test/providers/cpu/tensor/split_op_test.cc b/onnxruntime/test/providers/cpu/tensor/split_op_test.cc index f5967f29b8..b2dc3a44b3 100644 --- a/onnxruntime/test/providers/cpu/tensor/split_op_test.cc +++ b/onnxruntime/test/providers/cpu/tensor/split_op_test.cc @@ -94,6 +94,10 @@ static void SplitTestInt() { RunTest(axis, {}, input, outputs, false); //TensorRT parser: Assertion failed: axis != BATCH_DIM } +TEST(SplitOperatorTest, Axis0EqualSplitInt8) { + SplitTestInt(); +} + TEST(SplitOperatorTest, Axis0EqualSplitInt32) { SplitTestInt(); }