| 15 | |
| 16 | |
| 17 | class FuseViewCopyTransform(ExportPass): |
| 18 | _passes_required_after: Set[Type[ExportPass]] = set() |
| 19 | |
| 20 | VIEW_OP = exir_ops.edge.aten.view_copy.default |
| 21 | |
| 22 | UNARY_ELEMENTWISE_OPS = [ |
| 23 | exir_ops.edge.aten.alias_copy.default, |
| 24 | exir_ops.edge.aten.clone.default, |
| 25 | exir_ops.edge.dim_order_ops._clone_dim_order.default, |
| 26 | exir_ops.edge.aten._to_copy.default, |
| 27 | exir_ops.edge.dim_order_ops._to_dim_order_copy.default, |
| 28 | exir_ops.edge.quantized_decomposed.quantize_per_tensor.default, |
| 29 | exir_ops.edge.quantized_decomposed.dequantize_per_tensor.default, |
| 30 | exir_ops.edge.aten.abs.default, |
| 31 | exir_ops.edge.aten.clamp.default, |
| 32 | exir_ops.edge.aten.ceil.default, |
| 33 | exir_ops.edge.aten.floor.default, |
| 34 | exir_ops.edge.aten.neg.default, |
| 35 | exir_ops.edge.aten.relu.default, |
| 36 | exir_ops.edge.aten.round.default, |
| 37 | exir_ops.edge.aten.sigmoid.default, |
| 38 | exir_ops.edge.aten.silu.default, |
| 39 | exir_ops.edge.aten.sqrt.default, |
| 40 | exir_ops.edge.aten.tanh.default, |
| 41 | exir_ops.edge.aten.sign.default, |
| 42 | exir_ops.edge.aten.reciprocal.default, |
| 43 | exir_ops.edge.aten.gelu.default, |
| 44 | exir_ops.edge.aten.rsqrt.default, |
| 45 | exir_ops.edge.aten.exp.default, |
| 46 | exir_ops.edge.aten.log.default, |
| 47 | ] |
| 48 | |
| 49 | def merge_view_copy_chains( |
| 50 | self, graph: torch.fx.Graph |
| 51 | ) -> tuple[torch.fx.Graph, bool]: |
| 52 | """ |
| 53 | Find chains of view_copy nodes and unary elementwise ops and set all |
| 54 | view_copy nodes to have the final shape. The views will then be removed |
| 55 | by the remove_noop_view_copy call. |
| 56 | |
| 57 | Only merges view_copy nodes that are not used by any other nodes. |
| 58 | """ |
| 59 | view_op = self.VIEW_OP |
| 60 | modified = False |
| 61 | ops = self.UNARY_ELEMENTWISE_OPS + [view_op] |
| 62 | for node in graph.nodes: |
| 63 | if node.op == "call_function" and node.target == view_op: |
| 64 | # Find a chain of unary elementwise ops and save all view_copy nodes |
| 65 | end_node = node |
| 66 | view_ops = [node] |
| 67 | while ( |
| 68 | end_node.op == "call_function" |
| 69 | and end_node.target in ops |
| 70 | and len(end_node.users) == 1 |
| 71 | and list(end_node.users)[0].target in ops |
| 72 | ): |
| 73 | end_node = list(end_node.users)[0] |
| 74 | if end_node.target == view_op: |