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
80 changes: 63 additions & 17 deletions python/tvm/relax/frontend/torch/base_fx_graph_translator.py
Original file line number Diff line number Diff line change
Expand Up @@ -572,23 +572,69 @@ def call_binary_op(op, lhs, rhs):

def _pow(self, node: fx.Node) -> relax.Var:
lhs, rhs = self.retrieve_args(node)
# torch integer pow returns an integer tensor, but relax.op.power legalizes to
# TOPI power which requires floating-point inputs. Decompose an integer base with
# a constant non-negative integer exponent into repeated multiplication instead.
if (
isinstance(lhs, relax.Expr)
and isinstance(lhs.ty, relax.TensorType)
and lhs.ty.dtype.matches_code(DataTypeCode.INT, DataTypeCode.UINT)
and isinstance(rhs, int)
and not isinstance(rhs, bool)
and rhs >= 0
):
if rhs == 0:
return self.block_builder.emit(relax.op.ones_like(lhs))
result = lhs
for _ in range(rhs - 1):
result = self.block_builder.emit(relax.op.multiply(result, lhs))
return result
if isinstance(lhs, relax.Expr) and isinstance(lhs.ty, relax.TensorType):
lhs_dtype = lhs.ty.dtype
is_integer_base = lhs_dtype.matches_code(DataTypeCode.INT, DataTypeCode.UINT)
is_float_base = lhs_dtype.matches_code(DataTypeCode.FLOAT, DataTypeCode.BFLOAT)

# A Python float promotes an integer tensor to PyTorch's default floating-point
# dtype. ExportedProgram records the inferred dtype, while plain FX does not.
if is_integer_base and isinstance(rhs, float):
output_meta = node.meta.get("val")
output_dtype = self._convert_data_type(
output_meta.dtype
if isinstance(output_meta, self.torch.Tensor)
else self.torch.get_default_dtype()
)
lhs = self.block_builder.emit(relax.op.astype(lhs, output_dtype))
lhs_dtype = lhs.ty.dtype
is_integer_base = False
is_float_base = True

# Match the scalar conversion used by PyTorch's floating-point power kernels.
exponent_dtype = {
"float16": self.torch.float16,
"bfloat16": self.torch.bfloat16,
"float32": self.torch.float64,
"float64": self.torch.float64,
}.get(str(lhs_dtype))
if (
is_float_base
and exponent_dtype is not None
and isinstance(rhs, int | float)
and not isinstance(rhs, bool)
):
rhs = self.torch.scalar_tensor(rhs, dtype=exponent_dtype, device="cpu").item()

is_nonnegative_integral_exponent = (
isinstance(rhs, int) and not isinstance(rhs, bool) and rhs >= 0
) or (isinstance(rhs, float) and rhs >= 0 and rhs.is_integer())

# TOPI power requires floating-point inputs, and some backends do not preserve
# the sign of a negative base for integral exponents. Decompose after applying
# PyTorch's dtype promotion so Python integer and float scalars behave alike.
if (is_integer_base or is_float_base) and is_nonnegative_integral_exponent:
exponent = int(rhs)
if exponent == 0:
return self.block_builder.emit(relax.op.ones_like(lhs))

# Exponentiation by squaring avoids linear graph growth for large exponents.
result = None
factor = lhs
while exponent:
if exponent & 1:
result = (
factor
if result is None
else self.block_builder.emit(relax.op.multiply(result, factor))
)
exponent >>= 1
if exponent:
factor = self.block_builder.emit(relax.op.multiply(factor, factor))
return result

if is_float_base and isinstance(rhs, float):
return self.block_builder.emit(relax.op.power(lhs, relax.const(rhs, lhs_dtype)))
return self._binary_op(relax.op.power, operator.pow)(node)

def _div(self, node: fx.Node) -> relax.Var:
Expand Down
172 changes: 158 additions & 14 deletions tests/python/relax/test_frontend_from_exported_program.py
Original file line number Diff line number Diff line change
Expand Up @@ -1079,16 +1079,141 @@ def main(input: R.Tensor((4,), dtype="int64")) -> R.Tuple(R.Tensor((4,), dtype="
# block 0
with R.dataflow():
lv: R.Tensor((4,), dtype="int64") = R.multiply(input, input)
lv1: R.Tensor((4,), dtype="int64") = R.multiply(lv, input)
lv2: R.Tensor((4,), dtype="int64") = R.multiply(lv1, input)
gv: R.Tuple(R.Tensor((4,), dtype="int64")) = (lv2,)
lv1: R.Tensor((4,), dtype="int64") = R.multiply(lv, lv)
gv: R.Tuple(R.Tensor((4,), dtype="int64")) = (lv1,)
R.output(gv)
return gv

example_args = (torch.tensor([-1, 1, 2, 3], dtype=torch.int64),)
verify_model(Pow(), example_args, {}, expected)


@pytest.mark.parametrize("exponent", [3, 3.0])
def test_pow_float_integer_exponent(exponent):
class Pow(Module):
def forward(self, input):
return input.pow(exponent)

@tvm.script.ir_module
class expected:
@R.function
def main(
input: R.Tensor((4,), dtype="float32"),
) -> R.Tuple(R.Tensor((4,), dtype="float32")):
with R.dataflow():
lv: R.Tensor((4,), dtype="float32") = R.multiply(input, input)
lv1: R.Tensor((4,), dtype="float32") = R.multiply(input, lv)
gv: R.Tuple(R.Tensor((4,), dtype="float32")) = (lv1,)
R.output(gv)
return gv

example_args = (torch.tensor([-2.0, -1.0, 1.0, 2.0], dtype=torch.float32),)
verify_model(Pow(), example_args, {}, expected)
verify_model_numerically(Pow(), example_args)


@pytest.mark.parametrize("exponent", [0, 0.0, 1, 1.0])
def test_pow_float_integer_exponent_identity_cases(exponent):
class Pow(Module):
def forward(self, input):
return input.pow(exponent)

example_args = (torch.tensor([-2.0, -1.0, 1.0, 2.0], dtype=torch.float32),)
verify_model_numerically(Pow(), example_args)


def test_pow_float_integer_exponent_large():
class Pow(Module):
def forward(self, input):
return input.pow(17.0)

@tvm.script.ir_module
class expected:
@R.function
def main(input: R.Tensor((2,), dtype="float32")) -> R.Tuple(
R.Tensor((2,), dtype="float32")
):
with R.dataflow():
lv: R.Tensor((2,), dtype="float32") = R.multiply(input, input)
lv1: R.Tensor((2,), dtype="float32") = R.multiply(lv, lv)
lv2: R.Tensor((2,), dtype="float32") = R.multiply(lv1, lv1)
lv3: R.Tensor((2,), dtype="float32") = R.multiply(lv2, lv2)
lv4: R.Tensor((2,), dtype="float32") = R.multiply(input, lv3)
gv: R.Tuple(R.Tensor((2,), dtype="float32")) = (lv4,)
R.output(gv)
return gv

example_args = (torch.tensor([-1.25, 0.5], dtype=torch.float32),)
verify_model(Pow(), example_args, {}, expected)
verify_model_numerically(Pow(), example_args, rtol=1e-6, atol=1e-6)


def test_pow_integer_base_float_exponent():
class Pow(Module):
def forward(self, input):
return input.pow(3.0)

@tvm.script.ir_module
class expected:
@R.function
def main(input: R.Tensor((4,), dtype="int32")) -> R.Tuple(R.Tensor((4,), dtype="float32")):
with R.dataflow():
lv: R.Tensor((4,), dtype="float32") = R.astype(input, dtype="float32")
lv1: R.Tensor((4,), dtype="float32") = R.multiply(lv, lv)
lv2: R.Tensor((4,), dtype="float32") = R.multiply(lv, lv1)
gv: R.Tuple(R.Tensor((4,), dtype="float32")) = (lv2,)
R.output(gv)
return gv

example_args = (torch.tensor([-2, -1, 1, 2], dtype=torch.int32),)
verify_model(Pow(), example_args, {}, expected)
verify_model_numerically(Pow(), example_args)


def test_pow_integer_base_fractional_exponent():
class Pow(Module):
def forward(self, input):
return input.pow(0.5)

@tvm.script.ir_module
class expected:
@R.function
def main(input: R.Tensor((4,), dtype="int32")) -> R.Tuple(R.Tensor((4,), dtype="float32")):
with R.dataflow():
lv: R.Tensor((4,), dtype="float32") = R.astype(input, dtype="float32")
lv1: R.Tensor((4,), dtype="float32") = R.power(lv, R.const(0.5, "float32"))
gv: R.Tuple(R.Tensor((4,), dtype="float32")) = (lv1,)
R.output(gv)
return gv

example_args = (torch.tensor([1, 4, 9, 16], dtype=torch.int32),)
verify_model(Pow(), example_args, {}, expected)
verify_model_numerically(Pow(), example_args)


@pytest.mark.parametrize(
"dtype, exponent",
[
(torch.float16, 2049.0),
(torch.bfloat16, 257.0),
(torch.float32, 2**53 + 1),
(torch.float64, 2**53 + 1),
],
)
def test_pow_float_exponent_rounding(dtype, exponent):
class Pow(Module):
def forward(self, input):
return input.pow(exponent)

input = torch.tensor([-1.0], dtype=dtype)
expected = Pow()(input)
mod = from_exported_program(export(Pow(), args=(input,)))
vm = relax.VirtualMachine(relax.build(mod, "llvm"), tvm.cpu())
actual = torch.from_dlpack(vm["main"](tvm.runtime.from_dlpack(input))[0])

torch.testing.assert_close(actual, expected)


def test_logsoftmax():
class LogSoftmax(Module):
def forward(self, input):
Expand Down Expand Up @@ -1336,17 +1461,36 @@ def __init__(self, op):
def forward(self, lhs):
return self.op(lhs, 1.0)

@tvm.script.ir_module
class expected_binary2:
@R.function
def main(
lhs: R.Tensor((10, 10), dtype="float32"),
) -> R.Tuple(R.Tensor((10, 10), dtype="float32")):
with R.dataflow():
lv: R.Tensor((10, 10), dtype="float32") = relax_op(lhs, R.const(1.0))
gv: R.Tuple(R.Tensor((10, 10), dtype="float32")) = (lv,)
R.output(gv)
return gv
if op is operator.pow:

@tvm.script.ir_module
class expected_power:
@R.function
def main(
lhs: R.Tensor((10, 10), dtype="float32"),
) -> R.Tuple(R.Tensor((10, 10), dtype="float32")):
with R.dataflow():
gv: R.Tuple(R.Tensor((10, 10), dtype="float32")) = (lhs,)
R.output(gv)
return gv

expected_binary2 = expected_power

else:

@tvm.script.ir_module
class expected_other_binary:
@R.function
def main(
lhs: R.Tensor((10, 10), dtype="float32"),
) -> R.Tuple(R.Tensor((10, 10), dtype="float32")):
with R.dataflow():
lv: R.Tensor((10, 10), dtype="float32") = relax_op(lhs, R.const(1.0))
gv: R.Tuple(R.Tensor((10, 10), dtype="float32")) = (lv,)
R.output(gv)
return gv

expected_binary2 = expected_other_binary

# In-place ops (add_, mul_, ...) produce the same Relax program as their
# functional counterparts: mutation outputs are dropped by the importer.
Expand Down
Loading
Loading