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

Class incremental_save

lit_gpt/utils_old.py:365–404  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

363
364
365class 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
407T = TypeVar("T")

Callers 2

convert_lit_checkpointFunction · 0.90
convert_hf_checkpointFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected