mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
String Tensor SplitToSequence fix (#19942)
This commit is contained in:
parent
0af5eacc8b
commit
19ff4a6d6c
2 changed files with 14 additions and 1 deletions
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
|
|
|
|||
Loading…
Reference in a new issue