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
| 24 | |
| 25 | |
| 26 | class 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 |
no outgoing calls