From 8103ab37291d12c287317012ee4a3fa5429453a8 Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Thu, 27 Aug 2026 15:49:02 +0000 Subject: [PATCH 01/12] Fix Muon CPU optimizer offload Reconstruct full Muon gradients and momentum before Newton-Schulz updates while preserving CPUAdam auxiliary updates across ZeRO stages. Signed-off-by: Jin, Youzhi --- deepspeed/runtime/zero/stage3.py | 65 +++++++++++++++- deepspeed/runtime/zero/stage_1_and_2.py | 77 ++++++++++++++++++- tests/unit/ops/muon/test_muon.py | 58 ++++++++------ .../test_autotp_universal_checkpoint.py | 2 +- 4 files changed, 174 insertions(+), 28 deletions(-) diff --git a/deepspeed/runtime/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index d99a43ce601b..48dbce7a9018 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -1198,7 +1198,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'] = [] @@ -2387,6 +2387,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: @@ -2395,6 +2397,67 @@ 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] + state = self.optimizer.state.setdefault(fp32_param, {}) + momentum = state.get("momentum_buffer") + if momentum is None or momentum.numel() != fp32_param.numel(): + self._create_momentum_buffer(fp32_param.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)) + local_momentum_parts.append(momentum.narrow(0, dest_offset, numel).to(accelerator_device)) + + full_grads = self._partitioned_buffers_all_gather(muon_params, local_grad_parts, + self.gradient_accumulation_dtype) + full_momentums = self._partitioned_buffers_all_gather(muon_params, local_momentum_parts, + self.gradient_accumulation_dtype) + 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)] + fp32_param.grad.narrow(0, dest_offset, partition_numel).copy_(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)] = local_update.to(get_norm_dtype()).norm(2) + + if self._swappable_optimizer_subgroup(sub_group_id): + self._optimizer_states_and_gradient_swap_out(sub_group_id) + @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)) diff --git a/deepspeed/runtime/zero/stage_1_and_2.py b/deepspeed/runtime/zero/stage_1_and_2.py index 1fb27ee4dfbe..07f6391bd34c 100755 --- a/deepspeed/runtime/zero/stage_1_and_2.py +++ b/deepspeed/runtime/zero/stage_1_and_2.py @@ -228,10 +228,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 @@ -1647,6 +1643,77 @@ def complete_grad_norm_calculation_for_cpu_offload(self, params): return torch.tensor(total_norm, device=self.device, dtype=torch.float) + @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] + world_size = dist.get_world_size(group=process_group) + 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.to(accelerator_device) + full_grad = torch.empty(partition_size * world_size, dtype=local_grad.dtype, device=accelerator_device) + dist.all_gather_into_tensor(full_grad, local_grad, group=process_group) + + flat_param = self.single_partition_of_fp32_groups[group_idx] + state = self.optimizer.state.setdefault(flat_param, {}) + momentum = state.get("momentum_buffer") + if momentum is None or momentum.numel() != partition_size: + momentum = torch.zeros_like(flat_param) + state["momentum_buffer"] = momentum + local_momentum = momentum.to(accelerator_device) + full_momentum = torch.empty(partition_size * world_size, + dtype=local_momentum.dtype, + device=accelerator_device) + dist.all_gather_into_tensor(full_momentum, local_momentum, group=process_group) + + def reconstruct_param(buffer, param): + # A parameter can straddle partitions, so use the recorded slice map rather than + # deriving offsets from rank and partition size. + param_id = self.get_param_id(param) + full_param = torch.zeros(param.numel(), dtype=buffer.dtype, device=accelerator_device) + 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]) + num_elements = int(min(param.numel() - param_offset, partition_size - source_offset)) + if num_elements > 0: + source = buffer.narrow(0, partition_id * partition_size + source_offset, num_elements) + full_param.narrow(0, param_offset, num_elements).copy_(source) + return full_param.view_as(param) + + optimizer_group = self.optimizer.param_groups[group_idx] + for param in muon_params: + param_id = self.get_param_id(param) + grad = reconstruct_param(full_grad, param) + param_momentum = reconstruct_param(full_momentum, 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) + self.single_partition_of_fp32_groups[group_idx].grad.view(-1).narrow( + 0, dest_offset, num_elements).copy_( + local_update.to(self.single_partition_of_fp32_groups[group_idx].grad.dtype)) + self.norm_for_param_grads[param_id] = 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: @@ -2313,6 +2380,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..4bf650231409 100644 --- a/tests/unit/ops/muon/test_muon.py +++ b/tests/unit/ops/muon/test_muon.py @@ -180,13 +180,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 +200,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 +208,18 @@ 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 TestMuonZero12NumericalCorrectness(DistributedTest): @@ -228,21 +235,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 +295,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(), 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 fe996e267f3c..24e811d32876 100644 --- a/tests/unit/runtime/tensor_parallel/test_autotp_universal_checkpoint.py +++ b/tests/unit/runtime/tensor_parallel/test_autotp_universal_checkpoint.py @@ -469,7 +469,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) From 68973494be720c694b12f41531906c26efbe398c Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Thu, 27 Aug 2026 18:13:05 +0000 Subject: [PATCH 02/12] Improve the comm overhead, bound Muon offload gather buffers Limit cached GPU scratch buffers with LRU eviction and explicit cleanup, and cover buffer lifecycle behavior in tests. Signed-off-by: Jin, Youzhi --- deepspeed/runtime/zero/stage3.py | 144 +++++++++++++++++------- deepspeed/runtime/zero/stage_1_and_2.py | 128 ++++++++++++++++----- tests/unit/ops/muon/test_muon.py | 16 +++ 3 files changed, 222 insertions(+), 66 deletions(-) diff --git a/deepspeed/runtime/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index 48dbce7a9018..cc9ffac270d5 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -236,6 +236,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 @@ -525,6 +528,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: @@ -1933,57 +1937,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: @@ -2416,7 +2471,9 @@ def _apply_muon_updates_cpu_offload(self): fp32_param = self.fp32_partitioned_groups_flat[sub_group_id] state = self.optimizer.state.setdefault(fp32_param, {}) momentum = state.get("momentum_buffer") - if momentum is None or momentum.numel() != fp32_param.numel(): + momentum_was_created = momentum is None or momentum.numel() != fp32_param.numel() + if momentum_was_created: + # A newly allocated state is zero on every rank, so it needs no all-gather. self._create_momentum_buffer(fp32_param.numel(), sub_group_id, fp32_param.ds_id) momentum = state["momentum_buffer"] @@ -2426,12 +2483,19 @@ def _apply_muon_updates_cpu_offload(self): _, 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)) - local_momentum_parts.append(momentum.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)) - full_grads = self._partitioned_buffers_all_gather(muon_params, local_grad_parts, - self.gradient_accumulation_dtype) - full_momentums = self._partitioned_buffers_all_gather(muon_params, local_momentum_parts, + if momentum_was_created: + full_grads = self._partitioned_buffers_all_gather(muon_params, local_grad_parts, self.gradient_accumulation_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, + self.gradient_accumulation_dtype, + additional_buffers_to_allgather=local_momentum_parts) 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): diff --git a/deepspeed/runtime/zero/stage_1_and_2.py b/deepspeed/runtime/zero/stage_1_and_2.py index 07f6391bd34c..5c23f27e639e 100755 --- a/deepspeed/runtime/zero/stage_1_and_2.py +++ b/deepspeed/runtime/zero/stage_1_and_2.py @@ -360,6 +360,10 @@ def _enforce_cpu_offload(): else: self.use_grad_accum_attribute = False + self._muon_allgather_buffers = {} + 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 = [] @@ -1643,6 +1647,91 @@ 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) + return [reconstruct(gathered_buffer, index) for index in range(len(flat_buffers))] + + 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.""" @@ -1659,41 +1748,28 @@ def _apply_muon_updates_cpu_offload(self): world_size = dist.get_world_size(group=process_group) 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.to(accelerator_device) - full_grad = torch.empty(partition_size * world_size, dtype=local_grad.dtype, device=accelerator_device) - dist.all_gather_into_tensor(full_grad, local_grad, group=process_group) + 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") - if momentum is None or momentum.numel() != partition_size: + 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 - local_momentum = momentum.to(accelerator_device) - full_momentum = torch.empty(partition_size * world_size, - dtype=local_momentum.dtype, - device=accelerator_device) - dist.all_gather_into_tensor(full_momentum, local_momentum, group=process_group) - - def reconstruct_param(buffer, param): - # A parameter can straddle partitions, so use the recorded slice map rather than - # deriving offsets from rank and partition size. - param_id = self.get_param_id(param) - full_param = torch.zeros(param.numel(), dtype=buffer.dtype, device=accelerator_device) - 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]) - num_elements = int(min(param.numel() - param_offset, partition_size - source_offset)) - if num_elements > 0: - source = buffer.narrow(0, partition_id * partition_size + source_offset, num_elements) - full_param.narrow(0, param_offset, num_elements).copy_(source) - return full_param.view_as(param) + 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) optimizer_group = self.optimizer.param_groups[group_idx] - for param in muon_params: + for param, grad, param_momentum in zip(muon_params, full_grad, full_momentum): param_id = self.get_param_id(param) - grad = reconstruct_param(full_grad, param) - param_momentum = reconstruct_param(full_momentum, param) update = muon_update(grad, param_momentum, optimizer_group["momentum"], diff --git a/tests/unit/ops/muon/test_muon.py b/tests/unit/ops/muon/test_muon.py index 4bf650231409..3fb14b211f91 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) @@ -222,6 +224,20 @@ def test_muon_with_optimizer_offload(self, zero_stage): 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): """Numerical-correctness regression for #7807. From b139b8f82c6c2945d67dfcc0a61960970a11c663 Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Tue, 8 Sep 2026 16:31:34 +0000 Subject: [PATCH 03/12] Fix ZeRO-3 offload device mismatch, ZeRO-1/2 gather nesting, and cache eviction - Guard _apply_distributed_muon_update in ZeRO-3 when offload_optimizer is enabled to prevent device mismatch during backward and duplicate Newton-Schulz updates. - Unwrap single-buffer outputs in _muon_all_gather_partitions under multi-GPU ZeRO-1/2 to avoid returning nested lists. - Initialize _muon_allgather_buffers as an OrderedDict in ZeRO-1/2 so popitem(last=False) works during LRU cache eviction. - Clear cached all-gather buffers in DeepSpeedZeroOptimizer.destroy(). Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi --- deepspeed/runtime/zero/stage3.py | 2 +- deepspeed/runtime/zero/stage_1_and_2.py | 6 ++++-- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/deepspeed/runtime/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index cc9ffac270d5..609d4903166b 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -1629,7 +1629,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 = {} diff --git a/deepspeed/runtime/zero/stage_1_and_2.py b/deepspeed/runtime/zero/stage_1_and_2.py index 5c23f27e639e..23d9350076b6 100755 --- a/deepspeed/runtime/zero/stage_1_and_2.py +++ b/deepspeed/runtime/zero/stage_1_and_2.py @@ -360,7 +360,7 @@ def _enforce_cpu_offload(): else: self.use_grad_accum_attribute = False - self._muon_allgather_buffers = {} + self._muon_allgather_buffers = OrderedDict() self._muon_allgather_buffer_bytes = 0 self._muon_allgather_max_cached_bytes = 256 * 1024 * 1024 @@ -685,6 +685,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() self._unpin_offload_buffers() def _unpin_offload_buffers(self): @@ -1726,7 +1727,8 @@ def reconstruct(flat_buffer, buffer_index=0, compact=True): non_blocking=True) dist.all_gather_into_tensor(gathered_buffer, local_buffer, group=process_group) - return [reconstruct(gathered_buffer, index) for index in range(len(flat_buffers))] + 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() From ca25ef9662bdf8a4f7118120a0a2a992457624bc Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Tue, 8 Sep 2026 16:46:41 +0000 Subject: [PATCH 04/12] Keep Muon momentum buffer resident when save_muon_momentum_buffer_in_memory is set In ZeRO-3 CPU offload path, bypass NVMe swap-in/swap-out and use the resident muon_momentum_buffer_partitioned_groups_flat when save_muon_momentum_buffer_in_memory is enabled, avoiding unnecessary NVMe round-trips. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi --- deepspeed/runtime/zero/stage3.py | 28 +++++++++++++++++++--------- 1 file changed, 19 insertions(+), 9 deletions(-) diff --git a/deepspeed/runtime/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index 97eb4e58fed7..491781cf33d4 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -2471,17 +2471,24 @@ def _apply_muon_updates_cpu_offload(self): if not muon_params: continue - if self._swappable_optimizer_subgroup(sub_group_id): + if self._swappable_optimizer_subgroup(sub_group_id) and not self.save_muon_momentum_buffer_in_memory: self._optimizer_states_and_gradient_swap_in(sub_group_id) fp32_param = self.fp32_partitioned_groups_flat[sub_group_id] - state = self.optimizer.state.setdefault(fp32_param, {}) - momentum = state.get("momentum_buffer") - momentum_was_created = momentum is None or momentum.numel() != fp32_param.numel() - if momentum_was_created: - # A newly allocated state is zero on every rank, so it needs no all-gather. - self._create_momentum_buffer(fp32_param.numel(), sub_group_id, fp32_param.ds_id) - momentum = state["momentum_buffer"] + 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() != fp32_param.numel() + if momentum_was_created: + self._create_momentum_buffer(fp32_param.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() != fp32_param.numel() + if momentum_was_created: + # A newly allocated state is zero on every rank, so it needs no all-gather. + self._create_momentum_buffer(fp32_param.numel(), sub_group_id, fp32_param.ds_id) + momentum = state["momentum_buffer"] local_grad_parts = [] local_momentum_parts = [] @@ -2525,7 +2532,10 @@ def _apply_muon_updates_cpu_offload(self): momentum.narrow(0, dest_offset, partition_numel).copy_(local_momentum.to(momentum.dtype)) self.norm_for_param_grads[self.get_param_id(param)] = local_update.to(get_norm_dtype()).norm(2) - if self._swappable_optimizer_subgroup(sub_group_id): + 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) and not self.save_muon_momentum_buffer_in_memory: self._optimizer_states_and_gradient_swap_out(sub_group_id) @instrument_w_nvtx From b4ac49e41411374138f464a4e39a78495e4a3bd6 Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Tue, 8 Sep 2026 17:33:14 +0000 Subject: [PATCH 05/12] fix(zero): preserve Muon update scaling under CPU offload with loss scaling Signed-off-by: Jin, Youzhi --- deepspeed/runtime/zero/stage3.py | 13 +++++- deepspeed/runtime/zero/stage_1_and_2.py | 11 ++++- tests/unit/ops/muon/test_muon.py | 62 +++++++++++++++++++++++++ 3 files changed, 84 insertions(+), 2 deletions(-) diff --git a/deepspeed/runtime/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index 491781cf33d4..f47966720afa 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -2509,6 +2509,14 @@ def _apply_muon_updates_cpu_offload(self): local_grad_parts, self.gradient_accumulation_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): @@ -2525,7 +2533,10 @@ def _apply_muon_updates_cpu_offload(self): 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)] - fp32_param.grad.narrow(0, dest_offset, partition_numel).copy_(local_update.to(fp32_param.grad.dtype)) + # 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)) diff --git a/deepspeed/runtime/zero/stage_1_and_2.py b/deepspeed/runtime/zero/stage_1_and_2.py index 5d0cd573ff57..ec72f8e5f795 100644 --- a/deepspeed/runtime/zero/stage_1_and_2.py +++ b/deepspeed/runtime/zero/stage_1_and_2.py @@ -1778,6 +1778,13 @@ def _apply_muon_updates_cpu_offload(self): [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) @@ -1794,9 +1801,11 @@ def _apply_muon_updates_cpu_offload(self): 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_( - local_update.to(self.single_partition_of_fp32_groups[group_idx].grad.dtype)) + scaled_local_update.to(self.single_partition_of_fp32_groups[group_idx].grad.dtype)) self.norm_for_param_grads[param_id] = 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)) diff --git a/tests/unit/ops/muon/test_muon.py b/tests/unit/ops/muon/test_muon.py index 3fb14b211f91..630518f87d2f 100644 --- a/tests/unit/ops/muon/test_muon.py +++ b/tests/unit/ops/muon/test_muon.py @@ -390,3 +390,65 @@ 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.""" + + world_size = 2 + + @pytest.mark.parametrize("zero_stage", [1, 2]) + @pytest.mark.parametrize("loss_scale", [1.0, 1024.0]) + def test_offload_loss_scale_invariance(self, zero_stage, loss_scale): + from deepspeed.utils import safe_get_full_fp32_param + + hidden_dim, nlayers = 128, 2 + lr = 0.02 + torch.manual_seed(42) + model = SimpleModel(hidden_dim=hidden_dim, nlayers=nlayers) + + config_dict = { + "train_micro_batch_size_per_gpu": 4, + "gradient_clipping": 0.0, + "optimizer": { + "type": "muon", + "params": { + "lr": lr, + "momentum": 0.0 + } + }, + "fp16": { + "enabled": True, + "loss_scale": loss_scale + }, + "zero_optimization": { + "stage": zero_stage, + "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 p.ndim >= 2] + 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} + for n in pre: + applied_update_norm = ((pre[n] - post[n]) / lr).float().norm().item() + # Under double division bug, norm dropped to ~0.000024 for scale 1024. + # With fix, norm remains > 10.0 (Frobenius norm of ~128x128 orthogonal matrix is ~sqrt(128)~11.3). + assert applied_update_norm > 1.0, ( + f"Muon update vanished under loss_scale={loss_scale} (norm={applied_update_norm})! " + f"Double loss-scale division bug present.") From a12f7b67c81635f8b0b9d45b6fe58f1e7618d813 Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Tue, 8 Sep 2026 18:45:41 +0000 Subject: [PATCH 06/12] Fix Muon CPU offload gradient clipping norm and ZeRO-3 NVMe momentum residency - Fix gradient clipping norm accounting under loss scaling by storing scaled update norm in norm_for_param_grads for ZeRO-1/2/3. - Exclude resident ZeRO-3 Muon momentum buffer from OptimizerStateSwapInfo to prevent eviction by NVMe swapper. - Ensure swappable optimizer subgroups properly swap in and write back updated gradients and states in ZeRO-3 CPU offload. - Expand TestMuonOffloadLossScaling to ZeRO-1/2/3 with clipping equivalence validation across loss scales. - Add TestMuonZero3NVMeMomentumResidency for multi-step persistence of resident momentum under NVMe offload. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi --- .../runtime/swap_tensor/optimizer_utils.py | 2 + deepspeed/runtime/zero/stage3.py | 27 +-- deepspeed/runtime/zero/stage_1_and_2.py | 2 +- tests/unit/ops/muon/test_muon.py | 156 ++++++++++++++---- 4 files changed, 146 insertions(+), 41 deletions(-) diff --git a/deepspeed/runtime/swap_tensor/optimizer_utils.py b/deepspeed/runtime/swap_tensor/optimizer_utils.py index 191a85414b66..b52bf84401b4 100644 --- a/deepspeed/runtime/swap_tensor/optimizer_utils.py +++ b/deepspeed/runtime/swap_tensor/optimizer_utils.py @@ -211,6 +211,8 @@ def purge_state(self): 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: + 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) diff --git a/deepspeed/runtime/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index f47966720afa..36738bdaaa3f 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -1028,13 +1028,15 @@ def _create_momentum_buffer(self, num_elements, i, ds_id): device=self.device, dtype=self.communication_data_type) 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 _create_fp32_partitions(self): cpu_memory_usage = 0 @@ -2471,23 +2473,24 @@ def _apply_muon_updates_cpu_offload(self): if not muon_params: continue - if self._swappable_optimizer_subgroup(sub_group_id) and not self.save_muon_momentum_buffer_in_memory: + 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() != fp32_param.numel() + momentum_was_created = momentum is None or momentum.numel() != subgroup_numel if momentum_was_created: - self._create_momentum_buffer(fp32_param.numel(), sub_group_id, fp32_param.ds_id) + 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() != fp32_param.numel() + 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(fp32_param.numel(), sub_group_id, fp32_param.ds_id) + self._create_momentum_buffer(subgroup_numel, sub_group_id, fp32_param.ds_id) momentum = state["momentum_buffer"] local_grad_parts = [] @@ -2541,13 +2544,15 @@ def _apply_muon_updates_cpu_offload(self): 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)] = local_update.to(get_norm_dtype()).norm(2) + 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) and not self.save_muon_momentum_buffer_in_memory: - self._optimizer_states_and_gradient_swap_out(sub_group_id) + 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): diff --git a/deepspeed/runtime/zero/stage_1_and_2.py b/deepspeed/runtime/zero/stage_1_and_2.py index ec72f8e5f795..2cef7806b5ea 100644 --- a/deepspeed/runtime/zero/stage_1_and_2.py +++ b/deepspeed/runtime/zero/stage_1_and_2.py @@ -1806,7 +1806,7 @@ def _apply_muon_updates_cpu_offload(self): 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] = local_update.to(get_norm_dtype()).norm(2) + 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)) diff --git a/tests/unit/ops/muon/test_muon.py b/tests/unit/ops/muon/test_muon.py index 630518f87d2f..3084c9dce152 100644 --- a/tests/unit/ops/muon/test_muon.py +++ b/tests/unit/ops/muon/test_muon.py @@ -393,62 +393,160 @@ def test_update_matches_full_gradient_reference(self, zero_stage, ns_method, red class TestMuonOffloadLossScaling(DistributedTest): - """Verify Muon updates under CPU offload are invariant to loss_scale.""" + """Verify Muon updates under CPU offload are invariant to loss_scale with clipping.""" world_size = 2 - @pytest.mark.parametrize("zero_stage", [1, 2]) - @pytest.mark.parametrize("loss_scale", [1.0, 1024.0]) - def test_offload_loss_scale_invariance(self, zero_stage, loss_scale): + @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 - torch.manual_seed(42) - model = SimpleModel(hidden_dim=hidden_dim, nlayers=nlayers) + 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 p.ndim >= 2] + 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 + + def test_zero3_nvme_momentum_residency(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") + + hidden_dim, nlayers = 1024, 2 + lr = 0.01 + momentum = 0.95 config_dict = { - "train_micro_batch_size_per_gpu": 4, - "gradient_clipping": 0.0, + "train_micro_batch_size_per_gpu": 1, + "steps_per_print": 1, "optimizer": { "type": "muon", "params": { "lr": lr, - "momentum": 0.0 + "momentum": momentum } }, "fp16": { "enabled": True, - "loss_scale": loss_scale + "loss_scale": 1.0 }, "zero_optimization": { - "stage": zero_stage, + "stage": 3, + "reduce_scatter": False, + "save_muon_momentum_buffer_in_memory": True, "offload_optimizer": { - "device": "cpu", - "pin_memory": True - } + "device": "nvme", + "nvme_path": str(tmpdir) + }, + "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 - muon_named = [(n, p) for n, p in engine.module.named_parameters() if p.ndim >= 2] - pre = {n: safe_get_full_fp32_param(p).clone() for n, p in muon_named} + x = torch.randn(1, hidden_dim, device=device).half() + y = torch.randint(0, hidden_dim, (1, ), device=device) - 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) + opt = engine.optimizer + assert opt.swap_optimizer, "NVMe swap_optimizer must be enabled" + assert opt.save_muon_momentum_buffer_in_memory + + # Step 1 + engine.backward(engine(x, y)) engine.step() - post = {n: safe_get_full_fp32_param(p).clone() for n, p in muon_named} - for n in pre: - applied_update_norm = ((pre[n] - post[n]) / lr).float().norm().item() - # Under double division bug, norm dropped to ~0.000024 for scale 1024. - # With fix, norm remains > 10.0 (Frobenius norm of ~128x128 orthogonal matrix is ~sqrt(128)~11.3). - assert applied_update_norm > 1.0, ( - f"Muon update vanished under loss_scale={loss_scale} (norm={applied_update_norm})! " - f"Double loss-scale division bug present.") + 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]) From fce6c9718ac56ae0e9dbce1706275b2b1bbf12f1 Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Wed, 9 Sep 2026 06:00:11 +0000 Subject: [PATCH 07/12] Fix ZeRO-3 NVMe unswapped gradient fragment loss and support pipelined swapper - Retain unswapped gradient fragment ownership in OptimizerStateSwapInfo across Muon writeback and step until swap_out_optimizer_state. - Guard swapped gradient writing in writeback_optimizer_state_and_gradients when swapped_gradients is empty. - Implement writeback_optimizer_state_and_gradients and release_swap_buffers in PipelinedOptimizerSwapper. - Ensure synchronous swap-in without async prefetch during _apply_muon_updates_cpu_offload. - Expand TestMuonZero3NVMeMomentumResidency to cover both non-pipelined and pipelined NVMe swapping. - Add test_zero3_nvme_aggregate_unswapped_fragments for swappable subgroups composed of sub-MiB unswapped gradient fragments. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi --- .../runtime/swap_tensor/optimizer_utils.py | 4 +- .../partitioned_optimizer_swapper.py | 12 +-- .../pipelined_optimizer_swapper.py | 45 ++++++++++ deepspeed/runtime/zero/stage3.py | 3 +- tests/unit/ops/muon/test_muon.py | 90 +++++++++++++++++-- 5 files changed, 140 insertions(+), 14 deletions(-) diff --git a/deepspeed/runtime/swap_tensor/optimizer_utils.py b/deepspeed/runtime/swap_tensor/optimizer_utils.py index b52bf84401b4..244070fc96ec 100644 --- a/deepspeed/runtime/swap_tensor/optimizer_utils.py +++ b/deepspeed/runtime/swap_tensor/optimizer_utils.py @@ -207,6 +207,7 @@ 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 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" @@ -469,9 +470,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..41c02936f1a5 100644 --- a/deepspeed/runtime/swap_tensor/partitioned_optimizer_swapper.py +++ b/deepspeed/runtime/swap_tensor/partitioned_optimizer_swapper.py @@ -127,11 +127,12 @@ 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.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(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) @@ -149,6 +150,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..c96ed3102bc3 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,50 @@ 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(): + if param_info.swapped_gradients: + param_gradients = param_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.write_aio_handle, swap_buffers, swap_paths) + assert len(swap_buffers) == self.write_aio_handle.wait() + if param_info.unswapped_gradients: + param_info.write_unswapped_gradients(src_buffer=parameter.grad) + + 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 36738bdaaa3f..3bc2e5f36ee4 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -2598,7 +2598,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() diff --git a/tests/unit/ops/muon/test_muon.py b/tests/unit/ops/muon/test_muon.py index 3084c9dce152..d031cde29ff3 100644 --- a/tests/unit/ops/muon/test_muon.py +++ b/tests/unit/ops/muon/test_muon.py @@ -469,14 +469,26 @@ class TestMuonZero3NVMeMomentumResidency(DistributedTest): world_size = 1 - def test_zero3_nvme_momentum_residency(self, tmpdir): + @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, @@ -495,10 +507,7 @@ def test_zero3_nvme_momentum_residency(self, tmpdir): "stage": 3, "reduce_scatter": False, "save_muon_momentum_buffer_in_memory": True, - "offload_optimizer": { - "device": "nvme", - "nvme_path": str(tmpdir) - }, + "offload_optimizer": offload_optimizer_cfg, "sub_group_size": 100 }, "aio": { @@ -519,6 +528,10 @@ def test_zero3_nvme_momentum_residency(self, tmpdir): 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)) @@ -550,3 +563,70 @@ def test_zero3_nvme_momentum_residency(self, tmpdir): 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]) + + 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())) From a55e543812a792bb8d95e2398af6a61adf709924 Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Wed, 9 Sep 2026 06:21:56 +0000 Subject: [PATCH 08/12] Add ZeRO-3 NVMe mixed swapped/unswapped numerical equivalence test - Add test_zero3_nvme_mixed_fragments_numerical_equivalence to TestMuonZero3NVMeMomentumResidency. - Construct a single subgroup containing both >= 1 MiB (swapped) and < 1 MiB (unswapped) parameters. - Verify multi-step numerical equivalence against a non-NVMe CPU offload reference across both partitioned and pipelined swappers. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi --- tests/unit/ops/muon/test_muon.py | 103 +++++++++++++++++++++++++++++++ 1 file changed, 103 insertions(+) diff --git a/tests/unit/ops/muon/test_muon.py b/tests/unit/ops/muon/test_muon.py index d031cde29ff3..9af577a0bb3c 100644 --- a/tests/unit/ops/muon/test_muon.py +++ b/tests/unit/ops/muon/test_muon.py @@ -630,3 +630,106 @@ def forward(self, x): 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): + 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 + + 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): + return self.small(self.proj(self.large(x))).sum() + + lr = 0.01 + num_steps = 2 + + def _run_experiment(offload_cfg): + torch.manual_seed(42) + model = MixedMatrixModel() + config_dict = { + "train_micro_batch_size_per_gpu": 1, + "steps_per_print": 1, + "optimizer": { + "type": "muon", + "params": { + "lr": lr, + "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": offload_cfg, + "sub_group_size": 1000000 + }, + "aio": { + "block_size": 1048576 + } + } + 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 p.ndim >= 2] + init_params = {n: safe_get_full_fp32_param(p).clone() for n, p in muon_named} + + torch.manual_seed(1234) + for _ in range(num_steps): + x = torch.randn(2, 512, device=device).half() + loss = engine(x) + engine.backward(loss) + engine.step() + + final_params = {n: safe_get_full_fp32_param(p).clone() for n, p in muon_named} + return init_params, final_params, engine.optimizer + + # 1. Non-NVMe reference run (CPU offload in RAM, no swapping) + ref_cfg = {"device": "cpu", "pin_memory": True} + init_params, ref_final, _ = _run_experiment(ref_cfg) + + # 2. 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 + + _, nvme_final, opt = _run_experiment(nvme_cfg) + + assert opt.swap_optimizer + if pipeline: + assert isinstance(opt.optimizer_swapper, PipelinedOptimizerSwapper) + else: + assert isinstance(opt.optimizer_swapper, PartitionedOptimizerSwapper) + + # Verify numerical equivalence with non-NVMe reference and non-zero update + for n in ref_final: + update_norm = (init_params[n] - nvme_final[n]).norm().item() + assert update_norm > 0.01, f"Parameter {n} did not update" + max_diff = (ref_final[n] - nvme_final[n]).abs().max().item() + assert max_diff < 1e-4, ( + f"Numerical divergence for {n} under pipeline={pipeline}: " + f"max_diff={max_diff}, ref_norm={ref_final[n].norm().item()}, nvme_norm={nvme_final[n].norm().item()}") From 12fa5dedb6230d12cb40985e06dd6744f84aaf68 Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Wed, 9 Sep 2026 07:18:47 +0000 Subject: [PATCH 09/12] Fix ZeRO-3 Muon parameter selection and use independent full-gradient reference - Fix Muon parameter selection in ZeRO-3 by checking getattr(p, 'use_muon', False) and p.ds_shape instead of p.ndim. - Assert all 3 mixed-fragment parameters (large, proj, small) are identified and tracked. - Replace shared CPU-offload reference with an independent pure-PyTorch full-gradient reference maintaining momentum across steps. - Fix parameter selection in TestMuonOffloadLossScaling for ZeRO-3. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi --- tests/unit/ops/muon/test_muon.py | 140 ++++++++++++++++++------------- 1 file changed, 83 insertions(+), 57 deletions(-) diff --git a/tests/unit/ops/muon/test_muon.py b/tests/unit/ops/muon/test_muon.py index 9af577a0bb3c..2dafc47b6795 100644 --- a/tests/unit/ops/muon/test_muon.py +++ b/tests/unit/ops/muon/test_muon.py @@ -436,7 +436,9 @@ def _run_with_scale(loss_scale): model_parameters=model.parameters(), dist_init_required=False) device = engine.device - muon_named = [(n, p) for n, p in engine.module.named_parameters() if p.ndim >= 2] + 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) @@ -633,10 +635,12 @@ def forward(self, x): @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") @@ -656,59 +660,34 @@ def forward(self, x): return self.small(self.proj(self.large(x))).sum() lr = 0.01 + momentum = 0.95 num_steps = 2 - def _run_experiment(offload_cfg): - torch.manual_seed(42) - model = MixedMatrixModel() - config_dict = { - "train_micro_batch_size_per_gpu": 1, - "steps_per_print": 1, - "optimizer": { - "type": "muon", - "params": { - "lr": lr, - "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": offload_cfg, - "sub_group_size": 1000000 - }, - "aio": { - "block_size": 1048576 - } - } - 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 p.ndim >= 2] - init_params = {n: safe_get_full_fp32_param(p).clone() for n, p in muon_named} - - torch.manual_seed(1234) - for _ in range(num_steps): - x = torch.randn(2, 512, device=device).half() - loss = engine(x) - engine.backward(loss) - engine.step() - - final_params = {n: safe_get_full_fp32_param(p).clone() for n, p in muon_named} - return init_params, final_params, engine.optimizer - - # 1. Non-NVMe reference run (CPU offload in RAM, no swapping) - ref_cfg = {"device": "cpu", "pin_memory": True} - init_params, ref_final, _ = _run_experiment(ref_cfg) - - # 2. NVMe run (with mixed swapped/unswapped fragments in the same subgroup) + 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_model = MixedMatrixModel().to(device).half() + ref_model.load_state_dict({k: v.to(device).half() for k, v in init_state.items()}) + ref_momentums = {n: torch.zeros_like(p) for n, p in ref_model.named_parameters()} + + gen = torch.Generator().manual_seed(1234) + inputs = [torch.randn(2, 512, generator=gen).to(device).half() for _ in range(num_steps)] + + for step in range(num_steps): + 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.clone(), ref_momentums[n], beta=momentum, ns_method="gram") + p.add_(update.reshape(p.shape), alpha=-lr) + + ref_final = {n: p.clone().detach().cpu().float() for n, p in ref_model.named_parameters()} + + # 2. ZeRO-3 NVMe run with mixed swapped/unswapped fragments in the same subgroup nvme_cfg = { "device": "nvme", "nvme_path": str(tmpdir), @@ -717,19 +696,66 @@ def _run_experiment(offload_cfg): nvme_cfg["pipeline_read"] = True nvme_cfg["pipeline_write"] = True - _, nvme_final, opt = _run_experiment(nvme_cfg) + config_dict = { + "train_micro_batch_size_per_gpu": 1, + "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) - # Verify numerical equivalence with non-NVMe reference and non-zero update + nvme_final = {n: safe_get_full_fp32_param(p).clone().cpu() for n, p in muon_named} + + # Verify non-trivial update and numerical equivalence against independent reference for n in ref_final: update_norm = (init_params[n] - nvme_final[n]).norm().item() assert update_norm > 0.01, f"Parameter {n} did not update" - max_diff = (ref_final[n] - nvme_final[n]).abs().max().item() - assert max_diff < 1e-4, ( + rel_err = ((nvme_final[n] - ref_final[n]).norm() / (ref_final[n].norm() + 1e-8)).item() + assert rel_err < 0.05, ( f"Numerical divergence for {n} under pipeline={pipeline}: " - f"max_diff={max_diff}, ref_norm={ref_final[n].norm().item()}, nvme_norm={nvme_final[n].norm().item()}") + f"rel_err={rel_err:.4f}, ref_norm={ref_final[n].norm().item():.4f}, nvme_norm={nvme_final[n].norm().item():.4f}" + ) From 9f74ac9b9e9b4d74850feaf4b25183291b44de74 Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Wed, 9 Sep 2026 08:32:29 +0000 Subject: [PATCH 10/12] Add 2-rank ZeRO-3 NVMe Muon test with DP gradient averaging oracle - Add TestMuonZero3NVMeMultiRankMixedFragments with world_size=2 covering cross-rank all-gather, partition slicing, and DP gradient averaging. - Size parameters so per-rank partition includes both >= 1 MiB (swapped) and < 1 MiB (unswapped) fragments in the same subgroup. - Construct independent high-precision reference holding FP32 master weights/momentum and averaging per-rank FP16 gradients. - Compare actual update tensor (init - final) directly against reference update using ref_update.norm() as denominator. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi --- tests/unit/ops/muon/test_muon.py | 190 ++++++++++++++++++++++++++++--- 1 file changed, 177 insertions(+), 13 deletions(-) diff --git a/tests/unit/ops/muon/test_muon.py b/tests/unit/ops/muon/test_muon.py index 2dafc47b6795..a88ba90fe5b4 100644 --- a/tests/unit/ops/muon/test_muon.py +++ b/tests/unit/ops/muon/test_muon.py @@ -669,23 +669,26 @@ def forward(self, x): # 1. Independent full-gradient reference (pure PyTorch + canonical muon_update across 2 steps) device = get_accelerator().current_device_name() - ref_model = MixedMatrixModel().to(device).half() - ref_model.load_state_dict({k: v.to(device).half() for k, v in init_state.items()}) - ref_momentums = {n: torch.zeros_like(p) for n, p in ref_model.named_parameters()} + 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()} + ref_momentums = {n: torch.zeros_like(p).to(device).half() for n, p in ref_masters.items()} gen = torch.Generator().manual_seed(1234) inputs = [torch.randn(2, 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.clone(), ref_momentums[n], beta=momentum, ns_method="gram") - p.add_(update.reshape(p.shape), alpha=-lr) + ref_masters[n].add_(up := update.cpu().float(), alpha=-lr) - ref_final = {n: p.clone().detach().cpu().float() for n, p in ref_model.named_parameters()} + 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 = { @@ -749,13 +752,174 @@ def forward(self, x): 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_final: - update_norm = (init_params[n] - nvme_final[n]).norm().item() - assert update_norm > 0.01, f"Parameter {n} did not update" - rel_err = ((nvme_final[n] - ref_final[n]).norm() / (ref_final[n].norm() + 1e-8)).item() - assert rel_err < 0.05, ( - f"Numerical divergence for {n} under pipeline={pipeline}: " - f"rel_err={rel_err:.4f}, ref_norm={ref_final[n].norm().item():.4f}, nvme_norm={nvme_final[n].norm().item():.4f}" - ) + 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.25, (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, 128, bias=False) + + def forward(self, x): + return self.small(self.proj(self.large(x))).sum() + + lr = 0.01 + momentum = 0.95 + num_steps = 2 + micro_batch = 2 + 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()} + ref_momentums = {n: torch.zeros_like(p).to(device).half() 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()} + + avg_grads = {n: (grads0[n] + grads1[n]) / 2.0 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) == 3 + assert set(n for n, p in muon_named) == {"large.weight", "proj.weight", "small.weight"} + 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.25, ( + f"Update divergence for {n} under pipeline={pipeline}: " + f"rel_err={rel_err:.4f}, applied_norm={applied_norm:.4f}, ref_norm={ref_norm:.4f}") From 6cf979b565aac9bf5d000519eb177c7f9244406f Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Wed, 9 Sep 2026 10:22:06 +0000 Subject: [PATCH 11/12] Keep ZeRO-3 Muon momentum in fp32 and tighten the NVMe oracles ZeRO-3 allocated the Muon momentum buffer with communication_data_type, so under an fp16 config the accumulator was rounded to fp16 on every step even though the update itself was computed in fp32. ZeRO-1/2 derive the buffer from the fp32 master partition via zeros_like(flat_param), and the NVMe optimizer swapper already stores state in master_weights_and_grads_dtype, so fp16 momentum was both lossy and inconsistent with the rest of the stack. Allocate the momentum buffer in the master dtype, gather it (and the gradients it is blended with) in that dtype, and promote the gradient before Newton-Schulz in the non-offload path. The oracles in the ZeRO-3 NVMe tests created fp16 momentum too, so they reproduced and accepted the same drift; they now keep fp32 momentum. The 2-rank test could not observe ZeRO-3 padding because every Muon matrix had an element count divisible by the world size. Add an odd-numel matrix, assert that it really produces a padded final partition, and cover the padding sensitive reconstruction and slicing paths. Both NVMe tests also allowed 25% update-relative error, which was wide enough to hide a tail reconstruction or second-step momentum bug. The error was not inherent: a sum() loss over a linear stack yields rank-1 gradients, and Newton-Schulz amplifies fp16 noise in the near-null directions by roughly the fifth power of its slope at zero. Switching to a squared loss over a wide batch (scaled out of the fp16 subnormal range) and averaging DP gradients in the fp16 communication dtype, as DeepSpeed does, drops the observed error from 0.12-0.24 to under 0.02, so the bound is now 0.05. Verified on 8x Intel Battlemage (XPU): 161 passed in tests/unit/ops/muon, and 121 passed / 61 skipped in the ZeRO NVMe checkpointing and tensor fragment suites. Also refine the surrounding Muon work: drop an unused local in the ZeRO-1/2 offload path, collapse the triplicated momentum lookup and writeback branches in _apply_distributed_muon_update, and hoist the duplicated gradient writeback in the two optimizer swappers into a shared OptimizerSwapper helper. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi --- .../runtime/swap_tensor/optimizer_utils.py | 13 +++ .../partitioned_optimizer_swapper.py | 9 +- .../pipelined_optimizer_swapper.py | 9 +- deepspeed/runtime/zero/stage3.py | 88 +++++++++---------- deepspeed/runtime/zero/stage_1_and_2.py | 1 - tests/unit/ops/muon/test_muon.py | 51 +++++++---- 6 files changed, 92 insertions(+), 79 deletions(-) diff --git a/deepspeed/runtime/swap_tensor/optimizer_utils.py b/deepspeed/runtime/swap_tensor/optimizer_utils.py index 244070fc96ec..847d558e376e 100644 --- a/deepspeed/runtime/swap_tensor/optimizer_utils.py +++ b/deepspeed/runtime/swap_tensor/optimizer_utils.py @@ -209,9 +209,22 @@ def purge_state(self): 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) diff --git a/deepspeed/runtime/swap_tensor/partitioned_optimizer_swapper.py b/deepspeed/runtime/swap_tensor/partitioned_optimizer_swapper.py index 41c02936f1a5..46f34a684af3 100644 --- a/deepspeed/runtime/swap_tensor/partitioned_optimizer_swapper.py +++ b/deepspeed/runtime/swap_tensor/partitioned_optimizer_swapper.py @@ -127,14 +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(): - 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(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) diff --git a/deepspeed/runtime/swap_tensor/pipelined_optimizer_swapper.py b/deepspeed/runtime/swap_tensor/pipelined_optimizer_swapper.py index c96ed3102bc3..b00f6a75bc21 100644 --- a/deepspeed/runtime/swap_tensor/pipelined_optimizer_swapper.py +++ b/deepspeed/runtime/swap_tensor/pipelined_optimizer_swapper.py @@ -182,14 +182,7 @@ def writeback_optimizer_state_and_gradients(self, parameter, write_opt_state, wr self.swap_buffer_manager.free(swap_in_op.allocated_buffers) if write_gradients and param_info.has_gradients(): - if param_info.swapped_gradients: - param_gradients = param_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.write_aio_handle, swap_buffers, swap_paths) - assert len(swap_buffers) == self.write_aio_handle.wait() - if param_info.unswapped_gradients: - param_info.write_unswapped_gradients(src_buffer=parameter.grad) + self._writeback_gradients(param_info, parameter, self.write_aio_handle) param_info.release_memory() self.swap_ops[SYNC_SWAP_IN] = None diff --git a/deepspeed/runtime/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index 3bc2e5f36ee4..2a7880c2591c 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -1024,9 +1024,12 @@ 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 @@ -1658,31 +1661,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) @@ -1698,7 +1698,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], @@ -1719,28 +1724,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: @@ -2502,16 +2497,15 @@ def _apply_muon_updates_cpu_offload(self): 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, - self.gradient_accumulation_dtype) + 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, - self.gradient_accumulation_dtype, - additional_buffers_to_allgather=local_momentum_parts) + 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. diff --git a/deepspeed/runtime/zero/stage_1_and_2.py b/deepspeed/runtime/zero/stage_1_and_2.py index 2cef7806b5ea..d8ae60307771 100644 --- a/deepspeed/runtime/zero/stage_1_and_2.py +++ b/deepspeed/runtime/zero/stage_1_and_2.py @@ -1756,7 +1756,6 @@ def _apply_muon_updates_cpu_offload(self): continue process_group = self.real_dp_process_group[group_idx] - world_size = dist.get_world_size(group=process_group) 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 diff --git a/tests/unit/ops/muon/test_muon.py b/tests/unit/ops/muon/test_muon.py index a88ba90fe5b4..a15be0620aa6 100644 --- a/tests/unit/ops/muon/test_muon.py +++ b/tests/unit/ops/muon/test_muon.py @@ -657,11 +657,15 @@ def __init__(self): self.small = torch.nn.Linear(128, 128, bias=False) def forward(self, x): - return self.small(self.proj(self.large(x))).sum() + # 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() @@ -671,10 +675,11 @@ def forward(self, x): 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()} - ref_momentums = {n: torch.zeros_like(p).to(device).half() 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(2, 512, generator=gen).to(device).half() for _ in range(num_steps)] + 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() @@ -685,8 +690,8 @@ def forward(self, x): loss.backward() with torch.no_grad(): for n, p in ref_model.named_parameters(): - update = muon_update(p.grad.clone(), ref_momentums[n], beta=momentum, ns_method="gram") - ref_masters[n].add_(up := update.cpu().float(), alpha=-lr) + 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} @@ -700,7 +705,7 @@ def forward(self, x): nvme_cfg["pipeline_write"] = True config_dict = { - "train_micro_batch_size_per_gpu": 1, + "train_micro_batch_size_per_gpu": micro_batch, "steps_per_print": 1, "gradient_clipping": 0.0, "optimizer": { @@ -761,7 +766,7 @@ def forward(self, x): 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.25, (f"Numerical divergence for {n} under pipeline={pipeline}: " + 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}") @@ -794,15 +799,20 @@ def __init__(self): 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, 128, bias=False) + 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): - return self.small(self.proj(self.large(x))).sum() + # 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 = 2 + micro_batch = 256 rank = dist.get_rank() device = get_accelerator().current_device_name() @@ -820,7 +830,8 @@ def forward(self, x): 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()} - ref_momentums = {n: torch.zeros_like(p).to(device).half() 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] @@ -839,7 +850,9 @@ def forward(self, x): loss1.backward() grads1 = {n: p.grad.clone() for n, p in ref_model.named_parameters()} - avg_grads = {n: (grads0[n] + grads1[n]) / 2.0 for n in grads0} + # 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") @@ -891,8 +904,16 @@ def forward(self, x): 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 - assert set(n for n, p in muon_named) == {"large.weight", "proj.weight", "small.weight"} + 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 @@ -920,6 +941,6 @@ def forward(self, x): 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.25, ( + 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}") From 5a596a5da0d6f53e7b3f4f4cee0787ca02a7ed38 Mon Sep 17 00:00:00 2001 From: "Jin, Youzhi" Date: Wed, 9 Sep 2026 12:02:27 +0000 Subject: [PATCH 12/12] Restore ZeRO-3 resident Muon momentum when loading a checkpoint The in-memory Muon momentum buffer (save_muon_momentum_buffer_in_memory) is held twice: in optimizer.state and in the muon_momentum_buffer_partitioned_groups_flat cache that the Muon update path actually reads. Neither checkpoint load path kept the two in sync: * _rigid_load_state_dict() calls Optimizer.load_state_dict(), which rebinds optimizer.state to fresh tensors. The cache kept serving the pre-load (zero) momentum and wrote it straight back over the restored state on the next step. * NVMe optimizer offload skips the ZeRO optimizer state dict entirely and rebuilds state by copying swap files. A resident buffer is excluded from swapping by design, so it never reaches those files and the restored value was dropped outright. Copy the checkpointed momentum into the resident buffer and re-bind optimizer.state to it in both paths, so a resumed run continues from the momentum it saved. Add an interrupted-vs-uninterrupted checkpoint test over the non-offload path plus the partitioned and pipelined NVMe swappers. It asserts the restored momentum, the shared tensor identity, and that the post-resume weight update matches the uninterrupted run. Without the fix all three variants fail, and the end-to-end update error is 0.91 instead of 0. Validated on 8x Intel Battlemage XPU: tests/unit/ops/muon/test_muon.py 164 passed; test_nvme_checkpointing.py + test_zero_optimizer.py 59 passed (10 pre-existing missing-fixture collection errors). Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi --- deepspeed/runtime/engine.py | 18 ++++ deepspeed/runtime/zero/stage3.py | 62 +++++++++++++ tests/unit/ops/muon/test_muon.py | 152 +++++++++++++++++++++++++++++++ 3 files changed, 232 insertions(+) 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/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index 2a7880c2591c..31d44ef8824c 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -1041,6 +1041,67 @@ def _create_momentum_buffer(self, num_elements, i, ds_id): self.optimizer.state[ self.fp32_partitioned_groups_flat[i]]["momentum_buffer"] = unpinned_fp32_buffer_momentum + 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 cpu_memory_sub_groups = 0 @@ -3423,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/tests/unit/ops/muon/test_muon.py b/tests/unit/ops/muon/test_muon.py index a15be0620aa6..fdde19ea242d 100644 --- a/tests/unit/ops/muon/test_muon.py +++ b/tests/unit/ops/muon/test_muon.py @@ -566,6 +566,158 @@ def test_zero3_nvme_momentum_residency(self, tmpdir, pipeline): 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]: