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

Class GraphBuilder

backends/test/graph_builder.py:26–113  ·  view source on GitHub ↗

Utility class for creating a graph module with user-specified ops. This class allows us to create test graph modules with any ops we want directly, rather than relying on decomposition or passes. Usage: builder = GraphBuilder() # To insert placeholders, use builder.plac

Source from the content-addressed store, hash-verified

24
25
26class GraphBuilder(ExportPass):
27 """Utility class for creating a graph module with user-specified ops.
28
29 This class allows us to create test graph modules with any ops we want
30 directly, rather than relying on decomposition or passes.
31
32 Usage:
33 builder = GraphBuilder()
34 # To insert placeholders, use builder.placeholder.
35 x = builder.placeholder("x", torch.randn(1, 3, 224, 224))
36 # To insert an op, use builder.call_operator.
37 op = builder.call_operator(
38 some_op
39 (x, other_args, ...),
40 )
41 # Insert outputs as a list of ProxyValues using builder.output.
42 builder.output([op])
43 # Get GraphModule from builder.
44 gm = builder.get_graph_module()
45 """
46
47 def __init__(self, fake_tensor_mode: Optional[FakeTensorMode] = None) -> None:
48 self.exporter = ExportPass()
49 self.tracer: ExportPass.ExportTracer = self.ExportTracer(
50 self, torch.fx.graph.CodeGen()
51 )
52 self.fake_tensor_mode: FakeTensorMode = fake_tensor_mode or FakeTensorMode(
53 allow_fallback_kernels=False,
54 allow_non_fake_inputs=True,
55 )
56 self.tracer.fake_tensor_mode = self.fake_tensor_mode
57
58 # This will be called to create nodes in tracer.
59 self.interpreter = torch.fx.Interpreter(
60 torch.fx.GraphModule(torch.nn.Module(), torch.fx.Graph())
61 )
62
63 # pyre-ignore[14]: Inconsistent override.
64 def placeholder(
65 self, target: str, fake_tensor: Union[FakeTensor, torch.Tensor]
66 ) -> ProxyValue:
67 if not isinstance(fake_tensor, FakeTensor):
68 fake_tensor = self.fake_tensor_mode.from_tensor(fake_tensor)
69 logging.debug(f"Creating placeholder {target} => {fake_tensor.shape}")
70 placeholder = super().placeholder(target, fake_tensor, NodeMetadata({}))
71 return placeholder
72
73 # pyre-ignore[14]: Inconsistent override.
74 def output(self, results: list[ProxyValue]) -> ProxyValue:
75 logging.debug(f"Creating outputs {results}")
76 return super().output(results, NodeMetadata({}))
77
78 def get_graph_module(self) -> torch.fx.GraphModule:
79 return torch.fx.GraphModule(self.tracer.root, self.tracer.graph)
80
81 def call_operator(
82 self,
83 op, # pyre-ignore

Calls

no outgoing calls