(self, inp: Argument)
| 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}") |
no test coverage detected