Add long type support for SplitToSequence operator (#5367)

This commit is contained in:
Bowen Bao 2020-10-13 12:57:11 -07:00 committed by GitHub
parent e01d152464
commit 8e9afe1944
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 14 additions and 0 deletions

View file

@ -334,6 +334,7 @@ ONNX_CPU_OPERATOR_KERNEL(
DataTypeImpl::GetTensorType<float>(),
DataTypeImpl::GetTensorType<double>(),
DataTypeImpl::GetTensorType<int32_t>(),
DataTypeImpl::GetTensorType<int64_t>(),
DataTypeImpl::GetTensorType<std::string>()})
.TypeConstraint("S", DataTypeImpl::AllSequenceTensorTypes())
.TypeConstraint("I", std::vector<MLDataType>{
@ -358,6 +359,8 @@ Status SplitToSequence::Compute(OpKernelContext* context) const {
status = ComputeImpl<double>(*context, input, p_split_input);
else if (input.IsDataType<int32_t>())
status = ComputeImpl<int32_t>(*context, input, p_split_input);
else if (input.IsDataType<int64_t>())
status = ComputeImpl<int64_t>(*context, input, p_split_input);
else if (input.IsDataTypeString())
status = ComputeImpl<std::string>(*context, input, p_split_input);
else

View file

@ -317,6 +317,17 @@ TEST(SequenceOpsTest, SplitToSequence_DefaultAxis0EqualSplitFloat) {
test.Run();
}
TEST(SequenceOpsTest, SplitToSequence_DefaultAxis0EqualSplitLong) {
OpTester test("SplitToSequence", 11);
test.AddInput<int64_t>("input", {4, 2}, GetConsequtiveVector<int64_t>(1, 8));
test.AddInput<int64_t>("split", {1, 2}, {2, 2});
SeqTensors<int64_t> 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<float>("input", {4, 2}, GetConsequtiveVector<float>(1.f, 8));