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();
}