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 d600987cdd7b..8d007c39cfbf 100644 --- a/python/tvm/relax/frontend/torch/base_fx_graph_translator.py +++ b/python/tvm/relax/frontend/torch/base_fx_graph_translator.py @@ -1851,7 +1851,7 @@ def _argsort(self, node: fx.Node) -> relax.Var: x = self.env[node.args[0]] dim = node.args[1] if len(node.args) > 1 else node.kwargs.get("dim", -1) descending = node.args[2] if len(node.args) > 2 else node.kwargs.get("descending", False) - return self.block_builder.emit(relax.op.argsort(x, dim, descending)) + return self.block_builder.emit(relax.op.argsort(x, dim, descending, dtype="int64")) def _broadcast_to(self, node: fx.Node) -> relax.Var: args = self.retrieve_args(node) @@ -2316,7 +2316,7 @@ def _sort(self, node: fx.Node) -> relax.Var: dim = node.args[1] if len(node.args) > 1 else node.kwargs.get("dim", -1) descending = node.args[2] if len(node.args) > 2 else node.kwargs.get("descending", False) - indices = self.block_builder.emit(relax.op.argsort(x, dim, descending)) + indices = self.block_builder.emit(relax.op.argsort(x, dim, descending, dtype="int64")) values = self.block_builder.emit(relax.op.gather_elements(x, indices, axis=dim)) return self.block_builder.emit(relax.Tuple([values, indices])) diff --git a/tests/python/relax/test_frontend_from_exported_program.py b/tests/python/relax/test_frontend_from_exported_program.py index 7dc3c7356414..f08e1a8a3757 100644 --- a/tests/python/relax/test_frontend_from_exported_program.py +++ b/tests/python/relax/test_frontend_from_exported_program.py @@ -7784,18 +7784,18 @@ def forward(self, x): @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((5, 3), dtype="float32")) -> R.Tuple(R.Tensor((5, 3), dtype="int32")): + def main(x: R.Tensor((5, 3), dtype="float32")) -> R.Tuple(R.Tensor((5, 3), dtype="int64")): with R.dataflow(): - lv: R.Tensor((5, 3), dtype="int32") = R.argsort( - x, axis=1, descending=True, dtype="int32" + lv: R.Tensor((5, 3), dtype="int64") = R.argsort( + x, axis=1, descending=True, dtype="int64" ) lv1: R.Tensor((5, 3), dtype="float32") = R.gather_elements(x, lv, axis=1) - lv2: R.Tuple(R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int32")) = ( + lv2: R.Tuple(R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int64")) = ( lv1, lv, ) - lv3: R.Tensor((5, 3), dtype="int32") = lv2[1] - gv: R.Tuple(R.Tensor((5, 3), dtype="int32")) = (lv3,) + lv3: R.Tensor((5, 3), dtype="int64") = lv2[1] + gv: R.Tuple(R.Tensor((5, 3), dtype="int64")) = (lv3,) R.output(gv) return gv @@ -7803,6 +7803,39 @@ def main(x: R.Tensor((5, 3), dtype="float32")) -> R.Tuple(R.Tensor((5, 3), dtype verify_model(Argsort(), example_args, {}, Expected) +def test_sort(): + class Sort(Module): + def forward(self, x): + return torch.sort(x, dim=1, descending=True) + + @tvm.script.ir_module + class Expected: + @R.function + def main(x: R.Tensor((5, 3), dtype="float32")) -> R.Tuple( + R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int64") + ): + with R.dataflow(): + lv: R.Tensor((5, 3), dtype="int64") = R.argsort( + x, axis=1, descending=True, dtype="int64" + ) + lv1: R.Tensor((5, 3), dtype="float32") = R.gather_elements(x, lv, axis=1) + lv2: R.Tuple(R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int64")) = ( + lv1, + lv, + ) + lv3: R.Tensor((5, 3), dtype="float32") = lv2[0] + lv4: R.Tensor((5, 3), dtype="int64") = lv2[1] + gv: R.Tuple(R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int64")) = ( + lv3, + lv4, + ) + R.output(gv) + return gv + + example_args = (torch.randn(5, 3, dtype=torch.float32),) + verify_model(Sort(), example_args, {}, Expected) + + def test_topk(): class Topk(Module): def forward(self, x): diff --git a/tests/python/relax/test_frontend_from_fx.py b/tests/python/relax/test_frontend_from_fx.py index a489977958c7..595e02aa80de 100644 --- a/tests/python/relax/test_frontend_from_fx.py +++ b/tests/python/relax/test_frontend_from_fx.py @@ -6067,10 +6067,12 @@ class Expected: @R.function def main( inp_0: R.Tensor((5, 3), dtype="float32"), - ) -> R.Tensor((5, 3), dtype="int32"): + ) -> R.Tensor((5, 3), dtype="int64"): with R.dataflow(): - lv: R.Tensor((5, 3), dtype="int32") = R.argsort(inp_0, axis=1, descending=True) - gv: R.Tensor((5, 3), dtype="int32") = lv + lv: R.Tensor((5, 3), dtype="int64") = R.argsort( + inp_0, axis=1, descending=True, dtype="int64" + ) + gv: R.Tensor((5, 3), dtype="int64") = lv R.output(gv) return gv @@ -6086,18 +6088,18 @@ def forward(self, x): class Expected: @R.function def main(inp_0: R.Tensor((5, 3), dtype="float32")) -> R.Tuple( - R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int32") + R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int64") ): with R.dataflow(): - lv: R.Tensor((5, 3), dtype="int32") = R.argsort( - inp_0, axis=1, descending=True, dtype="int32" + lv: R.Tensor((5, 3), dtype="int64") = R.argsort( + inp_0, axis=1, descending=True, dtype="int64" ) lv1: R.Tensor((5, 3), dtype="float32") = R.gather_elements(inp_0, lv, axis=1) - lv2: R.Tuple(R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int32")) = ( + lv2: R.Tuple(R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int64")) = ( lv1, lv, ) - gv: R.Tuple(R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int32")) = ( + gv: R.Tuple(R.Tensor((5, 3), dtype="float32"), R.Tensor((5, 3), dtype="int64")) = ( lv2 ) R.output(gv)