diff --git a/backends/arm/_passes/__init__.py b/backends/arm/_passes/__init__.py index 695ab2eb06a..9e1d8fceca1 100644 --- a/backends/arm/_passes/__init__.py +++ b/backends/arm/_passes/__init__.py @@ -75,7 +75,6 @@ from .decompose_select import DecomposeSelectPass # noqa from .decompose_select_scatter_pass import DecomposeSelectScatterPass # noqa from .decompose_sign_pass import DecomposeSignPass # noqa -from .decompose_silu_pass import DecomposeSiluPass # noqa from .decompose_sinh_pass import DecomposeSinhPass # noqa from .decompose_softmax_pass import DecomposeSoftmaxPass # noqa from .decompose_softmax_unstable_pass import DecomposeSoftmaxUnstablePass # noqa diff --git a/backends/arm/_passes/arm_pass_manager.py b/backends/arm/_passes/arm_pass_manager.py index 5e4054e2a1d..ab1342f463c 100644 --- a/backends/arm/_passes/arm_pass_manager.py +++ b/backends/arm/_passes/arm_pass_manager.py @@ -76,7 +76,6 @@ DecomposeSelectPass, DecomposeSelectScatterPass, DecomposeSignPass, - DecomposeSiluPass, DecomposeSinhPass, DecomposeSoftmaxPass, DecomposeSoftmaxUnstablePass, @@ -434,7 +433,6 @@ def transform_for_annotation_pipeline(self, graph_module: GraphModule): DecomposeLeakyReLUPass(tfa_pass=True), DecomposeLinalgVectorNormPass(tfa_pass=True), DecomposeSqrtPass(tfa_pass=True), - DecomposeSiluPass(tfa_pass=True), DecomposeAvgPool2dPass(tfa_pass=True), DecomposeSoftmaxUnstablePass(tfa_pass=True), DecomposeSoftmaxPass(tfa_pass=True), diff --git a/backends/arm/_passes/decompose_silu_pass.py b/backends/arm/_passes/decompose_silu_pass.py deleted file mode 100644 index bf302715805..00000000000 --- a/backends/arm/_passes/decompose_silu_pass.py +++ /dev/null @@ -1,41 +0,0 @@ -# Copyright 2025 Arm Limited and/or its affiliates. -# -# This source code is licensed under the BSD-style license found in the -# LICENSE file in the root directory of this source tree. - - -from typing import Set, Type - -import torch -from executorch.backends.arm._passes import ArmPass -from executorch.backends.arm._passes.insert_table_ops import InsertTableOpsPass -from executorch.exir.pass_base import ExportPass - -aten_silu_ops = (torch.ops.aten.silu.default, torch.ops.aten.silu_.default) - - -class DecomposeSiluPass(ArmPass): - """ - This pass decomposes silu into a mul and a sigmoid node. - - Example: - y = silu(a) - Becomes: - x = sigmoid(a) - y = mul(a,x) - """ - - _passes_required_after: Set[Type[ExportPass]] = {InsertTableOpsPass} - - def call_operator(self, op, args, kwargs, meta): - if op not in (aten_silu_ops) or not self.allowed_to_transform(meta): - return super().call_operator(op, args, kwargs, meta) - sigmoid_op = torch.ops.aten.sigmoid.default - mul_op = torch.ops.aten.mul.Tensor - - original = args[0] - sigmoid = super().call_operator(sigmoid_op, (original,), {}, meta, updated=True) - - return super().call_operator( - mul_op, (original, sigmoid), {}, meta, updated=True - ) diff --git a/backends/arm/_passes/insert_table_ops.py b/backends/arm/_passes/insert_table_ops.py index 3daa4b9fcc7..98056d1b946 100644 --- a/backends/arm/_passes/insert_table_ops.py +++ b/backends/arm/_passes/insert_table_ops.py @@ -55,6 +55,7 @@ class TableOps: exir_ops.edge.aten.cosh.default: torch.cosh, exir_ops.edge.aten.acos.default: torch.acos, exir_ops.edge.aten.tan.default: torch.tan, + exir_ops.edge.aten.silu.default: torch.nn.functional.silu, } # Targets that must be treated explicitly diff --git a/backends/arm/operator_support/tosa_profile_supported_op_lists.py b/backends/arm/operator_support/tosa_profile_supported_op_lists.py index b4e0185150b..1bf91b4304a 100644 --- a/backends/arm/operator_support/tosa_profile_supported_op_lists.py +++ b/backends/arm/operator_support/tosa_profile_supported_op_lists.py @@ -123,6 +123,7 @@ exir_ops.edge.aten.copy.default, exir_ops.edge.aten.tan.default, exir_ops.edge.aten.index_put.default, + exir_ops.edge.aten.silu.default, } diff --git a/backends/arm/quantizer/quantization_annotator.py b/backends/arm/quantizer/quantization_annotator.py index 8d380833dd6..183a8423c03 100644 --- a/backends/arm/quantizer/quantization_annotator.py +++ b/backends/arm/quantizer/quantization_annotator.py @@ -388,6 +388,7 @@ def _match_pattern( torch.ops.aten.zeros_like.default, torch.ops.aten.pow.Tensor_Scalar, torch.ops.aten.gelu.default, + torch.ops.aten.silu.default, torch.ops.aten.sinh.default, torch.ops.aten.atan.default, torch.ops.aten.log1p.default, @@ -468,6 +469,7 @@ def _match_pattern( torch.ops.aten.hardtanh_.default, torch.ops.aten.relu.default, torch.ops.aten.relu_.default, + torch.ops.aten.silu_.default, torch.ops.aten.mean.default, torch.ops.aten.mean.dim, torch.ops.aten.permute.default, diff --git a/backends/arm/test/ops/test_silu.py b/backends/arm/test/ops/test_silu.py index 03dea738e9c..436cd91a7c1 100644 --- a/backends/arm/test/ops/test_silu.py +++ b/backends/arm/test/ops/test_silu.py @@ -1,6 +1,6 @@ # Copyright (c) Meta Platforms, Inc. and affiliates. # All rights reserved. -# Copyright 2025 Arm Limited and/or its affiliates. +# Copyright 2025-2026 Arm Limited and/or its affiliates. # # This source code is licensed under the BSD-style license found in the # LICENSE file in the root directory of this source tree. @@ -43,7 +43,6 @@ def forward( aten_op_FP = "torch.ops.aten.silu.default" aten_op_inplace_FP = "torch.ops.aten.silu_.default" - aten_op_INT = ["torch.ops.aten.sigmoid.default", "torch.ops.aten.mul.Tensor"] @common.parametrize("test_data", Silu.test_data) @@ -63,14 +62,22 @@ def test_silu_tosa_FP_inplace(test_data: input_t): @common.parametrize("test_data", Silu.test_data) def test_silu_tosa_INT(test_data: input_t): silu_data = (test_data(), False) - pipeline = TosaPipelineINT[input_t](Silu(), silu_data, Silu.aten_op_INT) + pipeline = TosaPipelineINT[input_t]( + Silu(), + silu_data, + [], + ) pipeline.run() @common.parametrize("test_data", Silu.test_data) def test_silu_tosa_INT_inplace(test_data: input_t): silu_data = (test_data(), True) - pipeline = TosaPipelineINT[input_t](Silu(), silu_data, Silu.aten_op_INT) + pipeline = TosaPipelineINT[input_t]( + Silu(), + silu_data, + [], + ) pipeline.run() @@ -81,7 +88,7 @@ def test_silu_u55_INT(test_data: input_t): pipeline = EthosU55PipelineINT[input_t]( Silu(), silu_data, - Silu.aten_op_INT, + [], ) pipeline.run() @@ -93,7 +100,7 @@ def test_silu_u55_INT_inplace(test_data: input_t): pipeline = EthosU55PipelineINT[input_t]( Silu(), silu_data, - Silu.aten_op_INT, + [], ) pipeline.run() @@ -105,7 +112,7 @@ def test_silu_u85_INT(test_data: input_t): pipeline = EthosU85PipelineINT[input_t]( Silu(), silu_data, - Silu.aten_op_INT, + [], ) pipeline.run() @@ -117,7 +124,7 @@ def test_silu_u85_INT_inplace(test_data: input_t): pipeline = EthosU85PipelineINT[input_t]( Silu(), silu_data, - Silu.aten_op_INT, + [], ) pipeline.run() @@ -155,7 +162,7 @@ def test_silu_vgf_quant(test_data: input_t): pipeline = VgfPipeline[input_t]( Silu(), silu_data, - Silu.aten_op_INT, + [], quantize=True, ) pipeline.run() @@ -168,7 +175,7 @@ def test_silu_vgf_quant_inplace(test_data: input_t): pipeline = VgfPipeline[input_t]( Silu(), silu_data, - Silu.aten_op_INT, + [], quantize=True, ) pipeline.run() diff --git a/backends/arm/tosa/partitioner.py b/backends/arm/tosa/partitioner.py index 62e110983f0..640dcd5761b 100644 --- a/backends/arm/tosa/partitioner.py +++ b/backends/arm/tosa/partitioner.py @@ -369,6 +369,8 @@ def ops_to_not_decompose( # noqa: C901 torch.ops.aten.hardswish.default, torch.ops.aten.linear.default, torch.ops.aten.linspace.default, + torch.ops.aten.silu.default, + torch.ops.aten.silu_.default, } ops_to_not_decompose_if_fp = { torch.ops.aten.eye.default, @@ -382,6 +384,8 @@ def ops_to_not_decompose( # noqa: C901 ops_to_not_decompose_if_integer = { torch.ops.aten.eye.default, torch.ops.aten.linspace.default, + torch.ops.aten.silu.default, + torch.ops.aten.silu_.default, } def filter_fn(node: torch.fx.Node) -> bool: