Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
194 changes: 108 additions & 86 deletions deepspeed/runtime/zero/stage3.py
Original file line number Diff line number Diff line change
Expand Up @@ -1665,95 +1665,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.
Expand Down
Loading