diff --git a/python/tvm/relax/frontend/torch/exported_program_translator.py b/python/tvm/relax/frontend/torch/exported_program_translator.py index 44804e4d2a98..29dc6a07e570 100644 --- a/python/tvm/relax/frontend/torch/exported_program_translator.py +++ b/python/tvm/relax/frontend/torch/exported_program_translator.py @@ -74,6 +74,27 @@ 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))) + negative = relax.op.multiply( + relax.const(alpha, dtype), + relax.op.subtract(relax.op.exp(scaled_x), relax.const(1, dtype)), + ) + 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: 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..6f5c28895d56 100644 --- a/tests/python/relax/test_frontend_from_exported_program.py +++ b/tests/python/relax/test_frontend_from_exported_program.py @@ -696,16 +696,20 @@ 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.where(lv, lv3, input) + lv5: R.Tensor((1, 3, 10, 10), dtype="float32") = R.multiply( + lv4, R.const(1.0507009873554805, "float32") ) - 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,) R.output(gv) return gv @@ -9206,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()