Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 15 additions & 0 deletions python/tvm/relax/frontend/torch/base_fx_graph_translator.py
Original file line number Diff line number Diff line change
Expand Up @@ -1942,6 +1942,21 @@ def _norm(self, node: fx.Node) -> relax.Var:
)
)

def _amax_amin(self, op: Callable) -> Callable:
"""torch.amax / torch.amin: reduce over ``dim`` (a list; empty means every axis)."""
from torch import fx

def convert(node: fx.Node) -> relax.Var:
args = self.retrieve_args(node)
x = args[0]
dim = args[1] if len(node.args) > 1 else node.kwargs.get("dim", [])
keepdim = args[2] if len(node.args) > 2 else node.kwargs.get("keepdim", False)
if isinstance(dim, list | tuple) and len(dim) == 0:
dim = None
return self.block_builder.emit(op(x, dim, keepdims=keepdim))

return convert

def _prod(self, node: fx.Node) -> relax.Var:
args = self.retrieve_args(node)
x = args[0]
Expand Down
36 changes: 22 additions & 14 deletions python/tvm/relax/frontend/torch/exported_program_translator.py
Original file line number Diff line number Diff line change
Expand Up @@ -1426,23 +1426,28 @@ def _exponential(self, node: fx.Node) -> relax.Var:
x = self.env[node.args[0]]
return self.block_builder.emit(relax.op.zeros_like(x))

def _max_dim(self, node: fx.Node) -> relax.Var:
x = self.env[node.args[0]]
dim = node.args[1]
keepdim = node.args[2] if len(node.args) > 2 else node.kwargs.get("keepdim", False)
def _max_min_dim(self, largest: bool) -> Callable:
"""torch.max(x, dim) / torch.min(x, dim): the (values, indices) pair along one axis."""

topk_res = self.block_builder.emit(
relax.op.topk(x, k=1, axis=dim, largest=True, ret_type="both", dtype="int64")
)
def convert(node: fx.Node) -> relax.Var:
x = self.env[node.args[0]]
dim = node.args[1]
keepdim = node.args[2] if len(node.args) > 2 else node.kwargs.get("keepdim", False)

values = topk_res[0]
indices = topk_res[1]
topk_res = self.block_builder.emit(
relax.op.topk(x, k=1, axis=dim, largest=largest, ret_type="both", dtype="int64")
)

if not keepdim:
values = self.block_builder.emit(relax.op.squeeze(values, axis=[dim]))
indices = self.block_builder.emit(relax.op.squeeze(indices, axis=[dim]))
values = topk_res[0]
indices = topk_res[1]

return self.block_builder.emit(relax.Tuple([values, indices]))
if not keepdim:
values = self.block_builder.emit(relax.op.squeeze(values, axis=[dim]))
indices = self.block_builder.emit(relax.op.squeeze(indices, axis=[dim]))

return self.block_builder.emit(relax.Tuple([values, indices]))

return convert

def _alias(self, node: fx.Node) -> relax.Var:
return self.env[node.args[0]]
Expand Down Expand Up @@ -1878,6 +1883,8 @@ def create_convert_map(
"min.other": self._binary_op(relax.op.minimum, min),
"max.default": self._unary_op(relax.op.max),
"min.default": self._unary_op(relax.op.min),
"amax.default": self._amax_amin(relax.op.max),
"amin.default": self._amax_amin(relax.op.min),
"maximum.default": self._binary_op(relax.op.maximum, torch.maximum),
"minimum.default": self._binary_op(relax.op.minimum, torch.minimum),
"remainder.Tensor": self._binary_op(relax.op.floor_mod, operator.mod),
Expand Down Expand Up @@ -1969,7 +1976,8 @@ def create_convert_map(
"sum.default": self._sum,
"sum.dim_IntList": self._sum,
"var.correction": self._var,
"max.dim": self._max_dim,
"max.dim": self._max_min_dim(largest=True),
"min.dim": self._max_min_dim(largest=False),
"median.dim": self._median,
"median.default": self._median,
# search
Expand Down
2 changes: 2 additions & 0 deletions python/tvm/relax/frontend/torch/fx_translator.py
Original file line number Diff line number Diff line change
Expand Up @@ -994,6 +994,8 @@ def create_convert_map(
"chunk": self._chunk,
"concat": self._cat,
"contiguous": lambda node: self.env[node.args[0]],
"amax": self._amax_amin(relax.op.max),
"amin": self._amax_amin(relax.op.min),
"cumprod": self._cumprod,
"cumsum": self._cumsum,
"expand": self._expand,
Expand Down
114 changes: 114 additions & 0 deletions tests/python/relax/test_frontend_from_exported_program.py
Original file line number Diff line number Diff line change
Expand Up @@ -8568,6 +8568,120 @@ def main(x: R.Tensor((4, 8), dtype="float32")) -> R.Tuple(
verify_model(Exponential(), example_args, {}, Expected)


def test_amax_amin():
class Amax(Module):
def forward(self, x):
return torch.amax(x, dim=1)

class AminKeep(Module):
def forward(self, x):
return torch.amin(x, dim=(0, 2), keepdim=True)

class AmaxAll(Module):
def forward(self, x):
return torch.amax(x)

@I.ir_module
class expected_amax:
@R.function
def main(x: R.Tensor((4, 8, 16), dtype="float32")) -> R.Tuple(
R.Tensor((4, 16), dtype="float32")
):
with R.dataflow():
lv: R.Tensor((4, 16), dtype="float32") = R.max(x, axis=[1], keepdims=False)
gv: R.Tuple(R.Tensor((4, 16), dtype="float32")) = (lv,)
R.output(gv)
return gv

@I.ir_module
class expected_amin_keep:
@R.function
def main(x: R.Tensor((4, 8, 16), dtype="float32")) -> R.Tuple(
R.Tensor((1, 8, 1), dtype="float32")
):
with R.dataflow():
lv: R.Tensor((1, 8, 1), dtype="float32") = R.min(x, axis=[0, 2], keepdims=True)
gv: R.Tuple(R.Tensor((1, 8, 1), dtype="float32")) = (lv,)
R.output(gv)
return gv

@I.ir_module
class expected_amax_all:
@R.function
def main(x: R.Tensor((4, 8, 16), dtype="float32")) -> R.Tuple(
R.Tensor((), dtype="float32")
):
with R.dataflow():
lv: R.Tensor((), dtype="float32") = R.max(x, axis=None, keepdims=False)
gv: R.Tuple(R.Tensor((), dtype="float32")) = (lv,)
R.output(gv)
return gv

example_args = (torch.randn(4, 8, 16, dtype=torch.float32),)
verify_model(Amax(), example_args, {}, expected_amax)
verify_model(AminKeep(), example_args, {}, expected_amin_keep)
verify_model(AmaxAll(), example_args, {}, expected_amax_all)
for model in (Amax(), AminKeep(), AmaxAll()):
verify_model_numerically(model, example_args)
# logsumexp decomposes through amax, so it is covered by the same converter.
verify_model_numerically(
type("LogSumExp", (Module,), {"forward": lambda self, x: torch.logsumexp(x, dim=1)})(),
example_args,
rtol=1e-5,
atol=1e-5,
)


def test_min_dim():
class MinDim(Module):
def forward(self, x):
return torch.min(x, dim=1)

class MinDimKeep(Module):
def forward(self, x):
return torch.min(x, dim=1, keepdim=True)

@I.ir_module
class expected1:
@R.function
def main(x: R.Tensor((4, 8, 16), dtype="float32")) -> R.Tuple(
R.Tensor((4, 16), dtype="float32"), R.Tensor((4, 16), dtype="int64")
):
with R.dataflow():
lv: R.Tuple(
R.Tensor((4, 1, 16), dtype="float32"), R.Tensor((4, 1, 16), dtype="int64")
) = R.topk(x, k=1, axis=1, ret_type="both", largest=False, dtype="int64")
lv1: R.Tensor((4, 1, 16), dtype="float32") = lv[0]
lv2: R.Tensor((4, 16), dtype="float32") = R.squeeze(lv1, axis=[1])
lv3: R.Tensor((4, 1, 16), dtype="int64") = lv[1]
lv4: R.Tensor((4, 16), dtype="int64") = R.squeeze(lv3, axis=[1])
lv5: R.Tuple(
R.Tensor((4, 16), dtype="float32"), R.Tensor((4, 16), dtype="int64")
) = (lv2, lv4)
lv6: R.Tensor((4, 16), dtype="float32") = lv5[0]
lv7: R.Tensor((4, 16), dtype="int64") = lv5[1]
gv: R.Tuple(
R.Tensor((4, 16), dtype="float32"), R.Tensor((4, 16), dtype="int64")
) = (lv6, lv7)
R.output(gv)
return gv

example_args = (torch.randn(4, 8, 16, dtype=torch.float32),)
verify_model(MinDim(), example_args, {}, expected1)
# Values and indices, both branches, against torch. Distinct values keep the
# argmin unambiguous.
x = torch.randperm(4 * 8 * 16).reshape(4, 8, 16).to(torch.float32)
for model in (MinDim(), MinDimKeep()):
with torch.no_grad():
want_v, want_i = model(x)
mod = from_exported_program(export(model, (x,)))
ex = relax.build(mod, tvm.target.Target("llvm"))
vm = relax.VirtualMachine(ex, tvm.cpu())
got = vm["main"](tvm.runtime.tensor(x.numpy()))
tvm.testing.assert_allclose(got[0].numpy(), want_v.numpy())
tvm.testing.assert_allclose(got[1].numpy(), want_i.numpy())


def test_max_dim():
class MaxDim1(Module):
def forward(self, x):
Expand Down
Loading