This is an automated email from the ASF dual-hosted git repository.

tlopex pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git


The following commit(s) were added to refs/heads/main by this push:
     new fc013fd528 [Fix][Relax][Frontend][Torch] Emit int64 indices for sort 
and argsort (#20254)
fc013fd528 is described below

commit fc013fd52848a789a37099260f558680edc9bf84
Author: Chen Yufan <[email protected]>
AuthorDate: Mon Sep 14 11:37:03 2026 +0800

    [Fix][Relax][Frontend][Torch] Emit int64 indices for sort and argsort 
(#20254)
    
    ### Problem
    
    `torch.sort` and `torch.argsort` return `int64` indices. The Relax
    PyTorch frontend emits `int32` for them.
    
    `relax.op.argsort` defaults to `dtype="int32"`, and both call sites in
    `base_fx_graph_translator.py` take that default:
    
    ```python
    # _argsort
    return self.block_builder.emit(relax.op.argsort(x, dim, descending))
    
    # _sort
    indices = self.block_builder.emit(relax.op.argsort(x, dim, descending))
    ```
    
    To be precise about the scope: the sorted values and the index *values*
    are already correct. Only the index dtype diverges from PyTorch.
    
    ### Why this looks like an oversight rather than a deliberate choice
    
    `relax.op.topk` has the same `dtype="int32"` default, and `_topk` in
    this very file explicitly overrides it:
    
    ```python
    relax.op.topk(x, k=k, axis=dim, largest=largest, ret_type="both", 
dtype="int64")
    ```
    
    So the frontend already establishes that the relax-level default has to
    be overridden to match PyTorch. Every index-producing torch op the
    frontend supports comes out as `int64` — except `sort` and `argsort`:
    
    | torch op | torch index dtype | frontend before this PR |
    | --- | --- | --- |
    | `torch.sort` | `int64` | **`int32`** |
    | `torch.argsort` | `int64` | **`int32`** |
    | `torch.topk` | `int64` | `int64` (explicit `dtype=`) |
    | `torch.argmax` / `torch.argmin` | `int64` | `int64` |
    | `torch.max(dim=)` / `torch.median(dim=)` | `int64` | `int64` |
    | `torch.bucketize` | `int64` | `int64` |
    
    A single imported graph that uses both `topk` and `sort` therefore
    carries two different index dtypes for the same kind of value.
    
    ### Reproduce
    
    ```python
    import torch
    from torch import nn
    from tvm.relax.frontend.torch import from_exported_program
    
    
    def relax_dtypes(fn, x):
        class M(nn.Module):
            def forward(self, t):
                return fn(t)
    
        ep = torch.export.export(M().eval(), (x,))
        return [str(f.dtype) for f in 
from_exported_program(ep)["main"].ret_ty.fields]
    
    
    x = torch.randn(3, 4)
    print("sort   ", relax_dtypes(lambda t: torch.sort(t, dim=1), x))
    print("argsort", relax_dtypes(lambda t: torch.argsort(t, dim=1), x))
    print("topk   ", relax_dtypes(lambda t: torch.topk(t, 2, dim=1), x))
    ```
    
    Before / after:
    
    ```
                     before                        after
    sort     ['float32', 'int32']          ['float32', 'int64']
    argsort  ['int32']                     ['int64']
    topk     ['float32', 'int64']          ['float32', 'int64']
    ```
    
    ### Fix
    
    Pass `dtype="int64"` at both `relax.op.argsort` call sites, matching
    what `_topk` already does.
    
    ### Verification
    
    Built with LLVM and ran the imported graph through the VM against
    PyTorch on `sort(dim=1)`, `sort(descending=True)`, `sort(dim=0)` and
    `argsort(dim=1)`:
    
    - before: values equal, indices equal, **index dtype `int32` vs torch
    `int64`**
    - after: values equal, indices equal, index dtype `int64` — matches
    torch on every case
    
    Test suites, `tests/python/relax/test_frontend_from_fx.py` +
    `tests/python/relax/test_frontend_from_exported_program.py`:
    
    - clean `main`: 24 failed, 412 passed, 3 skipped
    - with this change: 24 failed, **413** passed, 3 skipped
    
    The 24 failures are pre-existing on `main` in my environment
    (`test_dtypes` and friends) and are identical before and after; the one
    additional pass is the new `test_sort`.
    
    `ruff format --check` and `ruff check` are clean on the touched files.
    
    ### Tests
    
    - Updated the expected IR in `test_argsort` (both frontend test files)
    and `test_sort` (`test_frontend_from_fx.py`) — these had the `int32`
    result written into them.
    - Added `test_sort` to `test_frontend_from_exported_program.py`; that
    path had no coverage for `torch.sort`.
    
    ---
    
    This change was prepared with AI assistance (Claude). I have reviewed
    and verified it, and can speak to it in review.
    
    Co-authored-by: Claude <[email protected]>
---
 .../frontend/torch/base_fx_graph_translator.py     |  4 +-
 .../relax/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 b0682f1f4c..5b723350cd 100644
--- a/python/tvm/relax/frontend/torch/base_fx_graph_translator.py
+++ b/python/tvm/relax/frontend/torch/base_fx_graph_translator.py
@@ -2002,7 +2002,7 @@ class BaseFXGraphImporter(metaclass=abc.ABCMeta):
         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)
@@ -2495,7 +2495,7 @@ class BaseFXGraphImporter(metaclass=abc.ABCMeta):
         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 c74573c49f..b70774be65 100644
--- a/tests/python/relax/test_frontend_from_exported_program.py
+++ b/tests/python/relax/test_frontend_from_exported_program.py
@@ -8107,18 +8107,18 @@ def test_argsort():
     @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
 
@@ -8126,6 +8126,39 @@ def test_argsort():
     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 a0ba7971db..961d863dcc 100644
--- a/tests/python/relax/test_frontend_from_fx.py
+++ b/tests/python/relax/test_frontend_from_fx.py
@@ -6239,10 +6239,12 @@ def test_argsort():
         @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
 
@@ -6258,18 +6260,18 @@ def test_sort():
     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)

Reply via email to