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

Method deserialize_input

torch/_export/serde/serialize.py:1318–1379  ·  view source on GitHub ↗
(self, inp: Argument)

Source from the content-addressed store, hash-verified

1316 return tuple(args), kwargs
1317
1318 def deserialize_input(self, inp: Argument) -> Any:
1319 value = inp.value
1320 typ_ = inp.type
1321 if typ_ == "as_none":
1322 # None should converted as None, but is encoded as bool in serialized
1323 # Convert serialized object to torch equivalent
1324 return None
1325 elif typ_ == "as_scalar_type":
1326 return _SERIALIZE_TO_TORCH_DTYPE[value]
1327 elif typ_ == "as_memory_format":
1328 return _SERIALIZE_TO_TORCH_MEMORY_FORMAT[value]
1329 elif typ_ == "as_layout":
1330 return _SERIALIZE_TO_TORCH_LAYOUT[value]
1331 elif typ_ == "as_graph":
1332 assert isinstance(value, GraphArgument)
1333 with self.save_graph_module():
1334 self.deserialize_graph(value.graph)
1335 submodule = torch._export.exported_program._create_graph_module_for_export(self.module, self.graph)
1336 self.module.register_module(value.name, submodule)
1337 return self.graph.create_node(
1338 "get_attr",
1339 value.name,
1340 name=value.name,
1341 )
1342 elif isinstance(value, Device):
1343 return deserialize_device(value)
1344 elif isinstance(value, TensorArgument):
1345 return self.serialized_name_to_node[value.name]
1346 elif isinstance(value, (int, float, bool)):
1347 return value
1348 elif isinstance(value, str):
1349 return str(value)
1350 elif isinstance(value, (SymIntArgument, SymBoolArgument)):
1351 return self.deserialize_sym_argument(value)
1352 elif isinstance(value, list):
1353 if len(value) == 0:
1354 return []
1355 elif isinstance(value[0], TensorArgument):
1356 result = []
1357 for arg in value:
1358 result.append(self.serialized_name_to_node[arg.name])
1359 return result
1360 elif isinstance(value[0], (int, float, bool)):
1361 # convert from serialized.python.types.List to python list
1362 return list(value)
1363 elif isinstance(value[0], (SymIntArgument, SymBoolArgument)):
1364 return [self.deserialize_sym_argument(arg) for arg in value]
1365 elif isinstance(value[0], OptionalTensorArgument):
1366 def deserialize_optional_tensor_args(a):
1367 if a.type == "as_none":
1368 return None
1369 elif a.type == "as_tensor":
1370 return self.serialized_name_to_node[a.value]
1371 else:
1372 raise SerializeError(f"Unhandled argument {inp}")
1373 return list(map(deserialize_optional_tensor_args, value))
1374 else:
1375 raise SerializeError(f"Unhandled argument {inp}")

Callers 4

deserialize_nodeMethod · 0.95
deserialize_inputsMethod · 0.95

Calls 10

save_graph_moduleMethod · 0.95
deserialize_graphMethod · 0.95
isinstanceFunction · 0.85
deserialize_deviceFunction · 0.85
listFunction · 0.85
SerializeErrorClass · 0.85
register_moduleMethod · 0.80
create_nodeMethod · 0.45
appendMethod · 0.45

Tested by

no test coverage detected