From f10db0e3ceea4af4c8f8dc49ff01aff86614b84f Mon Sep 17 00:00:00 2001 From: Chen Yufan Date: Tue, 1 Sep 2026 18:19:09 +0800 Subject: [PATCH] [Fix][Relax][Frontend][Torch] Emit int64 indices for sort and argsort `torch.sort` and `torch.argsort` return int64 indices, but the Relax frontend emitted int32. `relax.op.argsort` defaults to `dtype="int32"`, and both call sites in `base_fx_graph_translator.py` took that default. `_topk` in the same file already overrides the identical int32 default on `relax.op.topk` to match PyTorch, and every other index-producing op the frontend supports (argmax, argmin, max.dim, median.dim, bucketize) comes out as int64. sort and argsort were the only exceptions, so a single imported graph could carry two different index dtypes for the same kind of value. The sorted values and the index values themselves were already correct; only the index dtype diverged. Pass `dtype="int64"` at both call sites, update the expectations that had the int32 result written into them, and add a `test_sort` to the exported-program tests, which had no coverage for `torch.sort`. Co-authored-by: Claude --- .../torch/base_fx_graph_translator.py | 4 +- .../test_frontend_from_exported_program.py | 45 ++++++++++++++++--- tests/python/relax/test_frontend_from_fx.py | 18 ++++---- 3 files changed, 51 insertions(+), 16 deletions(-) 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)