From 06111600fe88be30da5d5cb2619de3d206d256fa Mon Sep 17 00:00:00 2001 From: Arthur031221 Date: Fri, 2 Oct 2026 00:51:45 +0800 Subject: [PATCH 1/2] [Fix][Relax][Frontend][Torch] Apply the scale and input_scale arguments of aten.elu --- .../torch/exported_program_translator.py | 22 +++++++++++++++++++ .../test_frontend_from_exported_program.py | 20 ++++++++++++++++- 2 files changed, 41 insertions(+), 1 deletion(-) diff --git a/python/tvm/relax/frontend/torch/exported_program_translator.py b/python/tvm/relax/frontend/torch/exported_program_translator.py index 44804e4d2a98..8cf230fe6792 100644 --- a/python/tvm/relax/frontend/torch/exported_program_translator.py +++ b/python/tvm/relax/frontend/torch/exported_program_translator.py @@ -74,6 +74,28 @@ def _convert_pytorch_tensor_to_tvm(tensor_value: torch.Tensor) -> tvm.runtime.Te ########## Unary Ops ########## + def _elu(self, node: fx.Node) -> relax.Expr: + # aten.elu is elu(x, alpha, scale, input_scale); run_decompositions rewrites selu to it + scale = node.args[2] if len(node.args) > 2 else node.kwargs.get("scale", 1.0) + input_scale = node.args[3] if len(node.args) > 3 else node.kwargs.get("input_scale", 1.0) + if scale == 1 and input_scale == 1: + return super()._elu(node) + + x = self.env[node.args[0]] + alpha = node.args[1] if len(node.args) > 1 else node.kwargs.get("alpha", 1.0) + dtype = x.ty.dtype + bb = self.block_builder + scaled_x = x + if input_scale != 1: + scaled_x = bb.emit(relax.op.multiply(x, relax.const(input_scale, dtype))) + # scale * (ReLU(x) - alpha * ReLU(1 - exp(input_scale * x))) + negative = relax.op.multiply( + relax.const(-alpha, dtype), + relax.op.nn.relu(relax.op.subtract(relax.const(1, dtype), relax.op.exp(scaled_x))), + ) + out = bb.emit(relax.op.add(negative, relax.op.nn.relu(x))) + return out if scale == 1 else bb.emit(relax.op.multiply(out, relax.const(scale, dtype))) + def _hardtanh(self, node: fx.Node) -> relax.Expr: args = self.retrieve_args(node) x = args[0] diff --git a/tests/python/relax/test_frontend_from_exported_program.py b/tests/python/relax/test_frontend_from_exported_program.py index e4fbfacca826..ea4039a63ad7 100644 --- a/tests/python/relax/test_frontend_from_exported_program.py +++ b/tests/python/relax/test_frontend_from_exported_program.py @@ -236,6 +236,21 @@ def main(input_1: R.Tensor((1, 3), dtype="int32")) -> R.Tuple( verify_model(SqrtIntModel(), example_args_int32, {}, expected_int32) +def test_selu_and_elu_scale_arguments(): + # run_decompositions rewrites selu to aten.elu(x, alpha, scale), so scale must be applied. + class Selu(Module): + def forward(self, x): + return torch.nn.functional.selu(x) + + class EluScaled(Module): + def forward(self, x): + return torch.ops.aten.elu(x, 0.5, 2.0, 3.0) + + example_args = (torch.tensor([[-2.0, -0.5, 0.0, 1.5]], dtype=torch.float32),) + verify_model_numerically(Selu(), example_args, rtol=1e-5, atol=1e-5) + verify_model_numerically(EluScaled(), example_args, rtol=1e-5, atol=1e-5) + + def test_extended_unary_ops(): example_args = (torch.randn(1, 3, 10, 10, dtype=torch.float32),) @@ -706,7 +721,10 @@ def main(input: R.Tensor((1, 3, 10, 10), dtype="float32")) -> R.Tuple( ) lv4: R.Tensor((1, 3, 10, 10), dtype="float32") = R.nn.relu(input) lv5: R.Tensor((1, 3, 10, 10), dtype="float32") = R.add(lv3, lv4) - gv: R.Tuple(R.Tensor((1, 3, 10, 10), dtype="float32")) = (lv5,) + lv6: R.Tensor((1, 3, 10, 10), dtype="float32") = R.multiply( + lv5, R.const(1.0507009873554805, "float32") + ) + gv: R.Tuple(R.Tensor((1, 3, 10, 10), dtype="float32")) = (lv6,) R.output(gv) return gv From dfd32e712ddb132f22ada85b018a83798b3fe7fe Mon Sep 17 00:00:00 2001 From: Arthur031221 Date: Fri, 2 Oct 2026 11:27:38 +0800 Subject: [PATCH 2/2] [Fix][Relax][Frontend][Torch] Select ELU branch using the original input --- .../torch/exported_program_translator.py | 7 +-- .../test_frontend_from_exported_program.py | 63 +++++++++++-------- 2 files changed, 41 insertions(+), 29 deletions(-) diff --git a/python/tvm/relax/frontend/torch/exported_program_translator.py b/python/tvm/relax/frontend/torch/exported_program_translator.py index 8cf230fe6792..29dc6a07e570 100644 --- a/python/tvm/relax/frontend/torch/exported_program_translator.py +++ b/python/tvm/relax/frontend/torch/exported_program_translator.py @@ -88,12 +88,11 @@ def _elu(self, node: fx.Node) -> relax.Expr: scaled_x = x if input_scale != 1: scaled_x = bb.emit(relax.op.multiply(x, relax.const(input_scale, dtype))) - # scale * (ReLU(x) - alpha * ReLU(1 - exp(input_scale * x))) negative = relax.op.multiply( - relax.const(-alpha, dtype), - relax.op.nn.relu(relax.op.subtract(relax.const(1, dtype), relax.op.exp(scaled_x))), + relax.const(alpha, dtype), + relax.op.subtract(relax.op.exp(scaled_x), relax.const(1, dtype)), ) - out = bb.emit(relax.op.add(negative, relax.op.nn.relu(x))) + out = bb.emit(relax.op.where(relax.op.less(x, relax.const(0, dtype)), negative, x)) return out if scale == 1 else bb.emit(relax.op.multiply(out, relax.const(scale, dtype))) def _hardtanh(self, node: fx.Node) -> relax.Expr: diff --git a/tests/python/relax/test_frontend_from_exported_program.py b/tests/python/relax/test_frontend_from_exported_program.py index ea4039a63ad7..6f5c28895d56 100644 --- a/tests/python/relax/test_frontend_from_exported_program.py +++ b/tests/python/relax/test_frontend_from_exported_program.py @@ -236,21 +236,6 @@ def main(input_1: R.Tensor((1, 3), dtype="int32")) -> R.Tuple( verify_model(SqrtIntModel(), example_args_int32, {}, expected_int32) -def test_selu_and_elu_scale_arguments(): - # run_decompositions rewrites selu to aten.elu(x, alpha, scale), so scale must be applied. - class Selu(Module): - def forward(self, x): - return torch.nn.functional.selu(x) - - class EluScaled(Module): - def forward(self, x): - return torch.ops.aten.elu(x, 0.5, 2.0, 3.0) - - example_args = (torch.tensor([[-2.0, -0.5, 0.0, 1.5]], dtype=torch.float32),) - verify_model_numerically(Selu(), example_args, rtol=1e-5, atol=1e-5) - verify_model_numerically(EluScaled(), example_args, rtol=1e-5, atol=1e-5) - - def test_extended_unary_ops(): example_args = (torch.randn(1, 3, 10, 10, dtype=torch.float32),) @@ -711,20 +696,21 @@ def main(input: R.Tensor((1, 3, 10, 10), dtype="float32")) -> R.Tuple( R.Tensor((1, 3, 10, 10), dtype="float32") ): with R.dataflow(): - lv: R.Tensor((1, 3, 10, 10), dtype="float32") = R.exp(input) - lv1: R.Tensor((1, 3, 10, 10), dtype="float32") = R.subtract( - R.const(1.0, "float32"), lv + lv: R.Tensor((1, 3, 10, 10), dtype="bool") = R.less( + input, R.const(0.0, "float32") + ) + lv1: R.Tensor((1, 3, 10, 10), dtype="float32") = R.exp(input) + lv2: R.Tensor((1, 3, 10, 10), dtype="float32") = R.subtract( + lv1, R.const(1.0, "float32") ) - lv2: R.Tensor((1, 3, 10, 10), dtype="float32") = R.nn.relu(lv1) lv3: R.Tensor((1, 3, 10, 10), dtype="float32") = R.multiply( - R.const(-1.6732631921768188, "float32"), lv2 + R.const(1.6732631921768188, "float32"), lv2 ) - lv4: R.Tensor((1, 3, 10, 10), dtype="float32") = R.nn.relu(input) - lv5: R.Tensor((1, 3, 10, 10), dtype="float32") = R.add(lv3, lv4) - lv6: R.Tensor((1, 3, 10, 10), dtype="float32") = R.multiply( - lv5, R.const(1.0507009873554805, "float32") + lv4: R.Tensor((1, 3, 10, 10), dtype="float32") = R.where(lv, lv3, input) + lv5: R.Tensor((1, 3, 10, 10), dtype="float32") = R.multiply( + lv4, R.const(1.0507009873554805, "float32") ) - gv: R.Tuple(R.Tensor((1, 3, 10, 10), dtype="float32")) = (lv6,) + gv: R.Tuple(R.Tensor((1, 3, 10, 10), dtype="float32")) = (lv5,) R.output(gv) return gv @@ -9224,5 +9210,32 @@ def forward(self, theta): tvm.testing.assert_allclose(tvm_output_np, pytorch_output.numpy(), rtol=1e-5, atol=1e-5) +def test_elu_negative_input_scale(): + class EluScaled(Module): + def forward(self, x): + return torch.ops.aten.elu(x, 0.5, 2.0, -1.0) + + @tvm.script.ir_module + class expected: + @R.function + def main(x: R.Tensor((1, 4), dtype="float32")) -> R.Tuple( + R.Tensor((1, 4), dtype="float32") + ): + with R.dataflow(): + lv = R.multiply(x, R.const(-1.0, "float32")) + lv1 = R.less(x, R.const(0.0, "float32")) + lv2 = R.exp(lv) + lv3 = R.subtract(lv2, R.const(1.0, "float32")) + lv4 = R.multiply(R.const(0.5, "float32"), lv3) + lv5 = R.where(lv1, lv4, x) + lv6 = R.multiply(lv5, R.const(2.0, "float32")) + gv = (lv6,) + R.output(gv) + return gv + + example_args = (torch.tensor([[-2.0, -0.5, 0.0, 1.5]]),) + verify_model(EluScaled(), example_args, {}, expected) + + if __name__ == "__main__": tvm.testing.main()