SaverBuilder with support for custom save and restore.
| 738 | return io_ops.restore_v2(filename_tensor, names, slices, dtypes) |
| 739 | |
| 740 | class 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 | |
| 785 | def _get_saver_or_default(incremental_save_restore=False): |
| 786 | """Returns the saver from SAVERS collection, or creates a default one. |