From b36d5c6bf63f3626661420932a6896bda79ede54 Mon Sep 17 00:00:00 2001 From: Masahiro Masuda Date: Mon, 12 Dec 2022 09:02:24 +0900 Subject: [PATCH 1/3] add baddbmm conversion --- python/tvm/relay/frontend/pytorch.py | 9 +++++++++ tests/python/frontend/pytorch/test_forward.py | 12 ++++++++++++ 2 files changed, 21 insertions(+) diff --git a/python/tvm/relay/frontend/pytorch.py b/python/tvm/relay/frontend/pytorch.py index 30f14b490b1b..309f22dc0174 100644 --- a/python/tvm/relay/frontend/pytorch.py +++ b/python/tvm/relay/frontend/pytorch.py @@ -1863,6 +1863,13 @@ def chunk(self, inputs, input_types): return _op.split(data, indeces, axis) + def baddbmm(self, inputs, _): + input = inputs[0] + batch1, batch2 = inputs[1:3] + beta = inputs[3] + alpha = inputs[4] + return _expr.const(beta) * input + _expr.const(alpha) * (_op.nn.batch_matmul(batch1, batch2, transpose_b=False)) + def matmul(self, inputs, input_types): inputs_0 = inputs[0] @@ -3621,6 +3628,7 @@ def create_convert_map(self): "aten::unsafe_chunk": self.chunk, "aten::matmul": self.matmul, "aten::bmm": self.matmul, + "aten::baddbmm": self.baddbmm, "aten::expand": self.expand, "aten::Int": self.int, "prim::NumToTensor": self.numtotensor, @@ -4587,6 +4595,7 @@ def from_pytorch( if inp.type().kind() == "TupleType" or inp.type().kind() == "ListType": enable_lower_all_tuples = False break + _run_jit_passes(graph, enable_lower_all_tuples) if custom_convert_map: diff --git a/tests/python/frontend/pytorch/test_forward.py b/tests/python/frontend/pytorch/test_forward.py index 36bb5bede475..1d16e13f11f0 100755 --- a/tests/python/frontend/pytorch/test_forward.py +++ b/tests/python/frontend/pytorch/test_forward.py @@ -5038,5 +5038,17 @@ def _test_multinomial(num_samples): ) +@tvm.testing.uses_gpu +def test_baddbmm(): + def test_fn(alpha, beta): + return lambda inp, batch1, batch2: torch.baddbmm(inp, batch1, batch2, beta=beta, alpha=alpha) + + M = torch.randn(10, 3, 5) + batch1 = torch.randn(10, 3, 4) + batch2 = torch.randn(10, 4, 5) + + verify_model(test_fn(0.5, 1.0), [M, batch1, batch2]) + + if __name__ == "__main__": tvm.testing.main() From 8530f5aee02ff362384c7ec08195be38f35ece41 Mon Sep 17 00:00:00 2001 From: Masahiro Masuda Date: Mon, 12 Dec 2022 09:25:33 +0900 Subject: [PATCH 2/3] fix --- python/tvm/relay/frontend/pytorch.py | 15 +++++++++++---- tests/python/frontend/pytorch/test_forward.py | 4 +++- 2 files changed, 14 insertions(+), 5 deletions(-) diff --git a/python/tvm/relay/frontend/pytorch.py b/python/tvm/relay/frontend/pytorch.py index 309f22dc0174..b9d167ad2d86 100644 --- a/python/tvm/relay/frontend/pytorch.py +++ b/python/tvm/relay/frontend/pytorch.py @@ -1866,9 +1866,9 @@ def chunk(self, inputs, input_types): def baddbmm(self, inputs, _): input = inputs[0] batch1, batch2 = inputs[1:3] - beta = inputs[3] - alpha = inputs[4] - return _expr.const(beta) * input + _expr.const(alpha) * (_op.nn.batch_matmul(batch1, batch2, transpose_b=False)) + beta = _expr.const(float(inputs[3])) + alpha = _expr.const(float(inputs[4])) + return beta * input + alpha * _op.nn.batch_matmul(batch1, batch2, transpose_b=False) def matmul(self, inputs, input_types): @@ -2572,7 +2572,14 @@ def numel(self, inputs, input_types): return _op.ndarray_size(inputs[0]) def empty(self, inputs, input_types): - shape = inputs[0] + shape = [] + for s in inputs[0]: + if isinstance(s, _expr.Constant): + shape.append(s.data.numpy().item()) + else: + assert isinstance(s, int) + shape.append(s) + return _op.zeros(shape, _convert_dtype_value(inputs[1])) def empty_like(self, inputs, input_types): diff --git a/tests/python/frontend/pytorch/test_forward.py b/tests/python/frontend/pytorch/test_forward.py index 1d16e13f11f0..5daf9d4edb9f 100755 --- a/tests/python/frontend/pytorch/test_forward.py +++ b/tests/python/frontend/pytorch/test_forward.py @@ -5041,7 +5041,9 @@ def _test_multinomial(num_samples): @tvm.testing.uses_gpu def test_baddbmm(): def test_fn(alpha, beta): - return lambda inp, batch1, batch2: torch.baddbmm(inp, batch1, batch2, beta=beta, alpha=alpha) + return lambda inp, batch1, batch2: torch.baddbmm( + inp, batch1, batch2, beta=beta, alpha=alpha + ) M = torch.randn(10, 3, 5) batch1 = torch.randn(10, 3, 4) From d872003d280ff9c8c81ba51d158498e6800dbcf9 Mon Sep 17 00:00:00 2001 From: Masahiro Masuda Date: Mon, 12 Dec 2022 17:05:20 +0900 Subject: [PATCH 3/3] suppress lint --- tests/python/frontend/pytorch/test_forward.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/python/frontend/pytorch/test_forward.py b/tests/python/frontend/pytorch/test_forward.py index 5daf9d4edb9f..35242fbf7dde 100755 --- a/tests/python/frontend/pytorch/test_forward.py +++ b/tests/python/frontend/pytorch/test_forward.py @@ -14,7 +14,7 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -# pylint: disable=import-self, invalid-name, unused-argument +# pylint: disable=import-self, invalid-name, unused-argument, missing-function-docstring """Unit tests for various models and operators""" import os import platform