(self, serialized_graph: Graph)
| 1102 | raise SerializeError(f"Unable to deserialize output node {output}") |
| 1103 | |
| 1104 | def deserialize_graph(self, serialized_graph: Graph) -> torch.fx.Graph: |
| 1105 | # Handle the tensor metas. |
| 1106 | for name, tensor_value in serialized_graph.tensor_values.items(): |
| 1107 | meta_val = self.deserialize_tensor_meta(tensor_value, self.fake_tensor_mode) |
| 1108 | self.serialized_name_to_meta[name] = meta_val |
| 1109 | |
| 1110 | for name, sym_int_value in serialized_graph.sym_int_values.items(): |
| 1111 | self.serialized_name_to_meta[name] = self.deserialize_sym_int(sym_int_value) |
| 1112 | |
| 1113 | for name, sym_bool_value in serialized_graph.sym_bool_values.items(): |
| 1114 | self.serialized_name_to_meta[name] = self.deserialize_sym_bool(sym_bool_value) |
| 1115 | |
| 1116 | # Inputs: convert to placeholder nodes in FX. |
| 1117 | for input in serialized_graph.inputs: |
| 1118 | placeholder_node = self.graph.placeholder(input.as_tensor.name) |
| 1119 | self.sync_fx_node(input.as_tensor.name, placeholder_node) |
| 1120 | |
| 1121 | # Nodes: convert to call_function nodes. |
| 1122 | for serialized_node in serialized_graph.nodes: |
| 1123 | try: |
| 1124 | target = self.deserialize_operator(serialized_node.target) |
| 1125 | self.deserialize_node(serialized_node, target) |
| 1126 | |
| 1127 | except Exception as e: |
| 1128 | raise SerializeError(f"Failed deserializing node {serialized_node}") from e |
| 1129 | |
| 1130 | # Outputs: convert to a single `output` node. |
| 1131 | outputs = [] |
| 1132 | for output in serialized_graph.outputs: |
| 1133 | outputs.append(self.deserialize_graph_output(output)) |
| 1134 | |
| 1135 | if serialized_graph.is_single_tensor_return: |
| 1136 | assert len(outputs) == 1 |
| 1137 | outputs = outputs[0] # type: ignore[assignment] |
| 1138 | else: |
| 1139 | outputs = tuple(outputs) # type: ignore[assignment] |
| 1140 | |
| 1141 | output_node = self.graph.output(outputs) |
| 1142 | |
| 1143 | if serialized_graph.is_single_tensor_return: |
| 1144 | output_node.meta["val"] = output_node.args[0].meta["val"] |
| 1145 | else: |
| 1146 | output_node.meta["val"] = tuple( |
| 1147 | arg.meta["val"] for arg in output_node.args[0] |
| 1148 | ) |
| 1149 | |
| 1150 | return self.graph |
| 1151 | |
| 1152 | def deserialize_node(self, serialized_node: Node, target: Callable) -> None: |
| 1153 | if target.__module__ == "_operator": # TODO(zhxchen17) Follow up on this. |
no test coverage detected