From c63ae53d7f45cea93b64f45c70d2b8039b6d0552 Mon Sep 17 00:00:00 2001 From: Valery Chernov Date: Sun, 29 Jan 2023 16:24:32 +0300 Subject: [PATCH 1/3] add SequenceLength op --- python/tvm/relay/frontend/onnx.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/python/tvm/relay/frontend/onnx.py b/python/tvm/relay/frontend/onnx.py index 7b35d4a48135..6e0c7cc2dd3f 100644 --- a/python/tvm/relay/frontend/onnx.py +++ b/python/tvm/relay/frontend/onnx.py @@ -6148,6 +6148,15 @@ def _impl_v11(cls, inputs, attr, params): return _expr.Tuple(inputs) +class SequenceLength(OnnxOpConverter): + """Operator converter for sequence length op.""" + + @classmethod + def _impl_v11(cls, inputs, attr, params): + # Get length of input sequence + return _expr.const(len(inputs[0]), dtype="int64") + + class SequenceInsert(OnnxOpConverter): """Operator converter for sequence insert op.""" @@ -6483,6 +6492,7 @@ def _get_convert_map(opset): "LinearRegressor": LinearRegressor.get_converter(opset), # Sequence operators "SequenceConstruct": SequenceConstruct.get_converter(opset), + "SequenceLength": SequenceLength.get_converter(opset), "SequenceInsert": SequenceInsert.get_converter(opset), "ConcatFromSequence": ConcatFromSequence.get_converter(opset), "SplitToSequence": SplitToSequence.get_converter(opset), From 8cfc1b9acff0ee6017af86668907c6db7088bb00 Mon Sep 17 00:00:00 2001 From: Valery Chernov Date: Sun, 29 Jan 2023 16:38:35 +0300 Subject: [PATCH 2/3] add SequenceLength test --- tests/python/frontend/onnx/test_forward.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/tests/python/frontend/onnx/test_forward.py b/tests/python/frontend/onnx/test_forward.py index 4b17cfbbb3a5..af15c879d70a 100644 --- a/tests/python/frontend/onnx/test_forward.py +++ b/tests/python/frontend/onnx/test_forward.py @@ -7760,10 +7760,16 @@ def verify_sequence_ops(tensor_shape, num_tensors, axis=0, position=0, new_axis= "SplitToSequence", inputs=["concat_sequence"], outputs=["split_sequence"], axis=axis ) + # Test tensor extraction from sequence at_node = helper.make_node( "SequenceAt", inputs=["split_sequence", "position"], outputs=["output"] ) + # Test sequence length + split_node = helper.make_node( + "SequenceLength", inputs=["concat_sequence"], outputs=["output_2"] + ) + if new_axis is not None: new_axis_attr = helper.make_attribute("new_axis", new_axis) concat_node.attribute.append(new_axis_attr) @@ -7781,7 +7787,8 @@ def verify_sequence_ops(tensor_shape, num_tensors, axis=0, position=0, new_axis= output_shape[axis] = num_tensors + 1 else: output_shape[axis] = (num_tensors + 1) * output_shape[axis] - graph_outputs = [helper.make_tensor_value_info("output", TensorProto.FLOAT, output_shape)] + graph_outputs = [helper.make_tensor_value_info("output", TensorProto.FLOAT, output_shape), + helper.make_tensor_value_info("output_2", TensorProto.INT, ())] graph_nodes = [position_node, construct_node, insert_node, concat_node, split_node, at_node] From 7106a48d63fb30ea47a147775b6791d9f076c94f Mon Sep 17 00:00:00 2001 From: Valery Chernov Date: Sun, 29 Jan 2023 17:51:02 +0300 Subject: [PATCH 3/3] graph fix --- tests/python/frontend/onnx/test_forward.py | 20 +++++++++++++++----- 1 file changed, 15 insertions(+), 5 deletions(-) diff --git a/tests/python/frontend/onnx/test_forward.py b/tests/python/frontend/onnx/test_forward.py index af15c879d70a..6a780a632fb7 100644 --- a/tests/python/frontend/onnx/test_forward.py +++ b/tests/python/frontend/onnx/test_forward.py @@ -7766,8 +7766,8 @@ def verify_sequence_ops(tensor_shape, num_tensors, axis=0, position=0, new_axis= ) # Test sequence length - split_node = helper.make_node( - "SequenceLength", inputs=["concat_sequence"], outputs=["output_2"] + length_node = helper.make_node( + "SequenceLength", inputs=["split_sequence"], outputs=["output_2"] ) if new_axis is not None: @@ -7787,10 +7787,20 @@ def verify_sequence_ops(tensor_shape, num_tensors, axis=0, position=0, new_axis= output_shape[axis] = num_tensors + 1 else: output_shape[axis] = (num_tensors + 1) * output_shape[axis] - graph_outputs = [helper.make_tensor_value_info("output", TensorProto.FLOAT, output_shape), - helper.make_tensor_value_info("output_2", TensorProto.INT, ())] + graph_outputs = [ + helper.make_tensor_value_info("output", TensorProto.FLOAT, output_shape), + helper.make_tensor_value_info("output_2", TensorProto.INT64, []), + ] - graph_nodes = [position_node, construct_node, insert_node, concat_node, split_node, at_node] + graph_nodes = [ + position_node, + construct_node, + insert_node, + concat_node, + split_node, + at_node, + length_node, + ] graph = helper.make_graph( graph_nodes,