swap_manager = SwapManager()
def pack_to_cpu(tensor: ms.Tensor):
if not base_check_fn(tensor):
return tensor.detach()
if policy_fn is not None:
if policy_fn(tensor) == CheckpointPolicy.MUST_SAVE:
return tensor.detach()
if policy_fn(tensor) != CheckpointPolicy.MUST_SWAP:
raise RuntimeError(f"Swap :set an invalid policy {policy_fn(tensor)}")
group_name = swap_manager.get_current_group_name()
if not group_name:
return tensor.detach()
if not self.add_to_storage:
swap_manager.add_storage(group_name, self.storage)
self.add_to_storage = True
funcname = f"{group_name}::{tensor.shape}"