MCPcopy Create free account
hub / github.com/pytorch/pytorch / deserialize_graph

Method deserialize_graph

torch/_export/serde/serialize.py:1104–1150  ·  view source on GitHub ↗
(self, serialized_graph: Graph)

Source from the content-addressed store, hash-verified

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.

Callers 2

deserializeMethod · 0.95
deserialize_inputMethod · 0.95

Calls 12

deserialize_sym_intMethod · 0.95
deserialize_sym_boolMethod · 0.95
sync_fx_nodeMethod · 0.95
deserialize_operatorMethod · 0.95
deserialize_nodeMethod · 0.95
SerializeErrorClass · 0.85
itemsMethod · 0.45
placeholderMethod · 0.45
appendMethod · 0.45
outputMethod · 0.45

Tested by

no test coverage detected