diff --git a/deepspeed/runtime/engine.py b/deepspeed/runtime/engine.py index 86918bd71c5a..1934d8887dd5 100644 --- a/deepspeed/runtime/engine.py +++ b/deepspeed/runtime/engine.py @@ -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() @@ -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 diff --git a/deepspeed/runtime/swap_tensor/optimizer_utils.py b/deepspeed/runtime/swap_tensor/optimizer_utils.py index 191a85414b66..847d558e376e 100644 --- a/deepspeed/runtime/swap_tensor/optimizer_utils.py +++ b/deepspeed/runtime/swap_tensor/optimizer_utils.py @@ -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) @@ -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}' diff --git a/deepspeed/runtime/swap_tensor/partitioned_optimizer_swapper.py b/deepspeed/runtime/swap_tensor/partitioned_optimizer_swapper.py index a8d4a64b817f..46f34a684af3 100644 --- a/deepspeed/runtime/swap_tensor/partitioned_optimizer_swapper.py +++ b/deepspeed/runtime/swap_tensor/partitioned_optimizer_swapper.py @@ -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) @@ -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) diff --git a/deepspeed/runtime/swap_tensor/pipelined_optimizer_swapper.py b/deepspeed/runtime/swap_tensor/pipelined_optimizer_swapper.py index 9f76032cbeb6..b00f6a75bc21 100644 --- a/deepspeed/runtime/swap_tensor/pipelined_optimizer_swapper.py +++ b/deepspeed/runtime/swap_tensor/pipelined_optimizer_swapper.py @@ -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]) @@ -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: diff --git a/deepspeed/runtime/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index 0afd8c1b9f89..31d44ef8824c 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -237,6 +237,9 @@ def __init__( self.dtype = self.optimizer.param_groups[0]['params'][0].dtype self.gradient_accumulation_dtype = gradient_accumulation_dtype self._global_grad_norm = 0. + self._muon_allgather_buffers = collections.OrderedDict() + self._muon_allgather_buffer_bytes = 0 + self._muon_allgather_max_cached_bytes = 256 * 1024 * 1024 self.custom_loss_scaler = False self.external_loss_scale = None @@ -526,6 +529,7 @@ def _enforce_optimizer_offload(): def destroy(self): self.parameter_offload.destroy() + self._clear_muon_allgather_buffers() for hook in self._grad_acc_hooks: hook.remove() for hook in self._leaf_module_hooks: @@ -1020,17 +1024,83 @@ def _get_sub_group_partitions(self, sub_group_id): def _create_momentum_buffer(self, num_elements, i, ds_id): if self.use_muon and self.sub_groups_using_muon[i]: + # Momentum is an optimizer state that persists across steps, so it must keep the + # master (fp32) precision. Storing it in the reduced communication dtype would + # round the accumulator on every step and drift away from ZeRO-1/2 behavior. unpinned_fp32_buffer_momentum = torch.zeros(num_elements, device=self.device, - dtype=self.communication_data_type) + dtype=self.master_weights_and_grads_dtype) unpinned_fp32_buffer_momentum.requires_grad = False + if self.save_muon_momentum_buffer_in_memory: + unpinned_fp32_buffer_momentum.swappable = False + unpinned_fp32_buffer_momentum.is_resident = True + self.muon_momentum_buffer_partitioned_groups_flat[i] = unpinned_fp32_buffer_momentum + self.muon_momentum_buffer_partitioned_groups_flat[i].ds_id = ds_id if self.fp32_partitioned_groups_flat[i] not in self.optimizer.state: self.optimizer.state[self.fp32_partitioned_groups_flat[i]] = {} self.optimizer.state[ self.fp32_partitioned_groups_flat[i]]["momentum_buffer"] = unpinned_fp32_buffer_momentum - if self.save_muon_momentum_buffer_in_memory: - self.muon_momentum_buffer_partitioned_groups_flat[i] = unpinned_fp32_buffer_momentum - self.muon_momentum_buffer_partitioned_groups_flat[i].ds_id = ds_id + + def _adopt_restored_muon_momentum(self, sub_group_id, restored_momentum): + """Copy a checkpointed momentum into the resident buffer and re-bind the optimizer state. + + The resident buffer is deliberately excluded from swapping and is cached outside of + ``optimizer.state``, so both references must keep pointing at the same tensor. + """ + resident_momentum = self.muon_momentum_buffer_partitioned_groups_flat[sub_group_id] + if restored_momentum is resident_momentum: + return + if restored_momentum.numel() != resident_momentum.numel(): + raise RuntimeError("Muon momentum checkpoint size mismatch for subgroup " + f"{sub_group_id}: got {restored_momentum.numel()} elements, " + f"expected {resident_momentum.numel()}.") + resident_momentum.data.copy_(restored_momentum.data) + fp32_param = self.fp32_partitioned_groups_flat[sub_group_id] + self.optimizer.state.setdefault(fp32_param, {})["momentum_buffer"] = resident_momentum + + def _restore_muon_momentum_residency(self): + """Re-bind the in-memory Muon momentum cache to the freshly loaded optimizer state. + + ``Optimizer.load_state_dict()`` replaces the state tensors with new objects, so the + resident cache would otherwise keep serving the pre-load (usually zero) momentum and + overwrite the restored values on the very next step. + """ + if not (self.use_muon and self.save_muon_momentum_buffer_in_memory): + return + + for sub_group_id in self.muon_momentum_buffer_partitioned_groups_flat: + state = self.optimizer.state.get(self.fp32_partitioned_groups_flat[sub_group_id]) + restored_momentum = None if state is None else state.get("momentum_buffer") + if restored_momentum is not None: + self._adopt_restored_muon_momentum(sub_group_id, restored_momentum) + + def restore_resident_optimizer_states(self, state_dict): + """Restore optimizer state that is pinned in memory rather than swapped to NVMe. + + NVMe optimizer offload rebuilds its state by copying the swap files back, bypassing + ``Optimizer.load_state_dict()`` entirely. A resident Muon momentum buffer never reaches + those files, so it has to be pulled out of the saved optimizer state dict by hand. + """ + if not (self.use_muon and self.save_muon_momentum_buffer_in_memory): + return + + saved_state = state_dict[OPTIMIZER_STATE_DICT]["state"] + # Mirror how torch numbers parameters when packing an optimizer state dict. + self._set_fp32_optimizer_param_groups() + try: + saved_index_of_param = {} + for group in self.optimizer.param_groups: + for param in group["params"]: + saved_index_of_param[id(param)] = len(saved_index_of_param) + finally: + self._clear_fp32_optimizer_param_groups() + + for sub_group_id in self.muon_momentum_buffer_partitioned_groups_flat: + fp32_param = self.fp32_partitioned_groups_flat[sub_group_id] + saved_index = saved_index_of_param[id(fp32_param)] + restored_momentum = saved_state.get(saved_index, {}).get("momentum_buffer") + if restored_momentum is not None: + self._adopt_restored_muon_momentum(sub_group_id, restored_momentum) def _create_fp32_partitions(self): cpu_memory_usage = 0 @@ -1202,7 +1272,7 @@ def step_with_gradscaler(optimizer): if self.offload_optimizer: cur_device = self.subgroup_to_device[sub_group_id] - if cur_device == 'cpu': + if cur_device == 'cpu' or (self.use_muon and self.sub_groups_using_muon[sub_group_id]): self.optimizer.param_groups[param_group_id]['params'] = [fp32_param] step_with_gradscaler(self.optimizer) self.optimizer.param_groups[param_group_id]['params'] = [] @@ -1629,7 +1699,7 @@ def _apply_distributed_muon_update(self, communication_data_type: torch.dtype, b Returns: None """ - if not self.use_muon: + if not self.use_muon or self.offload_optimizer: return params_by_group = {} @@ -1652,31 +1722,28 @@ def _apply_distributed_muon_update(self, communication_data_type: torch.dtype, b if not params: continue - momentum_buffer = [] - if self._swappable_optimizer_subgroup(i) and not self.save_muon_momentum_buffer_in_memory: + fp32_param = self.fp32_partitioned_groups_flat[i] + swapped_momentum = self._swappable_optimizer_subgroup(i) and not self.save_muon_momentum_buffer_in_memory + if swapped_momentum: # swap-in once, keep resident through update + writeback - self.optimizer_swapper.swap_in_optimizer_state(parameter=self.fp32_partitioned_groups_flat[i]) - if "momentum_buffer" not in self.optimizer.state.get(self.fp32_partitioned_groups_flat[i], {}): - self._create_momentum_buffer(self.fp16_partitioned_groups_flat_numel[i], i, - self.fp32_partitioned_groups_flat[i].ds_id) - state_buffer = self.optimizer.state[self.fp32_partitioned_groups_flat[i]]["momentum_buffer"] - for param, dest_offset, _ in group_items: - momentum_buffer.append(state_buffer.narrow(0, dest_offset, param.partition_numel()).clone()) - elif self.save_muon_momentum_buffer_in_memory: + self.optimizer_swapper.swap_in_optimizer_state(parameter=fp32_param) + + if self.save_muon_momentum_buffer_in_memory: state_buffer = self.muon_momentum_buffer_partitioned_groups_flat[i] - for param, dest_offset, _ in group_items: - momentum_buffer.append(state_buffer.narrow(0, dest_offset, param.partition_numel()).clone()) else: - # Non-swappable optimizer (GPU/CPU): momentum buffer lives in optimizer state - if "momentum_buffer" not in self.optimizer.state.get(self.fp32_partitioned_groups_flat[i], {}): - self._create_momentum_buffer(self.fp16_partitioned_groups_flat_numel[i], i, - self.fp32_partitioned_groups_flat[i].ds_id) - state_buffer = self.optimizer.state[self.fp32_partitioned_groups_flat[i]]["momentum_buffer"] - for param, dest_offset, _ in group_items: - momentum_buffer.append(state_buffer.narrow(0, dest_offset, param.partition_numel()).clone()) + if "momentum_buffer" not in self.optimizer.state.get(fp32_param, {}): + self._create_momentum_buffer(self.fp16_partitioned_groups_flat_numel[i], i, fp32_param.ds_id) + state_buffer = self.optimizer.state[fp32_param]["momentum_buffer"] + + momentum_buffer = [ + state_buffer.narrow(0, dest_offset, param.partition_numel()).clone() + for param, dest_offset, _ in group_items + ] - gathered_params_momentums = self._partitioned_buffers_all_gather(params, momentum_buffer, - communication_data_type) + # Momentum is a persistent optimizer state, so gather and update it in the + # master (fp32) dtype instead of the reduced communication dtype. + muon_dtype = self.master_weights_and_grads_dtype + gathered_params_momentums = self._partitioned_buffers_all_gather(params, momentum_buffer, muon_dtype) process_group = self._get_sub_group_process_group(i) world_sz = dist.get_world_size(process_group) @@ -1692,7 +1759,12 @@ def _apply_distributed_muon_update(self, communication_data_type: torch.dtype, b param = params[base_i + rank] g = param.grad m = gathered_momentums_pad[base_i + rank] - update = muon_update(g, m, beta=self.muon_beta, ns_method=getattr(self, 'muon_ns_method', 'gram')) + # Promote the gradient so momentum tracking and Newton-Schulz run in fp32. + fp32_grad = g.to(muon_dtype) + update = muon_update(fp32_grad, + m, + beta=self.muon_beta, + ns_method=getattr(self, 'muon_ns_method', 'gram')) g.data.copy_(update, non_blocking=False) grad_handle = dist.all_gather(grads_pad[base_i:base_i + world_sz], grads_pad[base_i + rank], @@ -1713,28 +1785,18 @@ def _apply_distributed_muon_update(self, communication_data_type: torch.dtype, b start_offset = rank * chunk_sz end_offset = start_offset + chunk_sz if end_offset > param.grad.numel(): - buffer_to_update = torch.zeros(chunk_sz, - device=param.grad.device, - dtype=self.gradient_accumulation_dtype) + buffer_to_update = torch.zeros(chunk_sz, device=param.grad.device, dtype=gathered_momentum.dtype) buffer_to_update[:param.grad.numel() - start_offset] = gathered_momentum.view(-1).data[start_offset:param.grad.numel()] else: buffer_to_update = gathered_momentum.view(-1).data[start_offset:end_offset] - if self._swappable_optimizer_subgroup(i) and not self.save_muon_momentum_buffer_in_memory: - self.optimizer.state[self.fp32_partitioned_groups_flat[i]]["momentum_buffer"].narrow( - 0, dest_offset, param.partition_numel()).data.copy_(buffer_to_update, non_blocking=False) - elif self.save_muon_momentum_buffer_in_memory: - self.muon_momentum_buffer_partitioned_groups_flat[i].narrow( - 0, dest_offset, param.partition_numel()).data.copy_(buffer_to_update, non_blocking=False) - # update the momentum buffer in the optimizer state - self.optimizer.state[self.fp32_partitioned_groups_flat[i]][ - "momentum_buffer"] = self.muon_momentum_buffer_partitioned_groups_flat[i] - else: - # Non-swappable optimizer (GPU/CPU): write directly to optimizer state - self.optimizer.state[self.fp32_partitioned_groups_flat[i]]["momentum_buffer"].narrow( - 0, dest_offset, param.partition_numel()).data.copy_(buffer_to_update, non_blocking=False) - if self._swappable_optimizer_subgroup(i) and not self.save_muon_momentum_buffer_in_memory: - self.optimizer_swapper.swap_out_optimizer_state(parameter=self.fp32_partitioned_groups_flat[i]) + state_buffer.narrow(0, dest_offset, param.partition_numel()).data.copy_(buffer_to_update, + non_blocking=False) + if self.save_muon_momentum_buffer_in_memory: + # The resident buffer is not owned by the swapper, so re-publish it as the state + self.optimizer.state[fp32_param]["momentum_buffer"] = state_buffer + if swapped_momentum: + self.optimizer_swapper.swap_out_optimizer_state(parameter=fp32_param) for handle in grad_handles: handle.wait() for param, _, params_size_offset in group_items: @@ -1937,57 +1999,108 @@ def partition_grads(self, params_to_release: List[Parameter], grad_partitions: L self._swap_out_offload_fp32_gradients(offload_fp32_gradients, offload_fp32_offsets) return buffers - def _partitioned_buffers_all_gather(self, params: List[Parameter], buffers_to_allgather: List[Tensor], - communication_data_type: torch.dtype): + def _partitioned_buffers_all_gather(self, + params: List[Parameter], + buffers_to_allgather: List[Tensor], + communication_data_type: torch.dtype, + additional_buffers_to_allgather: List[Tensor] = None): """ Allgather the partitioned buffers of the parameters to the global buffer. Args: params: List[Parameter] buffers_to_allgather: List[Tensor] communication_data_type: torch.dtype + additional_buffers_to_allgather: Optional second buffer list to gather in the same collective. Returns: - List[Tensor] + List[Tensor], or one list per buffer list when an additional list is provided. """ - assert len(params) == len(buffers_to_allgather), "params and buffers_to_allgather must have the same length" - assert all(param.partition_numel() == buffer.numel() - for param, - buffer in zip(params, buffers_to_allgather)), \ + buffer_lists = [buffers_to_allgather] + if additional_buffers_to_allgather is not None: + buffer_lists.append(additional_buffers_to_allgather) + assert all(len(params) == len(buffers) for buffers in buffer_lists), \ + "params and buffers_to_allgather must have the same length" + assert all(param.partition_numel() == buffer.numel() for buffers in buffer_lists for param, buffer in zip(params, + buffers)), \ "params and buffers_to_allgather must have the same numel" self._assert_same_partition_group(params) process_group = self._get_param_partition_group(params[0]) partition_count = dist.get_world_size(group=process_group) - coalesced_buffer = instrument_w_nvtx(torch.cat)(buffers_to_allgather) - buffer_numel = coalesced_buffer.numel() - reduce_buffer = torch.empty(partition_count * buffer_numel, - dtype=communication_data_type, - device=params[0].device) - rearrange_buffer = torch.empty(partition_count * buffer_numel, - dtype=communication_data_type, - device=params[0].device) + if partition_count == 1: + outputs = [] + for buffers in buffer_lists: + list_outputs = [] + for param, buffer in zip(params, buffers): + full_numel = getattr(param, "ds_numel", param.numel()) + full_shape = getattr(param, "ds_shape", param.shape) + list_outputs.append( + buffer.to(communication_data_type).view(-1).narrow(0, 0, full_numel).view(full_shape)) + outputs.append(list_outputs) + return outputs[0] if additional_buffers_to_allgather is None else outputs + + buffer_numels = [sum(buffer.numel() for buffer in buffers) for buffers in buffer_lists] + local_numel = sum(buffer_numels) + output_numel = sum(numel * partition_count for numel in buffer_numels) + device = buffers_to_allgather[0].device + cache_key = (id(process_group), partition_count, local_numel, output_numel, communication_data_type, device) + cache = self._muon_allgather_buffers.pop(cache_key, None) + if cache is None: + reduce_buffer = torch.empty(partition_count * local_numel, dtype=communication_data_type, device=device) + rearrange_buffer = torch.empty(output_numel, dtype=communication_data_type, device=device) + local_buffer = torch.empty(local_numel, dtype=communication_data_type, device=device) + cache_bytes = (local_buffer.numel() + reduce_buffer.numel() + rearrange_buffer.numel()) * \ + communication_data_type.itemsize + if cache_bytes <= self._muon_allgather_max_cached_bytes: + while (self._muon_allgather_buffers + and self._muon_allgather_buffer_bytes + cache_bytes > self._muon_allgather_max_cached_bytes): + _, evicted = self._muon_allgather_buffers.popitem(last=False) + self._muon_allgather_buffer_bytes -= evicted[3] + self._muon_allgather_buffers[cache_key] = (local_buffer, reduce_buffer, rearrange_buffer, cache_bytes) + self._muon_allgather_buffer_bytes += cache_bytes + else: + local_buffer, reduce_buffer, rearrange_buffer, cache_bytes = cache + self._muon_allgather_buffers[cache_key] = cache + + buffer_offsets = [0] + for buffer_numel in buffer_numels: + buffer_offsets.append(buffer_offsets[-1] + buffer_numel) + for list_idx, buffers in enumerate(buffer_lists): + offset = buffer_offsets[list_idx] + copy_offset = offset + for buffer in buffers: + numel = buffer.numel() + local_buffer.narrow(0, copy_offset, numel).copy_(buffer, non_blocking=True) + copy_offset += numel my_rank = dist.get_rank(group=process_group) - partition = reduce_buffer.narrow(0, buffer_numel * my_rank, buffer_numel) - partition.data.copy_(coalesced_buffer.data, non_blocking=False) + partition = reduce_buffer.narrow(0, local_numel * my_rank, local_numel) + partition.copy_(local_buffer, non_blocking=False) dist.all_gather_into_tensor(reduce_buffer, partition, group=process_group) - param_partition_offsets = [0] + outputs = [] rearranged_offset = 0 - for idx, param in enumerate(params): - param_partition_offsets.append(param_partition_offsets[idx] + param.partition_numel()) - for idx, param in enumerate(params): - num_elements = param.partition_numel() - for partition_idx in range(partition_count): - sliced = reduce_buffer.narrow(0, buffer_numel * partition_idx + param_partition_offsets[idx], - num_elements) - rearrange_buffer.narrow(0, rearranged_offset, num_elements).copy_(sliced.data, non_blocking=False) - rearranged_offset += num_elements - param_full_offsets = [0] - for idx, param in enumerate(params): - # the offset is the sum of the numel of all the partitions of the parameter including padding - param_full_offsets.append(param_full_offsets[idx] + buffers_to_allgather[idx].numel() * partition_count) - output = [] - for idx, param in enumerate(params): - output.append(rearrange_buffer.narrow(0, param_full_offsets[idx], param.ds_numel).view(param.ds_shape)) - return output + for list_idx, buffers in enumerate(buffer_lists): + param_partition_offsets = [0] + for buffer in buffers: + param_partition_offsets.append(param_partition_offsets[-1] + buffer.numel()) + list_outputs = [] + for idx, param in enumerate(params): + num_elements = buffers[idx].numel() + for partition_idx in range(partition_count): + source_offset = (local_numel * partition_idx + buffer_offsets[list_idx] + + param_partition_offsets[idx]) + sliced = reduce_buffer.narrow(0, source_offset, num_elements) + rearrange_buffer.narrow(0, rearranged_offset, num_elements).copy_(sliced, non_blocking=False) + rearranged_offset += num_elements + full_numel = getattr(param, "ds_numel", param.numel()) + full_shape = getattr(param, "ds_shape", param.shape) + list_outputs.append( + rearrange_buffer.narrow(0, rearranged_offset - num_elements * partition_count, + full_numel).view(full_shape)) + outputs.append(list_outputs) + return outputs[0] if additional_buffers_to_allgather is None else outputs + + def _clear_muon_allgather_buffers(self): + self._muon_allgather_buffers.clear() + self._muon_allgather_buffer_bytes = 0 def reduce_ready_partitions_and_remove_grads(self, param): if self._coalesce_grad_reduction: @@ -2393,6 +2506,8 @@ def _pre_step(self): @instrument_w_nvtx def _get_norm_groups(self): + if self.offload_optimizer: + self._apply_muon_updates_cpu_offload() norm_groups = [] for i, group in enumerate(self.fp16_groups): if self.offload_optimizer: @@ -2401,6 +2516,99 @@ def _get_norm_groups(self): norm_groups.append(self.get_grad_norm_direct(self.averaged_gradients[i], self.fp16_groups[i])) return norm_groups + @instrument_w_nvtx + @torch.no_grad() + def _apply_muon_updates_cpu_offload(self): + """Orthogonalize full logical gradients before clipping CPU-offloaded updates.""" + if not self.use_muon: + return + + accelerator_device = get_accelerator().current_device_name() + for sub_group_id, params in enumerate(self.fp16_groups): + muon_params = [param for param in params if getattr(param, "use_muon", False)] + if not muon_params: + continue + + if self._swappable_optimizer_subgroup(sub_group_id): + self._optimizer_states_and_gradient_swap_in(sub_group_id) + + fp32_param = self.fp32_partitioned_groups_flat[sub_group_id] + subgroup_numel = int(self.fp16_partitioned_groups_flat_numel[sub_group_id]) + if self.save_muon_momentum_buffer_in_memory: + momentum = self.muon_momentum_buffer_partitioned_groups_flat.get(sub_group_id) + momentum_was_created = momentum is None or momentum.numel() != subgroup_numel + if momentum_was_created: + self._create_momentum_buffer(subgroup_numel, sub_group_id, fp32_param.ds_id) + momentum = self.muon_momentum_buffer_partitioned_groups_flat[sub_group_id] + else: + state = self.optimizer.state.setdefault(fp32_param, {}) + momentum = state.get("momentum_buffer") + momentum_was_created = momentum is None or momentum.numel() != subgroup_numel + if momentum_was_created: + # A newly allocated state is zero on every rank, so it needs no all-gather. + self._create_momentum_buffer(subgroup_numel, sub_group_id, fp32_param.ds_id) + momentum = state["momentum_buffer"] + + local_grad_parts = [] + local_momentum_parts = [] + for param in muon_params: + _, dest_offset, _ = self.grad_position[self.get_param_id(param)] + numel = param.partition_numel() + local_grad_parts.append(fp32_param.grad.narrow(0, dest_offset, numel).to(accelerator_device)) + if not momentum_was_created: + local_momentum_parts.append(momentum.narrow(0, dest_offset, numel).to(accelerator_device)) + + # Gather and run Muon in the master (fp32) dtype so momentum tracking and + # Newton-Schulz never see a half-precision round trip. + muon_dtype = self.master_weights_and_grads_dtype + if momentum_was_created: + full_grads = self._partitioned_buffers_all_gather(muon_params, local_grad_parts, muon_dtype) + full_momentums = [torch.zeros_like(full_grad) for full_grad in full_grads] + else: + full_grads, full_momentums = self._partitioned_buffers_all_gather( + muon_params, local_grad_parts, muon_dtype, additional_buffers_to_allgather=local_momentum_parts) + + # Unscale gathered gradients prior to Newton-Schulz and momentum tracking, + # since Newton-Schulz normalizes spectral norm and loses gradient scale. + loss_scale = float(self.loss_scale) + if loss_scale != 1.0: + for grad in full_grads: + grad.div_(loss_scale) + + optimizer_group = self.optimizer.param_groups[self.sub_group_to_group_id[sub_group_id]] + + for param, full_grad, full_momentum in zip(muon_params, full_grads, full_momentums): + update = muon_update(full_grad, + full_momentum, + beta=optimizer_group["momentum"], + ns_method=optimizer_group.get("ns_method", "gram"), + is_expert_group=getattr(param, "is_expert_group", False)) + partition_numel = param.partition_numel() + partition_rank = self._get_param_partition_rank(param) + start = partition_rank * partition_numel + real_numel = min(partition_numel, max(0, param.ds_numel - start)) + local_update = torch.zeros(partition_numel, dtype=update.dtype, device=accelerator_device) + if real_numel > 0: + local_update[:real_numel].copy_(update.view(-1).narrow(0, start, real_numel)) + _, dest_offset, _ = self.grad_position[self.get_param_id(param)] + # Rescale by loss_scale so downstream unscale_and_clip_grads cancels it cleanly + scaled_local_update = local_update * loss_scale if loss_scale != 1.0 else local_update + fp32_param.grad.narrow(0, dest_offset, + partition_numel).copy_(scaled_local_update.to(fp32_param.grad.dtype)) + local_momentum = torch.zeros(partition_numel, dtype=full_momentum.dtype, device=accelerator_device) + if real_numel > 0: + local_momentum[:real_numel].copy_(full_momentum.view(-1).narrow(0, start, real_numel)) + momentum.narrow(0, dest_offset, partition_numel).copy_(local_momentum.to(momentum.dtype)) + self.norm_for_param_grads[self.get_param_id(param)] = scaled_local_update.to(get_norm_dtype()).norm(2) + + if self.save_muon_momentum_buffer_in_memory and fp32_param in self.optimizer.state: + self.optimizer.state[fp32_param]["momentum_buffer"] = momentum + + if self._swappable_optimizer_subgroup(sub_group_id): + self._writeback_swap_state(sub_group_id, + write_opt_state=not self.save_muon_momentum_buffer_in_memory, + write_gradients=True) + @instrument_w_nvtx def _prepare_fp32_grad_for_sub_group(self, sub_group_id): partition_id = dist.get_rank(group=self._get_sub_group_process_group(sub_group_id)) @@ -2445,7 +2653,8 @@ def _optimizer_states_and_gradient_swap_in(self, sub_group_id, timer_names=None) self.optimizer_swapper.swap_in_optimizer_state( parameter=self.fp32_partitioned_groups_flat[sub_group_id], - async_parameter=self.next_swappable_fp32_partitioned_groups[sub_group_id]) + async_parameter=self.next_swappable_fp32_partitioned_groups[sub_group_id] + if timer_names is not None else None) if timer_names is not None: self.timers(OPTIMIZER_SWAP_IN_STATE_TIMER).stop() @@ -3275,6 +3484,7 @@ def _rigid_load_state_dict(self, state_dict, load_optimizer_states=True): self._set_fp32_optimizer_param_groups() self.optimizer.load_state_dict(state_dict[OPTIMIZER_STATE_DICT]) self._clear_fp32_optimizer_param_groups() + self._restore_muon_momentum_residency() if self.swap_optimizer: # Purge the swapped optimizer state, it was initialized to the freshly created model and not the checkpoint diff --git a/deepspeed/runtime/zero/stage_1_and_2.py b/deepspeed/runtime/zero/stage_1_and_2.py index f05a53867c93..d8ae60307771 100644 --- a/deepspeed/runtime/zero/stage_1_and_2.py +++ b/deepspeed/runtime/zero/stage_1_and_2.py @@ -229,10 +229,6 @@ def __init__(self, self.reduce_scatter = reduce_scatter - if isinstance(self.optimizer, MuonWithAuxAdam) and self.reduce_scatter and self.cpu_offload: - raise ValueError("Muon with reduce scatter does not support optimizer offload because offload retains " - "only partition slices; disable reduce scatter or optimizer offload") - self.overlap_comm = overlap_comm self.deepspeed_adam_offload = self.cpu_offload @@ -365,6 +361,10 @@ def _enforce_cpu_offload(): else: self.use_grad_accum_attribute = False + self._muon_allgather_buffers = OrderedDict() + self._muon_allgather_buffer_bytes = 0 + self._muon_allgather_max_cached_bytes = 256 * 1024 * 1024 + self.round_robin_bit16_groups = [] self.round_robin_bit16_indices = [] self.round_robin_bit16_meta = [] @@ -686,6 +686,7 @@ def destroy(self): for hook in self._grad_acc_hooks: hook.remove() self.print_rank_0("Removed grad acc hooks") + self._clear_muon_allgather_buffers() if get_accelerator().is_available(): get_accelerator().synchronize() self._unpin_offload_buffers() @@ -1656,6 +1657,158 @@ def complete_grad_norm_calculation_for_cpu_offload(self, params): return torch.tensor(total_norm, device=self.device, dtype=torch.float) + def _muon_all_gather_partitions(self, params, group_idx, flat_buffers, process_group, device): + """Gather only the partition slices needed by Muon parameters.""" + world_size = dist.get_world_size(group=process_group) + rank = dist.get_rank(group=process_group) + partition_size = int(self.partition_size[group_idx]) + + param_ids = [self.get_param_id(param) for param in params] + partition_numels = [] + for param, param_id in zip(params, param_ids): + partition_numel = 0 + for partition_id in self.param_to_partition_ids[group_idx][param_id]: + source_offset = int(self.grad_partition_insertion_offset[group_idx][partition_id][param_id]) + param_offset = int(self.grad_start_offset[group_idx][partition_id][param_id]) + partition_numel = max(partition_numel, min(param.numel() - param_offset, + partition_size - source_offset)) + partition_numels.append(partition_numel) + slot_offsets = [0] + for partition_numel in partition_numels: + slot_offsets.append(slot_offsets[-1] + partition_numel) + slots_numel = slot_offsets[-1] + + def reconstruct(flat_buffer, buffer_index=0, compact=True): + outputs = [] + for param, param_id, slot_offset in zip(params, param_ids, slot_offsets): + full_param = torch.zeros(param.numel(), dtype=flat_buffer.dtype, device=device) + partition_ids = self.param_to_partition_ids[group_idx][param_id] + for partition_id in partition_ids: + source_offset = int(self.grad_partition_insertion_offset[group_idx][partition_id][param_id]) + param_offset = int(self.grad_start_offset[group_idx][partition_id][param_id]) + num_elements = int(min(param.numel() - param_offset, partition_size - source_offset)) + if num_elements > 0: + if compact: + source_offset_in_buffer = (partition_id * (slots_numel * len(flat_buffers)) + + buffer_index * slots_numel + slot_offset) + else: + source_offset_in_buffer = source_offset + source = flat_buffer.narrow(0, source_offset_in_buffer, num_elements) + full_param.narrow(0, param_offset, num_elements).copy_(source) + outputs.append(full_param.view_as(param)) + return outputs + + if world_size == 1: + outputs = [reconstruct(buffer, index, compact=False) for index, buffer in enumerate(flat_buffers)] + return outputs[0] if len(outputs) == 1 else outputs + + cache_key = (group_idx, world_size, slots_numel, len(flat_buffers), flat_buffers[0].dtype, device) + cache = self._muon_allgather_buffers.pop(cache_key, None) + gathered_numel = slots_numel * len(flat_buffers) * world_size + cache_bytes = (slots_numel * len(flat_buffers) + gathered_numel) * flat_buffers[0].element_size() + if cache is None or cache[0].numel() != slots_numel * len(flat_buffers) or cache[1].numel() != gathered_numel: + local_buffer = torch.empty(slots_numel * len(flat_buffers), dtype=flat_buffers[0].dtype, device=device) + gathered_buffer = torch.empty(gathered_numel, dtype=flat_buffers[0].dtype, device=device) + if cache is not None: + self._muon_allgather_buffer_bytes -= cache[2] + while (self._muon_allgather_buffers + and self._muon_allgather_buffer_bytes + cache_bytes > self._muon_allgather_max_cached_bytes): + _, evicted = self._muon_allgather_buffers.popitem(last=False) + self._muon_allgather_buffer_bytes -= evicted[2] + if cache_bytes <= self._muon_allgather_max_cached_bytes: + self._muon_allgather_buffers[cache_key] = (local_buffer, gathered_buffer, cache_bytes) + self._muon_allgather_buffer_bytes += cache_bytes + else: + local_buffer, gathered_buffer, _ = cache + self._muon_allgather_buffers[cache_key] = cache + for index, flat_buffer in enumerate(flat_buffers): + for param_index, (param, param_id, slot_offset) in enumerate(zip(params, param_ids, slot_offsets)): + destination = local_buffer.narrow(0, index * slots_numel + slot_offset, partition_numels[param_index]) + destination.zero_() + partition_offset = self.grad_partition_insertion_offset[group_idx][rank].get(param_id) + if partition_offset is not None: + source_offset = int(partition_offset) + param_offset = int(self.grad_start_offset[group_idx][rank][param_id]) + num_elements = int(min(param.numel() - param_offset, partition_size - source_offset)) + if num_elements > 0: + destination.narrow(0, 0, + num_elements).copy_(flat_buffer.narrow(0, source_offset, num_elements), + non_blocking=True) + + dist.all_gather_into_tensor(gathered_buffer, local_buffer, group=process_group) + outputs = [reconstruct(gathered_buffer, index) for index in range(len(flat_buffers))] + return outputs[0] if len(outputs) == 1 else outputs + + def _clear_muon_allgather_buffers(self): + self._muon_allgather_buffers.clear() + self._muon_allgather_buffer_bytes = 0 + + @torch.no_grad() + def _apply_muon_updates_cpu_offload(self): + """Orthogonalize full Muon gradients before clipping CPU-offloaded updates.""" + if not isinstance(self.optimizer, MuonWithAuxAdam): + return + + accelerator_device = get_accelerator().current_device_name() + for group_idx, group in enumerate(self.round_robin_bit16_groups): + muon_params = [param for param in group if getattr(param, "use_muon", False)] + if not muon_params: + continue + + process_group = self.real_dp_process_group[group_idx] + rank = dist.get_rank(group=process_group) + partition_size = int(self.partition_size[group_idx]) + local_grad = self.single_partition_of_fp32_groups[group_idx].grad + + flat_param = self.single_partition_of_fp32_groups[group_idx] + state = self.optimizer.state.setdefault(flat_param, {}) + momentum = state.get("momentum_buffer") + momentum_was_created = momentum is None or momentum.numel() != local_grad.numel() + if momentum_was_created: + # A newly allocated state is zero on every rank, so it needs no all-gather. + momentum = torch.zeros_like(flat_param) + state["momentum_buffer"] = momentum + if momentum_was_created: + full_grad = self._muon_all_gather_partitions(muon_params, group_idx, [local_grad], process_group, + accelerator_device) + full_momentum = [torch.zeros_like(grad) for grad in full_grad] + else: + full_grad, full_momentum = self._muon_all_gather_partitions(muon_params, group_idx, + [local_grad, momentum], process_group, + accelerator_device) + + # Unscale gathered gradients prior to Newton-Schulz and momentum tracking, + # since Newton-Schulz normalizes spectral norm and loses gradient scale. + loss_scale = float(self.loss_scale) + if loss_scale != 1.0: + for grad in full_grad: + grad.div_(loss_scale) + + optimizer_group = self.optimizer.param_groups[group_idx] + for param, grad, param_momentum in zip(muon_params, full_grad, full_momentum): + param_id = self.get_param_id(param) + update = muon_update(grad, + param_momentum, + optimizer_group["momentum"], + ns_method=optimizer_group.get("ns_method", "gram"), + is_expert_group=getattr(param, "is_expert_group", False)) + + if rank not in self.param_to_partition_ids[group_idx][param_id]: + continue + source_offset = int(self.grad_start_offset[group_idx][rank][param_id]) + dest_offset = int(self.grad_partition_insertion_offset[group_idx][rank][param_id]) + num_elements = int(min(param.numel() - source_offset, partition_size - dest_offset)) + if num_elements > 0: + local_update = update.view(-1).narrow(0, source_offset, num_elements) + # Rescale by loss_scale so downstream unscale_and_clip_grads cancels it cleanly + scaled_local_update = local_update * loss_scale if loss_scale != 1.0 else local_update + self.single_partition_of_fp32_groups[group_idx].grad.view(-1).narrow( + 0, dest_offset, num_elements).copy_( + scaled_local_update.to(self.single_partition_of_fp32_groups[group_idx].grad.dtype)) + self.norm_for_param_grads[param_id] = scaled_local_update.to(get_norm_dtype()).norm(2) + momentum_update = param_momentum.view(-1).narrow(0, source_offset, num_elements) + momentum.narrow(0, dest_offset, num_elements).copy_(momentum_update.to(momentum.dtype)) + ############################################################################################ def copy_grads_in_partition(self, param): if self.cpu_offload: @@ -2324,6 +2477,8 @@ def step(self, closure=None): # Step 1:- Calculate gradient norm using bit-16 grads see_memory_usage('Before norm calculation') + if self.cpu_offload: + self._apply_muon_updates_cpu_offload() scaled_global_grad_norm = self.scaled_global_norm() self._global_grad_norm = scaled_global_grad_norm / prev_scale see_memory_usage('After norm before optimizer') diff --git a/tests/unit/ops/muon/test_muon.py b/tests/unit/ops/muon/test_muon.py index 654b256450a7..fdde19ea242d 100644 --- a/tests/unit/ops/muon/test_muon.py +++ b/tests/unit/ops/muon/test_muon.py @@ -11,6 +11,8 @@ from unit.common import DistributedTest from unit.simple_model import SimpleModel from deepspeed.accelerator import get_accelerator +from deepspeed.runtime.zero.stage_1_and_2 import DeepSpeedZeroOptimizer +from deepspeed.runtime.zero.stage3 import DeepSpeedZeroOptimizer_Stage3 if torch.half not in get_accelerator().supported_dtypes(): pytest.skip(f"fp16 not supported, valid dtype: {get_accelerator().supported_dtypes()}", allow_module_level=True) @@ -180,13 +182,13 @@ def test_ns_method_stage3(self, ns_method): engine.step() -class TestMuonRejectsReduceScatter(DistributedTest): - """Optimizer offload does not yet support Muon with reduce-scatter.""" +class TestMuonOptimizerOffload(DistributedTest): + """Muon remains correct when optimizer state is kept on the CPU.""" world_size = 1 - @pytest.mark.parametrize('zero_stage', [1, 2]) - def test_muon_reduce_scatter_with_optimizer_offload_raises(self, zero_stage): + @pytest.mark.parametrize('zero_stage', [1, 2, 3]) + def test_muon_with_optimizer_offload(self, zero_stage): config_dict = { "train_batch_size": 4, "optimizer": { @@ -200,7 +202,7 @@ def test_muon_reduce_scatter_with_optimizer_offload_raises(self, zero_stage): }, "zero_optimization": { "stage": zero_stage, - "reduce_scatter": True, + "reduce_scatter": False, "offload_optimizer": { "device": "cpu", "pin_memory": True, @@ -208,11 +210,32 @@ def test_muon_reduce_scatter_with_optimizer_offload_raises(self, zero_stage): }, } model = SimpleModel(hidden_dim=32, nlayers=2) - with pytest.raises(ValueError, match="Muon with reduce scatter does not support optimizer offload"): - deepspeed.initialize(config=config_dict, - model=model, - model_parameters=model.parameters(), - dist_init_required=False) + initial_params = [p.detach().clone().cpu() for p in model.parameters()] + engine, _, _, _ = deepspeed.initialize(config=config_dict, + model=model, + model_parameters=model.parameters(), + dist_init_required=False) + x = torch.randn(4, 32, device=engine.device, dtype=torch.half) + y = torch.randint(0, 32, (4, ), device=engine.device) + engine.backward(engine(x, y)) + engine.step() + assert any(not torch.equal(initial, + current.detach().cpu()) + for initial, current in zip(initial_params, model.parameters())) + + +class TestMuonAllGatherBufferLifecycle: + + @pytest.mark.parametrize("optimizer_class", [DeepSpeedZeroOptimizer, DeepSpeedZeroOptimizer_Stage3]) + def test_clear_muon_allgather_buffers(self, optimizer_class): + optimizer = optimizer_class.__new__(optimizer_class) + optimizer._muon_allgather_buffers = {"test": object()} + optimizer._muon_allgather_buffer_bytes = 128 + + optimizer._clear_muon_allgather_buffers() + + assert not optimizer._muon_allgather_buffers + assert optimizer._muon_allgather_buffer_bytes == 0 class TestMuonZero12NumericalCorrectness(DistributedTest): @@ -228,21 +251,23 @@ class TestMuonZero12NumericalCorrectness(DistributedTest): @pytest.mark.parametrize( "zero_stage,ns_method,reduce_scatter,gas,overlap_comm,use_multi_rank_bucket_allreduce," - "contiguous_gradients,reduce_bucket_size", [ - pytest.param(1, "gram", False, 1, False, True, True, 500000000, id="z1-gram-allreduce"), - pytest.param(1, "standard", False, 1, False, True, True, 500000000, id="z1-standard-allreduce"), - pytest.param(2, "gram", False, 1, False, True, True, 500000000, id="z2-gram-allreduce"), - pytest.param(2, "standard", False, 1, False, True, True, 500000000, id="z2-standard-allreduce"), - pytest.param(1, "gram", True, 1, False, True, True, 500000000, id="z1-reduce-scatter"), - pytest.param(2, "gram", True, 1, False, True, True, 500000000, id="z2-reduce-scatter"), - pytest.param(2, "gram", True, 2, True, True, True, 500000000, id="z2-rs-gas2-overlap"), - pytest.param(2, "gram", True, 2, False, False, True, 500000000, id="z2-rs-gas2-no-multi-rank"), - pytest.param(2, "gram", True, 1, False, True, True, 32768, id="z2-rs-extra-large-param"), - pytest.param(2, "gram", True, 1, False, True, False, 500000000, id="z2-rs-noncontiguous"), + "contiguous_gradients,reduce_bucket_size,offload_optimizer", [ + pytest.param(1, "gram", False, 1, False, True, True, 500000000, False, id="z1-gram-allreduce"), + pytest.param(1, "standard", False, 1, False, True, True, 500000000, False, id="z1-standard-allreduce"), + pytest.param(2, "gram", False, 1, False, True, True, 500000000, False, id="z2-gram-allreduce"), + pytest.param(2, "standard", False, 1, False, True, True, 500000000, False, id="z2-standard-allreduce"), + pytest.param(1, "gram", True, 1, False, True, True, 500000000, False, id="z1-reduce-scatter"), + pytest.param(2, "gram", True, 1, False, True, True, 500000000, False, id="z2-reduce-scatter"), + pytest.param(2, "gram", True, 2, True, True, True, 500000000, False, id="z2-rs-gas2-overlap"), + pytest.param(2, "gram", True, 2, False, False, True, 500000000, False, id="z2-rs-gas2-no-multi-rank"), + pytest.param(2, "gram", True, 1, False, True, True, 32768, False, id="z2-rs-extra-large-param"), + pytest.param(2, "gram", True, 1, False, True, False, 500000000, False, id="z2-rs-noncontiguous"), + pytest.param(1, "gram", False, 2, False, True, True, 500000000, True, id="z1-offload-gas2"), + pytest.param(2, "gram", True, 2, False, True, True, 500000000, True, id="z2-offload-rs-gas2"), ]) def test_update_matches_full_gradient_reference(self, zero_stage, ns_method, reduce_scatter, gas, overlap_comm, use_multi_rank_bucket_allreduce, contiguous_gradients, - reduce_bucket_size): + reduce_bucket_size, offload_optimizer): import copy from deepspeed.utils import safe_get_full_fp32_param from deepspeed.runtime.zero.muon.original_muon import muon_update @@ -286,6 +311,11 @@ def test_update_matches_full_gradient_reference(self, zero_stage, ns_method, red "reduce_bucket_size": reduce_bucket_size, }, } + if offload_optimizer: + config_dict["zero_optimization"]["offload_optimizer"] = { + "device": "cpu", + "pin_memory": True, + } engine, _, _, _ = deepspeed.initialize(config=config_dict, model=model, model_parameters=model.parameters(), @@ -360,3 +390,709 @@ def test_update_matches_full_gradient_reference(self, zero_stage, ns_method, red f"full-gradient reference -- orthogonalization likely ran on a partition slice rather than " f"the full averaged gradient (#7807)") assert changed, "optimizer step did not update any Muon weight (skipped step?)" + + +class TestMuonOffloadLossScaling(DistributedTest): + """Verify Muon updates under CPU offload are invariant to loss_scale with clipping.""" + + world_size = 2 + + @pytest.mark.parametrize("zero_stage", [1, 2, 3]) + def test_offload_loss_scale_equivalence(self, zero_stage): + from deepspeed.utils import safe_get_full_fp32_param + + hidden_dim, nlayers = 128, 2 + lr = 0.02 + clip_grad = 1.0 + + def _run_with_scale(loss_scale): + torch.manual_seed(42) + model = SimpleModel(hidden_dim=hidden_dim, nlayers=nlayers) + config_dict = { + "train_micro_batch_size_per_gpu": 4, + "gradient_clipping": clip_grad, + "optimizer": { + "type": "muon", + "params": { + "lr": lr, + "momentum": 0.0 + } + }, + "fp16": { + "enabled": True, + "loss_scale": loss_scale + }, + "zero_optimization": { + "stage": zero_stage, + "reduce_scatter": False, + "offload_optimizer": { + "device": "cpu", + "pin_memory": True + } + } + } + engine, _, _, _ = deepspeed.initialize(config=config_dict, + model=model, + model_parameters=model.parameters(), + dist_init_required=False) + device = engine.device + muon_named = [(n, p) for n, p in engine.module.named_parameters() + if getattr(p, "use_muon", False) or len(getattr(p, "ds_shape", p.shape)) >= 2] + assert len(muon_named) > 0, "No Muon parameters identified!" + pre = {n: safe_get_full_fp32_param(p).clone() for n, p in muon_named} + + torch.manual_seed(999) + x = torch.randn(4, hidden_dim, device=device).half() + y = torch.randint(0, hidden_dim, (4, ), device=device) + loss = engine(x, y) + engine.backward(loss) + engine.step() + + post = {n: safe_get_full_fp32_param(p).clone() for n, p in muon_named} + updates = {n: (pre[n] - post[n]).float() for n in pre} + return updates + + updates_scale_1 = _run_with_scale(1.0) + updates_scale_1024 = _run_with_scale(1024.0) + + for n in updates_scale_1: + norm_1 = updates_scale_1[n].norm().item() + norm_1024 = updates_scale_1024[n].norm().item() + assert norm_1 > 0.01, f"Update vanished for {n} under scale 1.0" + assert norm_1024 > 0.01, f"Update vanished for {n} under scale 1024.0" + max_diff = (updates_scale_1[n] - updates_scale_1024[n]).abs().max().item() + assert max_diff < 1e-3, ( + f"Update mismatch between scale 1.0 and 1024.0 under stage {zero_stage} " + f"with gradient clipping: max_diff={max_diff}, norm_1={norm_1}, norm_1024={norm_1024}") + + +class TestMuonZero3NVMeMomentumResidency(DistributedTest): + """Verify ZeRO-3 resident Muon momentum buffers persist across NVMe swapping steps.""" + + world_size = 1 + + @pytest.mark.parametrize("pipeline", [False, True]) + def test_zero3_nvme_momentum_residency(self, tmpdir, pipeline): + from deepspeed.ops.aio import AsyncIOBuilder + from deepspeed.runtime.swap_tensor.partitioned_optimizer_swapper import PartitionedOptimizerSwapper + from deepspeed.runtime.swap_tensor.pipelined_optimizer_swapper import PipelinedOptimizerSwapper + + if not deepspeed.ops.__compatible_ops__[AsyncIOBuilder.NAME]: + pytest.skip("Skip tests since async-io is not compatible") + + hidden_dim, nlayers = 1024, 2 + lr = 0.01 + momentum = 0.95 + offload_optimizer_cfg = { + "device": "nvme", + "nvme_path": str(tmpdir), + } + if pipeline: + offload_optimizer_cfg["pipeline_read"] = True + offload_optimizer_cfg["pipeline_write"] = True + + config_dict = { + "train_micro_batch_size_per_gpu": 1, + "steps_per_print": 1, + "optimizer": { + "type": "muon", + "params": { + "lr": lr, + "momentum": momentum + } + }, + "fp16": { + "enabled": True, + "loss_scale": 1.0 + }, + "zero_optimization": { + "stage": 3, + "reduce_scatter": False, + "save_muon_momentum_buffer_in_memory": True, + "offload_optimizer": offload_optimizer_cfg, + "sub_group_size": 100 + }, + "aio": { + "block_size": 1048576 + } + } + torch.manual_seed(42) + model = SimpleModel(hidden_dim=hidden_dim, nlayers=nlayers) + engine, _, _, _ = deepspeed.initialize(config=config_dict, + model=model, + model_parameters=model.parameters(), + dist_init_required=False) + + device = engine.device + x = torch.randn(1, hidden_dim, device=device).half() + y = torch.randint(0, hidden_dim, (1, ), device=device) + + opt = engine.optimizer + assert opt.swap_optimizer, "NVMe swap_optimizer must be enabled" + assert opt.save_muon_momentum_buffer_in_memory + if pipeline: + assert isinstance(opt.optimizer_swapper, PipelinedOptimizerSwapper) + else: + assert isinstance(opt.optimizer_swapper, PartitionedOptimizerSwapper) + + # Step 1 + engine.backward(engine(x, y)) + engine.step() + + assert len(opt.muon_momentum_buffer_partitioned_groups_flat) > 0 + step1_momentums = {} + for sub_group_id, buf in opt.muon_momentum_buffer_partitioned_groups_flat.items(): + expected_numel = int(opt.fp16_partitioned_groups_flat_numel[sub_group_id]) + assert buf.numel() == expected_numel, ( + f"Resident momentum for subgroup {sub_group_id} had numel={buf.numel()} " + f"instead of {expected_numel}; storage was evicted by NVMe swapper!") + assert not getattr(buf, "swappable", True) + assert getattr(buf, "is_resident", False) + step1_momentums[sub_group_id] = buf.clone() + + # Step 2: verify multi-step execution does not crash and momentum accumulates + engine.backward(engine(x, y)) + engine.step() + + for sub_group_id, buf in opt.muon_momentum_buffer_partitioned_groups_flat.items(): + expected_numel = int(opt.fp16_partitioned_groups_flat_numel[sub_group_id]) + assert buf.numel() == expected_numel + diff = (buf - step1_momentums[sub_group_id]).abs().max().item() + assert diff > 0.0, f"Momentum buffer did not accumulate changes at step 2 for subgroup {sub_group_id}" + + # Step 3 + engine.backward(engine(x, y)) + engine.step() + for sub_group_id, buf in opt.muon_momentum_buffer_partitioned_groups_flat.items(): + assert buf.numel() == int(opt.fp16_partitioned_groups_flat_numel[sub_group_id]) + + @pytest.mark.parametrize("offload", ["none", "nvme", "nvme_pipelined"]) + def test_zero3_momentum_survives_checkpoint(self, tmpdir, offload): + """An interrupted run must land on the same weights as an uninterrupted one. + + The resident momentum cache is a second reference to the tensor held in + ``optimizer.state``. ``Optimizer.load_state_dict()`` rebinds that entry to a fresh tensor, + and NVMe offload skips it entirely in favour of swap files that never hold the in-memory + buffer. Either way the resumed run would silently continue from a zero momentum. + """ + from deepspeed.runtime.zero.partition_parameters import Init + from deepspeed.runtime.swap_tensor.partitioned_optimizer_swapper import PartitionedOptimizerSwapper + from deepspeed.runtime.swap_tensor.pipelined_optimizer_swapper import PipelinedOptimizerSwapper + + uses_nvme = offload.startswith("nvme") + if uses_nvme: + from deepspeed.ops.aio import AsyncIOBuilder + if not deepspeed.ops.__compatible_ops__[AsyncIOBuilder.NAME]: + pytest.skip("Skip tests since async-io is not compatible") + + hidden_dim, nlayers = 1024, 2 + total_steps = 4 + interrupt_at = 2 + # A wide batch keeps the gradients full rank. Rank-deficient gradients make the + # Newton-Schulz iteration amplify fp16 noise in the near-null singular directions, which + # would swamp the effect this test is looking for. + micro_batch = 256 + offload_optimizer_cfg = None + if uses_nvme: + offload_optimizer_cfg = {"device": "nvme", "nvme_path": str(tmpdir.mkdir("nvme"))} + if offload == "nvme_pipelined": + offload_optimizer_cfg["pipeline_read"] = True + offload_optimizer_cfg["pipeline_write"] = True + config_dict = { + "train_micro_batch_size_per_gpu": micro_batch, + "steps_per_print": 1, + "optimizer": { + "type": "muon", + "params": { + "lr": 0.01, + "momentum": 0.95 + } + }, + "fp16": { + "enabled": True, + "loss_scale": 1.0 + }, + "zero_optimization": { + "stage": 3, + "reduce_scatter": False, + "save_muon_momentum_buffer_in_memory": True, + "sub_group_size": 100 + }, + "aio": { + "block_size": 1048576 + } + } + if offload_optimizer_cfg is not None: + config_dict["zero_optimization"]["offload_optimizer"] = offload_optimizer_cfg + + # Bias-free matrices so that every parameter takes the Muon path and the resident momentum + # buffer is the only optimizer state that has to survive the checkpoint. + class MatrixOnlyModel(torch.nn.Module): + + def __init__(self): + super().__init__() + self.layers = torch.nn.ModuleList( + [torch.nn.Linear(hidden_dim, hidden_dim, bias=False) for _ in range(nlayers)]) + + def forward(self, x): + for layer in self.layers: + x = layer(x) + # Scaled up to keep the fp16 gradients out of the subnormal range. + return x.pow(2).mean() * 1024.0 + + def build_engine(): + # NVMe swap files are named after parameter ids, so every engine must number its + # parameters the same way for the copied checkpoint files to be picked up. + Init.param_id = 0 + torch.manual_seed(42) + model = MatrixOnlyModel() + engine, _, _, _ = deepspeed.initialize(config=config_dict, + model=model, + model_parameters=model.parameters(), + dist_init_required=False) + if uses_nvme: + expected_swapper = (PipelinedOptimizerSwapper + if offload == "nvme_pipelined" else PartitionedOptimizerSwapper) + assert isinstance(engine.optimizer.optimizer_swapper, expected_swapper) + else: + assert not engine.optimizer.swap_optimizer + return engine + + def resident_momentums(engine): + buffers = engine.optimizer.muon_momentum_buffer_partitioned_groups_flat + assert len(buffers) > 0 + return {sub_group_id: buf.clone().float().cpu() for sub_group_id, buf in buffers.items()} + + def flat_weights(engine): + params = list(engine.module.parameters()) + with deepspeed.zero.GatheredParameters(params, modifier_rank=None): + return torch.cat([p.detach().float().cpu().reshape(-1) for p in params]) + + torch.manual_seed(7) + inputs = torch.randn(total_steps, micro_batch, hidden_dim).half() + + def run(engine, steps): + for step in steps: + engine.backward(engine(inputs[step].to(engine.device))) + engine.step() + + reference_engine = build_engine() + run(reference_engine, range(total_steps)) + reference_weights = flat_weights(reference_engine) + reference_engine.destroy() + + interrupted_engine = build_engine() + run(interrupted_engine, range(interrupt_at)) + checkpoint_weights = flat_weights(interrupted_engine) + saved_momentums = resident_momentums(interrupted_engine) + assert any(buf.abs().max().item() > 0.0 for buf in saved_momentums.values()) + ckpt_dir = str(tmpdir.mkdir("ckpt")) + interrupted_engine.save_checkpoint(ckpt_dir) + interrupted_engine.destroy() + + resumed_engine = build_engine() + resumed_engine.load_checkpoint(ckpt_dir) + assert torch.equal(flat_weights(resumed_engine), checkpoint_weights), \ + "Model weights were not restored from the checkpoint" + + restored_momentums = resident_momentums(resumed_engine) + assert restored_momentums.keys() == saved_momentums.keys() + for sub_group_id, saved in saved_momentums.items(): + restored = restored_momentums[sub_group_id] + assert torch.equal( + restored, + saved), (f"Resident Muon momentum for subgroup {sub_group_id} was not restored from the checkpoint; " + f"max diff {(restored - saved).abs().max().item()}") + fp32_param = resumed_engine.optimizer.fp32_partitioned_groups_flat[sub_group_id] + assert resumed_engine.optimizer.optimizer.state[fp32_param]["momentum_buffer"] is \ + resumed_engine.optimizer.muon_momentum_buffer_partitioned_groups_flat[sub_group_id], \ + ("Optimizer state and the resident cache must reference the same momentum tensor, " + "otherwise the next step silently discards the restored values") + + run(resumed_engine, range(interrupt_at, total_steps)) + # Newton-Schulz runs in reduced precision and is not bitwise reproducible, so compare the + # weight update accumulated after the checkpoint instead of requiring exact equality. + reference_update = reference_weights - checkpoint_weights + resumed_update = flat_weights(resumed_engine) - checkpoint_weights + relative_error = ((resumed_update - reference_update).norm() / reference_update.norm()).item() + assert relative_error < 0.05, ("Resuming from a checkpoint diverged from the uninterrupted run; " + f"relative update error {relative_error}") + + def test_zero3_nvme_aggregate_unswapped_fragments(self, tmpdir): + from deepspeed.ops.aio import AsyncIOBuilder + if not deepspeed.ops.__compatible_ops__[AsyncIOBuilder.NAME]: + pytest.skip("Skip tests since async-io is not compatible") + + # 20 layers of 128x128: each parameter is 16,384 elements (< 1 MiB in FP32). + # Subgroup aggregate is 327,680 elements (> 262,144 elements / 1 MiB in FP32), + # making the subgroup swappable while all individual gradient fragments are unswapped. + class SmallMatrixModel(torch.nn.Module): + + def __init__(self, num_layers=20, dim=128): + super().__init__() + self.layers = torch.nn.ModuleList([torch.nn.Linear(dim, dim, bias=False) for _ in range(num_layers)]) + + def forward(self, x): + for l in self.layers: + x = l(x) + return x.sum() + + config_dict = { + "train_micro_batch_size_per_gpu": 1, + "steps_per_print": 1, + "optimizer": { + "type": "muon", + "params": { + "lr": 0.01, + "momentum": 0.95 + } + }, + "fp16": { + "enabled": True, + "loss_scale": 1.0 + }, + "zero_optimization": { + "stage": 3, + "reduce_scatter": False, + "save_muon_momentum_buffer_in_memory": True, + "offload_optimizer": { + "device": "nvme", + "nvme_path": str(tmpdir) + }, + "sub_group_size": 1000000 + }, + "aio": { + "block_size": 1048576 + } + } + torch.manual_seed(42) + model = SmallMatrixModel() + engine, _, _, _ = deepspeed.initialize(config=config_dict, + model=model, + model_parameters=model.parameters(), + dist_init_required=False) + + device = engine.device + x = torch.randn(1, 128, device=device).half() + opt = engine.optimizer + assert opt.swap_optimizer + + initial_params = [p.clone().detach().cpu() for p in model.parameters()] + for step in range(2): + loss = engine(x) + engine.backward(loss) + engine.step() + + assert any(not torch.equal(init, p.detach().cpu()) for init, p in zip(initial_params, model.parameters())) + + @pytest.mark.parametrize("pipeline", [False, True]) + def test_zero3_nvme_mixed_fragments_numerical_equivalence(self, tmpdir, pipeline): + import copy + from deepspeed.ops.aio import AsyncIOBuilder + from deepspeed.utils import safe_get_full_fp32_param + from deepspeed.runtime.swap_tensor.partitioned_optimizer_swapper import PartitionedOptimizerSwapper + from deepspeed.runtime.swap_tensor.pipelined_optimizer_swapper import PipelinedOptimizerSwapper + from deepspeed.runtime.zero.muon.original_muon import muon_update + + if not deepspeed.ops.__compatible_ops__[AsyncIOBuilder.NAME]: + pytest.skip("Skip tests since async-io is not compatible") + + class MixedMatrixModel(torch.nn.Module): + + def __init__(self): + super().__init__() + # 512 x 512 = 262,144 elements (>= 1 MiB in FP32) -> SWAPPED + self.large = torch.nn.Linear(512, 512, bias=False) + # 512 x 128 = 65,536 elements (< 1 MiB in FP32) -> UNSWAPPED + self.proj = torch.nn.Linear(512, 128, bias=False) + # 128 x 128 = 16,384 elements (< 1 MiB in FP32) -> UNSWAPPED + self.small = torch.nn.Linear(128, 128, bias=False) + + def forward(self, x): + # A squared loss over a wide batch keeps the weight gradients well conditioned, + # so Newton-Schulz does not amplify fp16 noise into the comparison. The constant + # factor lifts the gradients out of the fp16 subnormal range. + return self.small(self.proj(self.large(x))).pow(2).mean() * 1024.0 + + lr = 0.01 + momentum = 0.95 + num_steps = 2 + micro_batch = 256 + + torch.manual_seed(42) + base_model = MixedMatrixModel() + init_state = copy.deepcopy(base_model.state_dict()) + + # 1. Independent full-gradient reference (pure PyTorch + canonical muon_update across 2 steps) + device = get_accelerator().current_device_name() + ref_masters = {n: p.clone().detach().cpu().float() for n, p in init_state.items()} + init_masters = {n: p.clone() for n, p in ref_masters.items()} + # Muon momentum is a persistent optimizer state and is kept in fp32, matching ZeRO-1/2/3. + ref_momentums = {n: torch.zeros_like(p).to(device).float() for n, p in ref_masters.items()} + + gen = torch.Generator().manual_seed(1234) + inputs = [torch.randn(micro_batch, 512, generator=gen).to(device).half() for _ in range(num_steps)] + + for step in range(num_steps): + ref_model = MixedMatrixModel().to(device).half() + ref_model.load_state_dict({k: v.to(device).half() for k, v in ref_masters.items()}) + + ref_model.zero_grad(set_to_none=True) + loss = ref_model(inputs[step]) + loss.backward() + with torch.no_grad(): + for n, p in ref_model.named_parameters(): + update = muon_update(p.grad.detach().float(), ref_momentums[n], beta=momentum, ns_method="gram") + ref_masters[n].add_(update.cpu().float(), alpha=-lr) + + ref_updates = {n: (init_masters[n] - ref_masters[n]) for n in ref_masters} + + # 2. ZeRO-3 NVMe run with mixed swapped/unswapped fragments in the same subgroup + nvme_cfg = { + "device": "nvme", + "nvme_path": str(tmpdir), + } + if pipeline: + nvme_cfg["pipeline_read"] = True + nvme_cfg["pipeline_write"] = True + + config_dict = { + "train_micro_batch_size_per_gpu": micro_batch, + "steps_per_print": 1, + "gradient_clipping": 0.0, + "optimizer": { + "type": "muon", + "params": { + "lr": lr, + "momentum": momentum + } + }, + "fp16": { + "enabled": True, + "loss_scale": 1.0 + }, + "zero_optimization": { + "stage": 3, + "reduce_scatter": False, + "save_muon_momentum_buffer_in_memory": True, + "offload_optimizer": nvme_cfg, + "sub_group_size": 1000000 + }, + "aio": { + "block_size": 1048576 + } + } + model = MixedMatrixModel() + model.load_state_dict({k: v.clone() for k, v in init_state.items()}) + engine, _, _, _ = deepspeed.initialize(config=config_dict, + model=model, + model_parameters=model.parameters(), + dist_init_required=False) + + # Explicitly verify all 3 parameters are identified as Muon parameters under ZeRO-3 + muon_named = [(n, p) for n, p in engine.module.named_parameters() + if getattr(p, "use_muon", False) or len(getattr(p, "ds_shape", p.shape)) >= 2] + assert len(muon_named) == 3, f"Expected 3 Muon parameters, got {len(muon_named)}: {[n for n, p in muon_named]}" + assert set(n for n, p in muon_named) == {"large.weight", "proj.weight", "small.weight"} + init_params = {n: safe_get_full_fp32_param(p).clone().cpu() for n, p in muon_named} + + for step in range(num_steps): + loss = engine(inputs[step]) + engine.backward(loss) + engine.step() + + opt = engine.optimizer + assert opt.swap_optimizer + if pipeline: + assert isinstance(opt.optimizer_swapper, PipelinedOptimizerSwapper) + else: + assert isinstance(opt.optimizer_swapper, PartitionedOptimizerSwapper) + + nvme_final = {n: safe_get_full_fp32_param(p).clone().cpu() for n, p in muon_named} + applied_updates = {n: (init_params[n] - nvme_final[n]) for n, p in muon_named} + + # Verify non-trivial update and numerical equivalence against independent reference + for n in ref_updates: + applied_norm = applied_updates[n].norm().item() + ref_norm = ref_updates[n].norm().item() + assert ref_norm > 0.01, f"Reference update vanished for {n}" + assert applied_norm > 0.01, f"Engine update vanished for {n}" + rel_err = ((applied_updates[n] - ref_updates[n]).norm() / (ref_norm + 1e-8)).item() + assert rel_err < 0.05, (f"Numerical divergence for {n} under pipeline={pipeline}: " + f"rel_err={rel_err:.4f}, applied_norm={applied_norm:.4f}, ref_norm={ref_norm:.4f}") + + +class TestMuonZero3NVMeMultiRankMixedFragments(DistributedTest): + """Verify ZeRO-3 NVMe Muon multi-rank reconstruction, DP gradient averaging, and mixed fragments.""" + + world_size = 2 + + @pytest.mark.parametrize("pipeline", [False, True]) + def test_zero3_nvme_multirank_mixed_fragments(self, tmpdir, pipeline): + import copy + from deepspeed.ops.aio import AsyncIOBuilder + from deepspeed.utils import safe_get_full_fp32_param + from deepspeed.runtime.swap_tensor.partitioned_optimizer_swapper import PartitionedOptimizerSwapper + from deepspeed.runtime.swap_tensor.pipelined_optimizer_swapper import PipelinedOptimizerSwapper + from deepspeed.runtime.zero.muon.original_muon import muon_update + + if not deepspeed.ops.__compatible_ops__[AsyncIOBuilder.NAME]: + pytest.skip("Skip tests since async-io is not compatible") + + class MixedMatrixModel(torch.nn.Module): + + def __init__(self): + super().__init__() + # 512 x 1024 = 524,288 elements (2 MiB in FP32). + # Partition per rank (world_size=2) = 262,144 elements (>= 1 MiB in FP32) -> SWAPPED + self.large = torch.nn.Linear(512, 1024, bias=False) + # 1024 x 128 = 131,072 elements (512 KiB in FP32). + # Partition per rank = 65,536 elements (< 1 MiB in FP32) -> UNSWAPPED + self.proj = torch.nn.Linear(1024, 128, bias=False) + # 128 x 128 = 16,384 elements (64 KiB in FP32). + # Partition per rank = 8,192 elements (< 1 MiB in FP32) -> UNSWAPPED + self.small = torch.nn.Linear(128, 129, bias=False) + # 127 x 129 = 16,383 elements: odd numel, so ZeRO-3 pads the flat partition + # and the final rank owns a partially out-of-range slice -> UNSWAPPED + self.odd = torch.nn.Linear(129, 127, bias=False) + + def forward(self, x): + # A squared loss keeps the per-sample output error distinct, so the weight + # gradients stay well conditioned and Newton-Schulz does not amplify fp16 noise. + return self.odd(self.small(self.proj(self.large(x)))).pow(2).mean() * 1024.0 + + lr = 0.01 + momentum = 0.95 + num_steps = 2 + micro_batch = 256 + rank = dist.get_rank() + device = get_accelerator().current_device_name() + + torch.manual_seed(42) + base_model = MixedMatrixModel() + init_state = copy.deepcopy(base_model.state_dict()) + + # Deterministic global inputs across steps: each rank receives its own slice + gen = torch.Generator().manual_seed(1234) + inputs_step = [torch.randn(2 * micro_batch, 512, generator=gen).to(device).half() for _ in range(num_steps)] + + # 1. Independent high-precision oracle (on rank 0): + # Maintains FP32 master weights & FP32 momentum, computes FP16 forward/backward per rank, + # explicitly averages DP gradients across ranks, and steps Muon. + if rank == 0: + ref_masters = {n: p.clone().detach().cpu().float() for n, p in init_state.items()} + init_masters = {n: p.clone() for n, p in ref_masters.items()} + # Muon momentum is a persistent optimizer state and is kept in fp32, matching ZeRO-1/2/3. + ref_momentums = {n: torch.zeros_like(p).to(device).float() for n, p in ref_masters.items()} + + for step in range(num_steps): + x0 = inputs_step[step][:micro_batch] + x1 = inputs_step[step][micro_batch:] + + ref_model = MixedMatrixModel().to(device).half() + ref_model.load_state_dict({k: v.to(device).half() for k, v in ref_masters.items()}) + + ref_model.zero_grad(set_to_none=True) + loss0 = ref_model(x0) + loss0.backward() + grads0 = {n: p.grad.clone() for n, p in ref_model.named_parameters()} + + ref_model.zero_grad(set_to_none=True) + loss1 = ref_model(x1) + loss1.backward() + grads1 = {n: p.grad.clone() for n, p in ref_model.named_parameters()} + + # DeepSpeed averages DP gradients in the fp16 communication dtype, so the oracle + # must do the same before promoting to fp32 for momentum and Newton-Schulz. + avg_grads = {n: ((grads0[n] + grads1[n]) / 2.0).float() for n in grads0} + with torch.no_grad(): + for n in ref_masters: + up = muon_update(avg_grads[n], ref_momentums[n], beta=momentum, ns_method="gram") + ref_masters[n].add_(up.cpu().float(), alpha=-lr) + + ref_updates = {n: (init_masters[n] - ref_masters[n]) for n in ref_masters} + + # 2. DeepSpeed ZeRO-3 NVMe run with mixed swapped/unswapped fragments across 2 ranks + nvme_cfg = { + "device": "nvme", + "nvme_path": str(tmpdir), + } + if pipeline: + nvme_cfg["pipeline_read"] = True + nvme_cfg["pipeline_write"] = True + + config_dict = { + "train_micro_batch_size_per_gpu": micro_batch, + "steps_per_print": 1, + "gradient_clipping": 0.0, + "optimizer": { + "type": "muon", + "params": { + "lr": lr, + "momentum": momentum + } + }, + "fp16": { + "enabled": True, + "loss_scale": 1.0 + }, + "zero_optimization": { + "stage": 3, + "reduce_scatter": False, + "save_muon_momentum_buffer_in_memory": True, + "offload_optimizer": nvme_cfg, + "sub_group_size": 1000000 + }, + "aio": { + "block_size": 1048576 + } + } + model = MixedMatrixModel() + model.load_state_dict({k: v.clone() for k, v in init_state.items()}) + engine, _, _, _ = deepspeed.initialize(config=config_dict, + model=model, + model_parameters=model.parameters(), + dist_init_required=False) + + muon_named = [(n, p) for n, p in engine.module.named_parameters() + if getattr(p, "use_muon", False) or len(getattr(p, "ds_shape", p.shape)) >= 2] + assert len(muon_named) == 4 + assert set(n for n, p in muon_named) == {"large.weight", "proj.weight", "small.weight", "odd.weight"} + + # Exercise the padded final partition: odd.weight cannot be split evenly across 2 ranks, + # so its last partition is padded and reconstruction must drop the out-of-range tail. + odd_param = dict(muon_named)["odd.weight"] + world_size = dist.get_world_size() + assert odd_param.ds_numel % world_size != 0, "odd.weight must not divide evenly across ranks" + assert odd_param.partition_numel() * world_size > odd_param.ds_numel, "expected a padded final partition" + + init_engine_params = {n: safe_get_full_fp32_param(p).clone().cpu() for n, p in muon_named} + + # Each rank executes on its own micro-batch slice + rank_slice = slice(rank * micro_batch, (rank + 1) * micro_batch) + for step in range(num_steps): + rank_x = inputs_step[step][rank_slice] + loss = engine(rank_x) + engine.backward(loss) + engine.step() + + opt = engine.optimizer + assert opt.swap_optimizer + if pipeline: + assert isinstance(opt.optimizer_swapper, PipelinedOptimizerSwapper) + else: + assert isinstance(opt.optimizer_swapper, PartitionedOptimizerSwapper) + + final_engine_params = {n: safe_get_full_fp32_param(p).clone().cpu() for n, p in muon_named} + + if rank == 0: + applied_updates = {n: (init_engine_params[n] - final_engine_params[n]) for n, p in muon_named} + for n in ref_updates: + applied_norm = applied_updates[n].norm().item() + ref_norm = ref_updates[n].norm().item() + assert ref_norm > 0.01, f"Reference update vanished for {n}" + assert applied_norm > 0.01, f"Engine update vanished for {n}" + rel_err = ((applied_updates[n] - ref_updates[n]).norm() / (ref_norm + 1e-8)).item() + assert rel_err < 0.05, ( + f"Update divergence for {n} under pipeline={pipeline}: " + f"rel_err={rel_err:.4f}, applied_norm={applied_norm:.4f}, ref_norm={ref_norm:.4f}") diff --git a/tests/unit/runtime/tensor_parallel/test_autotp_universal_checkpoint.py b/tests/unit/runtime/tensor_parallel/test_autotp_universal_checkpoint.py index 2a1159359241..8b893b997de4 100644 --- a/tests/unit/runtime/tensor_parallel/test_autotp_universal_checkpoint.py +++ b/tests/unit/runtime/tensor_parallel/test_autotp_universal_checkpoint.py @@ -471,7 +471,7 @@ def test_sub_param_layer_materializes_zero_width_final_dimension(layer_cls): layer._tp_partition([layer.weight, None]) assert layer.weight.shape == (4, 0) - output = layer(torch.empty(2, 0)) + output = layer(torch.empty(2, 0, device=layer.weight.device)) assert output.shape == (2, 4)