diff --git a/deepspeed/runtime/zero/stage3.py b/deepspeed/runtime/zero/stage3.py index 97fc13d29e5f..c58af3ba6449 100644 --- a/deepspeed/runtime/zero/stage3.py +++ b/deepspeed/runtime/zero/stage3.py @@ -1633,95 +1633,117 @@ def _muon_update_sub_group(self, i, group_items, communication_data_type: torch. """ params = [param for param, _, _ in group_items] grads = [grad for _, _, grad in group_items] - if params: - momentum_buffer = [] - if self._swappable_optimizer_subgroup(i) and not self.save_muon_momentum_buffer_in_memory: - # 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: - 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()) + if not params: + return + + momentum_buffer = [] + if self._swappable_optimizer_subgroup(i) and not self.save_muon_momentum_buffer_in_memory: + # 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: + 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()) + + gathered_params_momentums = self._partitioned_buffers_all_gather(params, momentum_buffer, + communication_data_type) + + process_group = self._get_sub_group_process_group(i) + world_sz = dist.get_world_size(process_group) + rank = dist.get_rank(process_group) + # A round can mix projection shapes under GQA, so use padded flat tensors when ranks do not + # contribute equally sized grads. Process one round at a time to bound temporary memory. + for base_i in range(0, len(params), world_sz): + round_params = params[base_i:base_i + world_sz] + round_grads = grads[base_i:base_i + world_sz] + round_momentums = gathered_params_momentums[base_i:base_i + world_sz] + round_count = len(round_grads) + numels = [g.numel() for g in round_grads] + uniform_shape = round_count == world_sz and all(g.shape == round_grads[0].shape for g in round_grads) + + if rank < round_count: + param = round_params[rank] + g = round_grads[rank] + m = round_momentums[rank] + update = muon_update(g, + m, + beta=self.muon_beta, + ns_method=getattr(self, 'muon_ns_method', 'gram'), + num_heads=getattr(param, 'muon_num_heads', None)) + g.data.copy_(update, non_blocking=False) + + if uniform_shape: + local_grad = round_grads[rank] + local_momentum = round_momentums[rank] + grad_out = round_grads + momentum_out = round_momentums 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()) - - gathered_params_momentums = self._partitioned_buffers_all_gather(params, momentum_buffer, - communication_data_type) - - process_group = self._get_sub_group_process_group(i) - world_sz = dist.get_world_size(process_group) - rank = dist.get_rank(process_group) - grads_pad = grads + [torch.empty_like(grads[-1])] * ((world_sz - len(params) % world_sz) % world_sz) - gathered_momentums_pad = gathered_params_momentums + [torch.empty_like(gathered_params_momentums[-1])] * ( - (world_sz - len(gathered_params_momentums) % world_sz) % world_sz) - grad_handles = [] - momentum_handles = [] - for base_i in range(len(params))[::world_sz]: - if base_i + rank < len(params): - param = params[base_i + rank] - g = grads[base_i + rank] - m = gathered_momentums_pad[base_i + rank] - update = muon_update(g, - m, - beta=self.muon_beta, - ns_method=getattr(self, 'muon_ns_method', 'gram'), - num_heads=getattr(param, 'muon_num_heads', None)) - 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], - group=process_group, - async_op=True) - grad_handles.append(grad_handle) - momentum_handle = dist.all_gather(gathered_momentums_pad[base_i:base_i + world_sz], - gathered_momentums_pad[base_i + rank], - group=process_group, - async_op=True) - momentum_handles.append(momentum_handle) - - for handle in momentum_handles: - handle.wait() - for idx, (param, dest_offset, grad) in enumerate(group_items): - gathered_momentum = gathered_params_momentums[idx] - chunk_sz = math.ceil(grad.numel() / world_sz) - start_offset = rank * chunk_sz - end_offset = start_offset + chunk_sz - if end_offset > grad.numel(): - buffer_to_update = torch.zeros(chunk_sz, - device=grad.device, - dtype=self.gradient_accumulation_dtype) - buffer_to_update[:grad.numel() - - start_offset] = gathered_momentum.view(-1).data[start_offset: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] + max_numel = max(numels) + if rank < round_count: + local_grad = torch.zeros(max_numel, dtype=g.dtype, device=g.device) + local_grad[:g.numel()] = g.view(-1) + local_momentum = torch.zeros(max_numel, dtype=m.dtype, device=m.device) + local_momentum[:m.numel()] = m.view(-1) 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) + # This rank has no grad to contribute in this round; participate with a dummy + # tensor of the agreed-upon shape so the collective still completes. + ref_grad = round_grads[0] + ref_momentum = round_momentums[0] + local_grad = torch.zeros(max_numel, dtype=ref_grad.dtype, device=ref_grad.device) + local_momentum = torch.zeros(max_numel, dtype=ref_momentum.dtype, device=ref_momentum.device) + grad_out = [torch.empty_like(local_grad) for _ in range(world_sz)] + momentum_out = [torch.empty_like(local_momentum) for _ in range(world_sz)] + + grad_handle = dist.all_gather(grad_out, local_grad, group=process_group, async_op=True) + momentum_handle = dist.all_gather(momentum_out, local_momentum, group=process_group, async_op=True) + momentum_handle.wait() + grad_handle.wait() + if not uniform_shape: + for grad, momentum, gathered_grad, gathered_momentum in zip(round_grads, round_momentums, grad_out, + momentum_out): + grad.data.copy_(gathered_grad[:grad.numel()].view_as(grad), non_blocking=False) + momentum.view(-1).data.copy_(gathered_momentum[:momentum.numel()], non_blocking=False) + + for idx, (param, dest_offset, grad) in enumerate(group_items): + gathered_momentum = gathered_params_momentums[idx] + chunk_sz = math.ceil(grad.numel() / world_sz) + start_offset = rank * chunk_sz + end_offset = start_offset + chunk_sz + if end_offset > grad.numel(): + buffer_to_update = torch.zeros(chunk_sz, device=grad.device, dtype=self.gradient_accumulation_dtype) + buffer_to_update[:grad.numel() - + start_offset] = gathered_momentum.view(-1).data[start_offset: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_swapper.swap_out_optimizer_state(parameter=self.fp32_partitioned_groups_flat[i]) - for handle in grad_handles: - handle.wait() + 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]) def _apply_muon_to_accumulated_grads(self): """Orthogonalize each Muon parameter's accumulated gradient once per optimizer step.