siyiweigeHEW opened a new issue, #20319:
URL: https://github.com/apache/tvm/issues/20319
### Expected behavior
`tvm.relax.frontend.torch._one_hot`
(`python/tvm/relax/frontend/torch/fx_translator.py:725-735` and
`python/tvm/relax/frontend/torch/exported_program_translator.py:1198-1210`)
reads the
`num_classes` argument of a PyTorch `F.one_hot` call and forwards it
verbatim to
`relax.op.one_hot`. It should validate that `num_classes > 0` and raise a
clear,
frontend-level error when it is not — instead of letting an arbitrary value
reach the
C++ op and surface as an opaque low-level assertion.
### Actual behavior
`_one_hot` performs **no validation of `num_classes`**. Any non-positive
value
(`0`, `-1`, `-2`, …) is passed straight to `relax.op.one_hot`, whose C++
builder
(`src/relax/op/tensor/manipulate.cc:3112`) enforces `TVM_FFI_ICHECK(depth >
0)` and
raises an unhelpful `InternalError`:
```
InternalError: Check failed: (depth > 0) is false: one_hot: depth must be
positive, but got 0
```
The failure message gives no hint that the `num_classes` argument is at
fault or what
the fix is. Notably, `num_classes=0` — which is a plain constant — is
**accepted by both
`torch.export.export` and `fx.symbolic_trace`**, so a data-driven model
(e.g. a class
count that evaluates to 0) converts through the torch toolchain without
complaint and
then crashes TVM:
```
# torch.export succeeds:
ep = torch.export.export(lambda x: F.one_hot(x, num_classes=0), (x,)) # OK
# fx.symbolic_trace succeeds:
gm = fx.symbolic_trace(M) # OK
# ...but TVM's legacy from_fx path crashes on both:
mod = from_fx(gm, [((3,), "int64")])
# InternalError: Check failed: (depth > 0) is false: one_hot: depth must be
positive, but got 0
```
Additional context:
- In the **modern path** (`from_exported_program`), `_one_hot` is
effectively **dead
code**: `from_exported_program` runs
`exported_program.run_decompositions()` by
default (`exported_program_translator.py:2386-2387`), which rewrites
`aten.one_hot` into `arange + unsqueeze + eq + _to_copy` before dispatch —
so the
recommended path never reaches `relax.op.one_hot` (verified by
instrumentation,
`_one_hot` call count = 0). The validation gap is therefore only
triggerable via the
legacy `from_fx` path, but it lives in shared converter logic and the
recommended
path itself relies on torch's decomposition for every `num_classes` value.
- The documented PyTorch default `num_classes=-1` (auto-infer `max+1`) is a
data-dependent value and is rejected by `torch.export` itself
(`GuardOnDataDependentSymNode`) before TVM is reached; the legacy
`from_fx` path on
the same default reports the misleading
`ValueError: num_classes not found in node.args or node.kwargs` (the
argument is
absent rather than `-1`).
### Environment
- OS: Linux
- TVM: v0.24.dev0 (main branch, commit `390af87345`, built 2026-08-18)
- Python: 3.11
- torch: 2.10.0+cu128
- target: `llvm`
### Steps to reproduce
```python
"""Repro: TVM torch frontend crashes on non-positive num_classes (no
validation)."""
import torch
import torch.nn.functional as F
from tvm import relax
from tvm.relax.frontend.torch import from_fx
x = torch.tensor([0, 2, 1], dtype=torch.int64)
class M(torch.nn.Module):
def forward(self, x):
return F.one_hot(x, num_classes=0) # non-positive constant
# 1) The torch toolchain accepts the model:
import torch.fx as fx
gm = fx.symbolic_trace(M()) # OK
print("fx.symbolic_trace:", " ".join(str(n) for n in gm.graph.nodes))
# 2) TVM legacy from_fx path crashes at conversion:
try:
mod = from_fx(gm, [((3,), "int64")])
ex = relax.build(mod, target="llvm")
print("TVM from_fx: OK")
except Exception as e:
print(f"TVM from_fx: {type(e).__name__}: {e}")
# 3) torch.export also accepts num_classes=0 (context):
ep = torch.export.export(M(), (x,))
print("torch.export:", " ".join(str(n.target) for n in ep.graph.nodes if
n.op == "call_function"))
```
Actual output:
```
fx.symbolic_trace: x one_hot output
TVM from_fx: InternalError: Check failed: (depth > 0) is false: one_hot:
depth must be positive, but got 0
torch.export: one_hot.default
```
The same crash occurs for `num_classes=-1` and `num_classes=-2`
(`... depth must be positive, but got -1` / `-2`).
### Triage
* needs-triage
* bug
* relax
* frontend/torch
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]