String Tensor SplitToSequence fix (#19942)

This commit is contained in:
Adam Pocock 2024-03-20 13:52:00 -04:00 committed by GitHub
parent 0af5eacc8b
commit 19ff4a6d6c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 14 additions and 1 deletions

View file

@ -453,7 +453,7 @@ Status SplitToSequence::ComputeImpl(OpKernelContext& context, const Tensor& inpu
int num_remaining_splits = 0;
InlinedVector<int64_t> split_sizes;
const bool is_string_type = input.IsDataTypeString();
const size_t element_size = (is_string_type) ? 0U : input.DataType()->Size();
const size_t element_size = input.DataType()->Size();
// figure out split_scalar or split_sizes
if (p_split_input) {

View file

@ -442,6 +442,19 @@ TEST(SequenceOpsTest, SplitToSequence_PositiveAxisScalarSplit) {
test.Run();
}
TEST(SequenceOpsTest, SplitToSequence_StringSplit) {
OpTester test("SplitToSequence", 11);
test.AddInput<std::string>("input", {3}, std::vector<std::string>({"Test string", "Another string", "A third and much longer string"}));
int64_t axis = 0;
test.AddAttribute("axis", axis);
SeqTensors<std::string> output;
output.AddTensor({1}, {"Test string"});
output.AddTensor({1}, {"Another string"});
output.AddTensor({1}, {"A third and much longer string"});
test.AddSeqOutput("S2", output);
test.Run();
}
TEST(SequenceOpsTest, SplitToSequence_DefaultAxis0UnevenSplitFloat) {
OpTester test("SplitToSequence", 11);
test.AddInput<float>("input", {5, 2}, GetConsecutiveVector<float>(1.f, 10));