diff --git a/onnxruntime/core/providers/cpu/sequence/sequence_ops.cc b/onnxruntime/core/providers/cpu/sequence/sequence_ops.cc index e9e15b8b99..2ee3908a8b 100644 --- a/onnxruntime/core/providers/cpu/sequence/sequence_ops.cc +++ b/onnxruntime/core/providers/cpu/sequence/sequence_ops.cc @@ -334,6 +334,7 @@ ONNX_CPU_OPERATOR_KERNEL( DataTypeImpl::GetTensorType(), DataTypeImpl::GetTensorType(), DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), DataTypeImpl::GetTensorType()}) .TypeConstraint("S", DataTypeImpl::AllSequenceTensorTypes()) .TypeConstraint("I", std::vector{ @@ -358,6 +359,8 @@ Status SplitToSequence::Compute(OpKernelContext* context) const { status = ComputeImpl(*context, input, p_split_input); else if (input.IsDataType()) status = ComputeImpl(*context, input, p_split_input); + else if (input.IsDataType()) + status = ComputeImpl(*context, input, p_split_input); else if (input.IsDataTypeString()) status = ComputeImpl(*context, input, p_split_input); else diff --git a/onnxruntime/test/providers/cpu/sequence/sequence_ops_test.cc b/onnxruntime/test/providers/cpu/sequence/sequence_ops_test.cc index d8b906a1e9..d46c0626e0 100644 --- a/onnxruntime/test/providers/cpu/sequence/sequence_ops_test.cc +++ b/onnxruntime/test/providers/cpu/sequence/sequence_ops_test.cc @@ -317,6 +317,17 @@ TEST(SequenceOpsTest, SplitToSequence_DefaultAxis0EqualSplitFloat) { test.Run(); } +TEST(SequenceOpsTest, SplitToSequence_DefaultAxis0EqualSplitLong) { + OpTester test("SplitToSequence", 11); + test.AddInput("input", {4, 2}, GetConsequtiveVector(1, 8)); + test.AddInput("split", {1, 2}, {2, 2}); + SeqTensors output; + output.AddTensor({2, 2}, {1, 2, 3, 4}); + output.AddTensor({2, 2}, {5, 6, 7, 8}); + test.AddSeqOutput("S2", output); + test.Run(); +} + TEST(SequenceOpsTest, SplitToSequence_DefaultAxis0EqualSplitFloatScalarSplit) { OpTester test("SplitToSequence", 11); test.AddInput("input", {4, 2}, GetConsequtiveVector(1.f, 8));