Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions deepspeed/runtime/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -4451,6 +4451,8 @@ def load_checkpoint(self,
_, _, free = disk_usage(offload_dir)
logger.info(f"Copying complete! {free / 1e9:,.2f} GB free on target filesystem")
self.optimizer.reset_swap_buffers()
if load_optimizer_states and not load_module_only:
self._load_resident_optimizer_states(load_dir, tag)

if self._optimizer_has_ckpt_event_epilogue():
self.optimizer.checkpoint_event_epilogue()
Expand Down Expand Up @@ -4705,6 +4707,22 @@ def get_sparse_tensor_module_names(original_set, loaded_set, original_parameters

return load_path, client_state

def _load_resident_optimizer_states(self, load_dir, tag):
"""Restore optimizer state that NVMe offload keeps in memory instead of in swap files.

The NVMe restore path above only copies the swap files back, so any state the optimizer
deliberately pins in memory has to come from the regular ZeRO optimizer checkpoint.
"""
if not hasattr(self.optimizer, 'restore_resident_optimizer_states'):
return

zero_sd_list = self._get_all_zero_checkpoints(load_dir, tag)
if zero_sd_list is None:
return

rank = dist.get_rank(group=self.optimizer.dp_process_group)
self.optimizer.restore_resident_optimizer_states(zero_sd_list[rank])

def _load_zero_checkpoint(self, load_dir, tag, load_optimizer_states=True):

load_serial = None
Expand Down
19 changes: 16 additions & 3 deletions deepspeed/runtime/swap_tensor/optimizer_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -207,10 +207,26 @@ def purge_state(self):
for swap_info in self.swap_params_info.values():
swap_info.tensors = [swap_info.tensors[0]]
swap_info.has_state_tensors = False
swap_info.release_unswapped_gradients()

def _writeback_gradients(self, swap_info, parameter, aio_handle):
"""Persist a parameter's gradient partitions, whichever way they are stored."""
if swap_info.swapped_gradients:
param_gradients = swap_info.swapped_gradients.values()
swap_buffers = [parameter.grad.narrow(0, grad.offset, grad.length) for grad in param_gradients]
swap_paths = [grad.path for grad in param_gradients]
swap_out_tensors(aio_handle, swap_buffers, swap_paths)
assert len(swap_buffers) == aio_handle.wait()
if swap_info.unswapped_gradients:
swap_info.write_unswapped_gradients(src_buffer=parameter.grad)

def is_swappable_tensor(self, tensor=None, numel=None):
assert tensor is not None or numel is not None, "Either tensor or numel must be provided"
if tensor is not None:
# Callers can pin an optimizer state in memory (e.g. the Muon momentum buffer under
# save_muon_momentum_buffer_in_memory) by tagging the tensor, which excludes it from swapping.
if not getattr(tensor, "swappable", True) or getattr(tensor, "is_resident", False):
return False
return self.min_aio_bytes <= (tensor.numel() * self.swap_element_size)
return self.min_aio_bytes <= (numel * self.swap_element_size)

Expand Down Expand Up @@ -467,9 +483,6 @@ def _retrieve_unswapped_grad_partitions(self, swap_info, dest_buffer):
self._stop_timer(UNSWAPPED_READ_GRADIENTS)
self._log_timers([UNSWAPPED_READ_GRADIENTS])

# It should be safe to discard unswapped gradient partitions
swap_info.release_unswapped_gradients()

if SWAPPER_DEBUG_MODE:
logger.info(
f'optimizer_retrieve_unswapped_gradients: param={swap_info.param_id} tensor_count={tensor_count} elem_count={num_elem_count}'
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -127,13 +127,7 @@ def writeback_optimizer_state_and_gradients(self, parameter, write_opt_state, wr
self._swap_out_optimizer_state(swap_info)

if write_gradients and swap_info.has_gradients():
param_gradients = swap_info.swapped_gradients.values()
swap_buffers = [parameter.grad.narrow(0, grad.offset, grad.length) for grad in param_gradients]
swap_paths = [grad.path for grad in param_gradients]
swap_out_tensors(self.aio_handle, swap_buffers, swap_paths)
assert len(swap_buffers) == self.aio_handle.wait()
if swap_info.unswapped_gradients:
swap_info.write_unswapped_gradients(src_buffer=parameter.grad)
self._writeback_gradients(swap_info, parameter, self.aio_handle)

self.release_swap_buffers(parameter)

Expand All @@ -149,6 +143,7 @@ def swap_out_optimizer_state(self, parameter, async_swap=False):
self._start_timer(SWAP_OUT_PARAM_TIMER)
self._swap_out_optimizer_state(swap_info)
self.release_swap_buffers(parameter)
swap_info.release_unswapped_gradients()
self._stop_timer(SWAP_OUT_PARAM_TIMER)
self.timer_names.add(SWAP_OUT_PARAM_TIMER)

Expand Down
38 changes: 38 additions & 0 deletions deepspeed/runtime/swap_tensor/pipelined_optimizer_swapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,6 +137,7 @@ def swap_out_optimizer_state(self, parameter, async_swap):

assert self.swap_ops[SYNC_SWAP_IN] is not None
assert not self.swap_ops[SYNC_SWAP_IN].wait_required
self.swap_ops[SYNC_SWAP_IN].param_info.release_unswapped_gradients()
swap_op = self._swap_out_optimizer_state(aio_handle=self.write_aio_handle,
parameter=parameter,
swap_in_op=self.swap_ops[SYNC_SWAP_IN])
Expand All @@ -157,6 +158,43 @@ def swap_out_gradients(self, parameter, gradient_offsets, gradient_tensors):
gradient_tensors=gradient_tensors,
gradient_swapper=self.gradient_swapper)

def writeback_optimizer_state_and_gradients(self, parameter, write_opt_state, write_gradients):
swap_in_op = self.swap_ops[SYNC_SWAP_IN]
assert swap_in_op is not None and swap_in_op.is_parameter(parameter)
param_info = swap_in_op.param_info

if self.swap_ops[ASYNC_SWAP_OUT]:
self._start_timer(ASYNC_SWAP_OUT_STATE_TIMER)
self._complete_swap_out(ASYNC_SWAP_OUT)
self._stop_timer(ASYNC_SWAP_OUT_STATE_TIMER)
self.timer_names.add(ASYNC_SWAP_OUT_STATE_TIMER)

if write_opt_state:
self._start_timer(SWAP_OUT_STATE_TIMER)
swap_op = self._swap_out_optimizer_state(aio_handle=self.write_aio_handle,
parameter=parameter,
swap_in_op=swap_in_op)
self.swap_ops[SYNC_SWAP_OUT] = swap_op
self._complete_swap_out(SYNC_SWAP_OUT)
self._stop_timer(SWAP_OUT_STATE_TIMER)
self.timer_names.add(SWAP_OUT_STATE_TIMER)
else:
self.swap_buffer_manager.free(swap_in_op.allocated_buffers)

if write_gradients and param_info.has_gradients():
self._writeback_gradients(param_info, parameter, self.write_aio_handle)

param_info.release_memory()
self.swap_ops[SYNC_SWAP_IN] = None

def release_swap_buffers(self, parameter):
swap_in_op = self.swap_ops[SYNC_SWAP_IN]
if swap_in_op is not None and swap_in_op.is_parameter(parameter):
param_info = swap_in_op.param_info
param_info.release_memory()
self.swap_buffer_manager.free(swap_in_op.allocated_buffers)
self.swap_ops[SYNC_SWAP_IN] = None

def _complete_swap_out(self, swap_out_type):
self.swap_ops[swap_out_type].wait()
for buffer in self.swap_ops[swap_out_type].state_buffers:
Expand Down
Loading
Loading