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)