| 363 | |
| 364 | |
| 365 | class incremental_save: |
| 366 | def __init__(self, name): |
| 367 | self.name = name |
| 368 | self.zipfile = torch._C.PyTorchFileWriter(str(name)) |
| 369 | self.has_saved = False |
| 370 | self.next_key = 0 |
| 371 | |
| 372 | def __enter__(self): |
| 373 | return self |
| 374 | |
| 375 | def store_early(self, tensor): |
| 376 | if isinstance(tensor, torch.Tensor): |
| 377 | return SavingProxyForTensor(tensor, self) |
| 378 | raise TypeError(f"can only store tensors early, not {type(tensor)}") |
| 379 | |
| 380 | def save(self, obj): |
| 381 | if self.has_saved: |
| 382 | raise RuntimeError("have already saved") |
| 383 | # Write the pickle data for `obj` |
| 384 | data_buf = BytesIO() |
| 385 | pickler = IncrementalPyTorchPickler(self, data_buf, protocol=5) |
| 386 | pickler.dump(obj) |
| 387 | data_value = data_buf.getvalue() |
| 388 | self.zipfile.write_record("data.pkl", data_value, len(data_value)) |
| 389 | self.has_saved = True |
| 390 | |
| 391 | def _write_storage_and_return_key(self, storage): |
| 392 | if self.has_saved: |
| 393 | raise RuntimeError("have already saved") |
| 394 | key = self.next_key |
| 395 | self.next_key += 1 |
| 396 | name = f"data/{key}" |
| 397 | if storage.device.type != "cpu": |
| 398 | storage = storage.cpu() |
| 399 | num_bytes = storage.nbytes() |
| 400 | self.zipfile.write_record(name, storage.data_ptr(), num_bytes) |
| 401 | return key |
| 402 | |
| 403 | def __exit__(self, type, value, traceback): |
| 404 | self.zipfile.write_end_of_file() |
| 405 | |
| 406 | |
| 407 | T = TypeVar("T") |
no outgoing calls
no test coverage detected