| 226 | |
| 227 | |
| 228 | class incremental_save: |
| 229 | def __init__(self, name): |
| 230 | self.name = name |
| 231 | self.zipfile = torch._C.PyTorchFileWriter(str(name)) |
| 232 | self.has_saved = False |
| 233 | self.next_key = 0 |
| 234 | |
| 235 | def __enter__(self): |
| 236 | return self |
| 237 | |
| 238 | def store_early(self, tensor): |
| 239 | if isinstance(tensor, torch.Tensor): |
| 240 | return SavingProxyForTensor(tensor, self) |
| 241 | raise TypeError(f'can only store tensors early, not {type(tensor)}') |
| 242 | |
| 243 | def save(self, obj): |
| 244 | if self.has_saved: |
| 245 | raise RuntimeError('have already saved') |
| 246 | # Write the pickle data for `obj` |
| 247 | data_buf = BytesIO() |
| 248 | pickler = IncrementalPyTorchPickler(self, data_buf, protocol=5) |
| 249 | pickler.dump(obj) |
| 250 | data_value = data_buf.getvalue() |
| 251 | self.zipfile.write_record('data.pkl', data_value, len(data_value)) |
| 252 | self.has_saved = True |
| 253 | |
| 254 | def _write_storage_and_return_key(self, storage): |
| 255 | if self.has_saved: |
| 256 | raise RuntimeError('have already saved') |
| 257 | key = self.next_key |
| 258 | self.next_key += 1 |
| 259 | name = f'data/{key}' |
| 260 | if storage.device.type != 'cpu': |
| 261 | storage = storage.cpu() |
| 262 | num_bytes = storage.nbytes() |
| 263 | self.zipfile.write_record(name, storage.data_ptr(), num_bytes) |
| 264 | return key |
| 265 | |
| 266 | def __exit__(self, type, value, traceback): |
| 267 | self.zipfile.write_end_of_file() |
| 268 | |
| 269 | |
| 270 | T = TypeVar('T') |
no outgoing calls
no test coverage detected