1111import torch
1212from executorch .backends .arm ._passes .arm_pass import ArmOpTargetedPass
1313from executorch .backends .arm ._passes .arm_pass_utils import refresh_permute_view_meta
14- from executorch .backends .arm ._passes .dim_maps import PermuteMap , same_numel , ViewMap
14+ from executorch .backends .arm ._passes .dim_maps import (
15+ _dim_equals ,
16+ PermuteMap ,
17+ same_numel ,
18+ ViewMap ,
19+ )
1520from executorch .exir .dialects ._ops import ops as exir_ops
1621from executorch .exir .pass_base import ExportPass , PassResult
1722from torch .export .exported_program import ExportedProgram
@@ -142,8 +147,12 @@ class FuseIdenticalInputTransformsPass(ArmOpTargetedPass):
142147 exir_ops .edge .aten .bitwise_xor .Tensor ,
143148 exir_ops .edge .aten .remainder .Tensor ,
144149 }
150+ _NARY_ELEMENTWISE_OPS = {
151+ exir_ops .edge .aten .where .self ,
152+ }
153+ _ELEMENTWISE_OPS = _BINARY_ELEMENTWISE_OPS | _NARY_ELEMENTWISE_OPS
145154
146- target_ops = _BINARY_ELEMENTWISE_OPS | _CONCAT_OPS
155+ target_ops = _ELEMENTWISE_OPS | _CONCAT_OPS
147156
148157 def __init__ (self , exported_program : ExportedProgram | None = None ) -> None :
149158 super ().__init__ ()
@@ -180,31 +189,43 @@ def _sink_identical_input_transforms(self, node: Node) -> bool:
180189 if node .target not in self .target_ops :
181190 return False
182191
183- input_transforms = self . _input_transforms (node )
184- if input_transforms is None :
192+ input_nodes = list (node . all_input_nodes )
193+ if len ( input_nodes ) < 2 :
185194 return False
186195
187196 node_val = node .meta .get ("val" , None )
188197 if node_val is None :
189198 return False
190199
191- transform = input_transforms [0 ]
192- updated_args = self ._updated_node_args (
193- node , transform , node_val , input_transforms
194- )
200+ transforms = [n for n in input_nodes if n .target in self ._TARGETS ]
201+ if not transforms :
202+ return False
203+ transform = transforms [0 ]
204+ if not self ._inputs_share_transform_or_are_layout_invariant (
205+ node , transform , input_nodes
206+ ):
207+ return False
208+
209+ updated_args = self ._updated_node_args (node , transform , node_val , input_nodes )
195210 if updated_args is None :
196211 return False
197212 node_args , node_kwargs , transform_args , node_output_shape = updated_args
198213
199214 # Remove input transforms
200- producers = [n .all_input_nodes [0 ] for n in input_transforms ]
215+ producers = [
216+ n .all_input_nodes [0 ] if n .target in self ._TARGETS else n
217+ for n in input_nodes
218+ ]
201219
202220 node .args = node_args
203221 node .kwargs = node_kwargs
204- for input_transform , producer in zip (input_transforms , producers ):
222+ for input_transform , producer in zip (input_nodes , producers ):
205223 node .replace_input_with (input_transform , producer )
206- for input_transform in dict .fromkeys (input_transforms ):
207- if len (input_transform .users ) == 0 :
224+ for input_transform in dict .fromkeys (input_nodes ):
225+ if (
226+ input_transform .target in self ._TARGETS
227+ and len (input_transform .users ) == 0
228+ ):
208229 node .graph .erase_node (input_transform )
209230
210231 node .meta = copy .copy (node .meta )
@@ -235,27 +256,52 @@ def _new_transform_meta(self, node: Node, transform: Node) -> dict[str, Any]:
235256 return meta
236257
237258 def _updated_node_args (
238- self , node : Node , transform : Node , node_val : Any , input_transforms : list [Node ]
259+ self , node : Node , transform : Node , node_val : Any , input_nodes : list [Node ]
239260 ) -> (
240261 tuple [tuple [Any , ...], dict [str , Any ], tuple [Any , ...], tuple [Any , ...]] | None
241262 ):
242- if not self ._transforms_are_identical (input_transforms ):
243- return None
244- if not self ._transforms_only_used_by_node (node , input_transforms ):
245- return None
246-
247263 if node .target in self ._BINARY_ELEMENTWISE_OPS :
248- return self ._update_node_args_binary (
249- node , transform , node_val , input_transforms
250- )
264+ return self ._update_node_args_binary (node , transform , node_val , input_nodes )
251265
252266 if node .target in self ._CONCAT_OPS :
253- return self ._update_node_args_concat (
254- node , transform , node_val , input_transforms
255- )
267+ return self ._update_node_args_concat (node , transform , node_val , input_nodes )
268+
269+ if node .target in self ._NARY_ELEMENTWISE_OPS :
270+ return self ._update_node_args_binary (node , transform , node_val , input_nodes )
256271
257272 return None
258273
274+ def _inputs_share_transform_or_are_layout_invariant (
275+ self , node : Node , transform : Node , input_nodes : list [Node ]
276+ ) -> bool :
277+ transforms = [n for n in input_nodes if n .target in self ._TARGETS ]
278+ if not self ._transforms_are_identical (transforms ):
279+ return False
280+ if not self ._transforms_only_used_by_node (node , transforms ):
281+ return False
282+ if len (transforms ) == len (input_nodes ):
283+ return True
284+ if node .target not in self ._ELEMENTWISE_OPS :
285+ return False
286+
287+ transform_val = transform .meta .get ("val" )
288+ if not isinstance (transform_val , torch .Tensor ):
289+ return False
290+ rank = len (transform_val .shape )
291+ return all (
292+ input_node in transforms or self .is_layout_invariant (input_node , rank )
293+ for input_node in input_nodes
294+ )
295+
296+ @staticmethod
297+ def is_layout_invariant (node : Node , rank : int ) -> bool :
298+ value = node .meta .get ("val" )
299+ return (
300+ isinstance (value , torch .Tensor )
301+ and len (value .shape ) == rank
302+ and all (_dim_equals (dim , 1 ) for dim in value .shape )
303+ )
304+
259305 def _transforms_are_identical (self , input_transforms : list [Node ]) -> bool :
260306 target = input_transforms [0 ].target
261307 if target not in self ._TARGETS :
@@ -281,7 +327,15 @@ def _transforms_only_used_by_node(
281327
282328 def _update_node_args_binary (self , node , transform , node_val , input_transforms ):
283329 producer_shapes = [
284- tuple (input_node .all_input_nodes [0 ].meta ["val" ].shape )
330+ tuple (
331+ (
332+ input_node .all_input_nodes [0 ]
333+ if input_node .target in self ._TARGETS
334+ else input_node
335+ )
336+ .meta ["val" ]
337+ .shape
338+ )
285339 for input_node in input_transforms
286340 ]
287341
@@ -293,6 +347,9 @@ def _update_node_args_binary(self, node, transform, node_val, input_transforms):
293347 transform_args = (node , * transform .args [1 :])
294348 if transform .target == self ._VIEW_TARGET :
295349 transform_args = (node , list (node_val .shape ))
350+ # Reshaping before an elementwise op can change which dimensions
351+ # broadcast. Sinking is safe only when the broadcast in the source
352+ # layout already has exactly one of the producer shapes.
296353 if node_output_shape not in producer_shapes :
297354 return None
298355
@@ -382,20 +439,3 @@ def _mapped_concat_dim(self, transform: Node, concat_dim: int) -> int | None:
382439 if mapped_dims is None or len (mapped_dims ) != 1 :
383440 return None
384441 return mapped_dims [0 ]
385-
386- def _input_transforms (self , node : Node ) -> list [Node ] | None :
387- if node .target in self ._BINARY_ELEMENTWISE_OPS :
388- input_transforms = list (node .args [:2 ])
389- elif node .target in self ._CONCAT_OPS :
390- if len (node .args ) == 0 or not isinstance (node .args [0 ], Sequence ):
391- return None
392- input_transforms = list (node .args [0 ])
393- else :
394- return None
395-
396- if len (input_transforms ) < 2 or not all (
397- isinstance (n , Node ) for n in input_transforms
398- ):
399- return None
400-
401- return cast (list [Node ], input_transforms )
0 commit comments