MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / PartialRestoreSaverBuilder

Class PartialRestoreSaverBuilder

tensorflow/python/training/saver.py:740–783  ·  view source on GitHub ↗

SaverBuilder with support for custom save and restore.

Source from the content-addressed store, hash-verified

738 return io_ops.restore_v2(filename_tensor, names, slices, dtypes)
739
740class PartialRestoreSaverBuilder(BulkSaverBuilder):
741 """SaverBuilder with support for custom save and restore.
742 """
743 def __init__(self, *args, **kwargs):
744 self._ckpt_dir_or_file = kwargs.pop("checkpoint_dir", True)
745 super(PartialRestoreSaverBuilder, self).__init__(*args, **kwargs)
746
747 def restore_op(self, filename_tensor, saveable, preferred_shard):
748 """Create ops to restore 'saveable'.
749 """
750 if hasattr(saveable, "load_from_checkpoint"):
751 return saveable.load_from_checkpoint(
752 self._ckpt_dir_or_file, filename_tensor, preferred_shard)
753 return super(PartialRestoreSaverBuilder, self).restore_op(
754 filename_tensor, saveable, preferred_shard)
755
756 def bulk_restore(self, filename_tensor, saveables, preferred_shard,
757 restore_sequentially):
758 """Restore all tensors contained in saveables.
759 """
760 custom_saveables = []
761 other_saveables = []
762 offset = 0
763 for saveable in saveables:
764 if hasattr(saveable, "load_from_checkpoint"):
765 custom_saveables.append(
766 (offset,
767 saveable.load_from_checkpoint(
768 self._ckpt_dir_or_file, filename_tensor, preferred_shard)))
769 else:
770 other_saveables.append(saveable)
771 offset += len(saveable.specs)
772 other_tensors = super(PartialRestoreSaverBuilder, self).bulk_restore(
773 filename_tensor, other_saveables, preferred_shard, restore_sequentially)
774 all_tensors = []
775 prev_offset = 0
776 for offset, tensors in custom_saveables:
777 if offset > prev_offset:
778 all_tensors.extend(other_tensors[prev_offset:offset])
779 prev_offset = offset
780 all_tensors.extend(tensors)
781 if len(other_tensors) > prev_offset:
782 all_tensors.extend(other_tensors[prev_offset: len(other_tensors)])
783 return all_tensors
784
785def _get_saver_or_default(incremental_save_restore=False):
786 """Returns the saver from SAVERS collection, or creates a default one.

Callers 1

restore_fnMethod · 0.90

Calls

no outgoing calls

Tested by 1

restore_fnMethod · 0.72