MCPcopy Create free account
hub / github.com/Pints-AI/1.5-Pints / incremental_save

Class incremental_save

lit_gpt/utils.py:228–267  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

226
227
228class 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
270T = TypeVar('T')

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected