diff --git a/python/tvm/relax/frontend/torch/base_fx_graph_translator.py b/python/tvm/relax/frontend/torch/base_fx_graph_translator.py index 3ff3b596af3b..f86043952cb3 100644 --- a/python/tvm/relax/frontend/torch/base_fx_graph_translator.py +++ b/python/tvm/relax/frontend/torch/base_fx_graph_translator.py @@ -1942,6 +1942,21 @@ def _norm(self, node: fx.Node) -> relax.Var: ) ) + def _amax_amin(self, op: Callable) -> Callable: + """torch.amax / torch.amin: reduce over ``dim`` (a list; empty means every axis).""" + from torch import fx + + def convert(node: fx.Node) -> relax.Var: + args = self.retrieve_args(node) + x = args[0] + dim = args[1] if len(node.args) > 1 else node.kwargs.get("dim", []) + keepdim = args[2] if len(node.args) > 2 else node.kwargs.get("keepdim", False) + if isinstance(dim, list | tuple) and len(dim) == 0: + dim = None + return self.block_builder.emit(op(x, dim, keepdims=keepdim)) + + return convert + def _prod(self, node: fx.Node) -> relax.Var: args = self.retrieve_args(node) x = args[0] diff --git a/python/tvm/relax/frontend/torch/exported_program_translator.py b/python/tvm/relax/frontend/torch/exported_program_translator.py index 86c936723bf2..dd7ca398beaf 100644 --- a/python/tvm/relax/frontend/torch/exported_program_translator.py +++ b/python/tvm/relax/frontend/torch/exported_program_translator.py @@ -1426,23 +1426,28 @@ def _exponential(self, node: fx.Node) -> relax.Var: x = self.env[node.args[0]] return self.block_builder.emit(relax.op.zeros_like(x)) - def _max_dim(self, node: fx.Node) -> relax.Var: - x = self.env[node.args[0]] - dim = node.args[1] - keepdim = node.args[2] if len(node.args) > 2 else node.kwargs.get("keepdim", False) + def _max_min_dim(self, largest: bool) -> Callable: + """torch.max(x, dim) / torch.min(x, dim): the (values, indices) pair along one axis.""" - topk_res = self.block_builder.emit( - relax.op.topk(x, k=1, axis=dim, largest=True, ret_type="both", dtype="int64") - ) + def convert(node: fx.Node) -> relax.Var: + x = self.env[node.args[0]] + dim = node.args[1] + keepdim = node.args[2] if len(node.args) > 2 else node.kwargs.get("keepdim", False) - values = topk_res[0] - indices = topk_res[1] + topk_res = self.block_builder.emit( + relax.op.topk(x, k=1, axis=dim, largest=largest, ret_type="both", dtype="int64") + ) - if not keepdim: - values = self.block_builder.emit(relax.op.squeeze(values, axis=[dim])) - indices = self.block_builder.emit(relax.op.squeeze(indices, axis=[dim])) + values = topk_res[0] + indices = topk_res[1] - return self.block_builder.emit(relax.Tuple([values, indices])) + if not keepdim: + values = self.block_builder.emit(relax.op.squeeze(values, axis=[dim])) + indices = self.block_builder.emit(relax.op.squeeze(indices, axis=[dim])) + + return self.block_builder.emit(relax.Tuple([values, indices])) + + return convert def _alias(self, node: fx.Node) -> relax.Var: return self.env[node.args[0]] @@ -1878,6 +1883,8 @@ def create_convert_map( "min.other": self._binary_op(relax.op.minimum, min), "max.default": self._unary_op(relax.op.max), "min.default": self._unary_op(relax.op.min), + "amax.default": self._amax_amin(relax.op.max), + "amin.default": self._amax_amin(relax.op.min), "maximum.default": self._binary_op(relax.op.maximum, torch.maximum), "minimum.default": self._binary_op(relax.op.minimum, torch.minimum), "remainder.Tensor": self._binary_op(relax.op.floor_mod, operator.mod), @@ -1969,7 +1976,8 @@ def create_convert_map( "sum.default": self._sum, "sum.dim_IntList": self._sum, "var.correction": self._var, - "max.dim": self._max_dim, + "max.dim": self._max_min_dim(largest=True), + "min.dim": self._max_min_dim(largest=False), "median.dim": self._median, "median.default": self._median, # search diff --git a/python/tvm/relax/frontend/torch/fx_translator.py b/python/tvm/relax/frontend/torch/fx_translator.py index 2e35ce6ce704..de313b78d235 100644 --- a/python/tvm/relax/frontend/torch/fx_translator.py +++ b/python/tvm/relax/frontend/torch/fx_translator.py @@ -994,6 +994,8 @@ def create_convert_map( "chunk": self._chunk, "concat": self._cat, "contiguous": lambda node: self.env[node.args[0]], + "amax": self._amax_amin(relax.op.max), + "amin": self._amax_amin(relax.op.min), "cumprod": self._cumprod, "cumsum": self._cumsum, "expand": self._expand, diff --git a/tests/python/relax/test_frontend_from_exported_program.py b/tests/python/relax/test_frontend_from_exported_program.py index d78629f5737b..86564759d63f 100644 --- a/tests/python/relax/test_frontend_from_exported_program.py +++ b/tests/python/relax/test_frontend_from_exported_program.py @@ -8568,6 +8568,120 @@ def main(x: R.Tensor((4, 8), dtype="float32")) -> R.Tuple( verify_model(Exponential(), example_args, {}, Expected) +def test_amax_amin(): + class Amax(Module): + def forward(self, x): + return torch.amax(x, dim=1) + + class AminKeep(Module): + def forward(self, x): + return torch.amin(x, dim=(0, 2), keepdim=True) + + class AmaxAll(Module): + def forward(self, x): + return torch.amax(x) + + @I.ir_module + class expected_amax: + @R.function + def main(x: R.Tensor((4, 8, 16), dtype="float32")) -> R.Tuple( + R.Tensor((4, 16), dtype="float32") + ): + with R.dataflow(): + lv: R.Tensor((4, 16), dtype="float32") = R.max(x, axis=[1], keepdims=False) + gv: R.Tuple(R.Tensor((4, 16), dtype="float32")) = (lv,) + R.output(gv) + return gv + + @I.ir_module + class expected_amin_keep: + @R.function + def main(x: R.Tensor((4, 8, 16), dtype="float32")) -> R.Tuple( + R.Tensor((1, 8, 1), dtype="float32") + ): + with R.dataflow(): + lv: R.Tensor((1, 8, 1), dtype="float32") = R.min(x, axis=[0, 2], keepdims=True) + gv: R.Tuple(R.Tensor((1, 8, 1), dtype="float32")) = (lv,) + R.output(gv) + return gv + + @I.ir_module + class expected_amax_all: + @R.function + def main(x: R.Tensor((4, 8, 16), dtype="float32")) -> R.Tuple( + R.Tensor((), dtype="float32") + ): + with R.dataflow(): + lv: R.Tensor((), dtype="float32") = R.max(x, axis=None, keepdims=False) + gv: R.Tuple(R.Tensor((), dtype="float32")) = (lv,) + R.output(gv) + return gv + + example_args = (torch.randn(4, 8, 16, dtype=torch.float32),) + verify_model(Amax(), example_args, {}, expected_amax) + verify_model(AminKeep(), example_args, {}, expected_amin_keep) + verify_model(AmaxAll(), example_args, {}, expected_amax_all) + for model in (Amax(), AminKeep(), AmaxAll()): + verify_model_numerically(model, example_args) + # logsumexp decomposes through amax, so it is covered by the same converter. + verify_model_numerically( + type("LogSumExp", (Module,), {"forward": lambda self, x: torch.logsumexp(x, dim=1)})(), + example_args, + rtol=1e-5, + atol=1e-5, + ) + + +def test_min_dim(): + class MinDim(Module): + def forward(self, x): + return torch.min(x, dim=1) + + class MinDimKeep(Module): + def forward(self, x): + return torch.min(x, dim=1, keepdim=True) + + @I.ir_module + class expected1: + @R.function + def main(x: R.Tensor((4, 8, 16), dtype="float32")) -> R.Tuple( + R.Tensor((4, 16), dtype="float32"), R.Tensor((4, 16), dtype="int64") + ): + with R.dataflow(): + lv: R.Tuple( + R.Tensor((4, 1, 16), dtype="float32"), R.Tensor((4, 1, 16), dtype="int64") + ) = R.topk(x, k=1, axis=1, ret_type="both", largest=False, dtype="int64") + lv1: R.Tensor((4, 1, 16), dtype="float32") = lv[0] + lv2: R.Tensor((4, 16), dtype="float32") = R.squeeze(lv1, axis=[1]) + lv3: R.Tensor((4, 1, 16), dtype="int64") = lv[1] + lv4: R.Tensor((4, 16), dtype="int64") = R.squeeze(lv3, axis=[1]) + lv5: R.Tuple( + R.Tensor((4, 16), dtype="float32"), R.Tensor((4, 16), dtype="int64") + ) = (lv2, lv4) + lv6: R.Tensor((4, 16), dtype="float32") = lv5[0] + lv7: R.Tensor((4, 16), dtype="int64") = lv5[1] + gv: R.Tuple( + R.Tensor((4, 16), dtype="float32"), R.Tensor((4, 16), dtype="int64") + ) = (lv6, lv7) + R.output(gv) + return gv + + example_args = (torch.randn(4, 8, 16, dtype=torch.float32),) + verify_model(MinDim(), example_args, {}, expected1) + # Values and indices, both branches, against torch. Distinct values keep the + # argmin unambiguous. + x = torch.randperm(4 * 8 * 16).reshape(4, 8, 16).to(torch.float32) + for model in (MinDim(), MinDimKeep()): + with torch.no_grad(): + want_v, want_i = model(x) + mod = from_exported_program(export(model, (x,))) + ex = relax.build(mod, tvm.target.Target("llvm")) + vm = relax.VirtualMachine(ex, tvm.cpu()) + got = vm["main"](tvm.runtime.tensor(x.numpy())) + tvm.testing.assert_allclose(got[0].numpy(), want_v.numpy()) + tvm.testing.assert_allclose(got[1].numpy(), want_i.numpy()) + + def test_max_dim(): class MaxDim1(Module): def forward(self, x):