MCPcopy Create free account
hub / github.com/pytorch/executorch / FuseViewCopyTransform

Class FuseViewCopyTransform

backends/transforms/fuse_view_copy.py:17–118  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

15
16
17class 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:

Callers 1

preprocessMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected