diff --git a/deepspeed/runtime/zero/stage_1_and_2.py b/deepspeed/runtime/zero/stage_1_and_2.py index 922fe18b5184..8a4db6454928 100644 --- a/deepspeed/runtime/zero/stage_1_and_2.py +++ b/deepspeed/runtime/zero/stage_1_and_2.py @@ -139,6 +139,13 @@ def clear(self): self.copy_streams.clear() +class GradPartitionBufferGroup(list): + + def __init__(self, grad_buffers, flat_partition): + super().__init__(grad_buffers) + self.flat_partition = flat_partition + + class DeepSpeedZeroOptimizer(ZeROOptimizer): """ DeepSpeedZeroOptimizer designed to reduce the memory footprint @@ -1015,6 +1022,10 @@ def independent_gradient_partition_epilogue(self): if self.cpu_offload is False: for i, _ in enumerate(self.bit16_groups): + if self._use_compact_grad_partition_buffer(i): + if self.is_gradient_accumulation_boundary() and self.averaged_gradients.get(i) is None: + self.averaged_gradients[i] = self._build_compact_grad_partition_buffer(i) + continue if i not in self.all_grad_tensors or self.all_grad_tensors[i] is None: self.all_grad_tensors[i] = self.get_all_grad_tensors(self.params_in_partition[i], dtype=self.gradient_accumulation_dtype) @@ -1057,6 +1068,11 @@ def finalize_gradient_accumulation_boundary(self): self._finalize_cpu_offload_gradient_accumulation() return for i, _ in enumerate(self.bit16_groups): + if self._use_compact_grad_partition_buffer(i): + if self.averaged_gradients.get(i) is None: + self.averaged_gradients[i] = self._build_compact_grad_partition_buffer(i) + self.all_grad_tensors[i] = None + continue self.averaged_gradients[i] = self.get_flat_partition(self.params_in_partition[i], self.first_offset[i], self.partition_size[i], @@ -1218,13 +1234,103 @@ def fill_grad_accum_attribute(self): def get_gradient_for_reduction(self, param): if self.use_grad_accum_attribute: - return param.grad_accum.to(self.dtype) if param.grad_accum is not None else None + return param.grad_accum if param.grad_accum is not None else None else: return param.grad + def _gradient_reduction_buffer_dtype(self): + return self.gradient_accumulation_dtype if self.use_grad_accum_attribute else self.dtype + def get_param_gradient_attribute(self, param): return param.grad_accum if self.use_grad_accum_attribute else param.grad + def _use_compact_grad_partition_buffer(self, group_idx): + if (self.cpu_offload or self.gradient_accumulation_dtype != torch.float32 + or self.gradient_accumulation_dtype == self.dtype): + return False + if not self.contiguous_gradients: + return False + params_in_partition = self.params_in_partition[group_idx] + if len(params_in_partition) > 0 and self._is_muon_group(params_in_partition): + return False + return True + + def _build_compact_grad_partition_buffer(self, group_idx): + partition_size = int(self.partition_size[group_idx]) + flat_partition = torch.zeros(partition_size, + dtype=self.gradient_accumulation_dtype, + device=get_accelerator().current_device_name()) + partition_id = dist.get_rank(group=self.real_dp_process_group[group_idx]) + + grad_buffers = [] + current_size = 0 + for tensor in self.params_in_partition[group_idx]: + param_id = self.get_param_id(tensor) + source_offset = int(self.grad_start_offset[group_idx][partition_id][param_id]) + dest_offset = int(self.grad_partition_insertion_offset[group_idx][partition_id][param_id]) + num_elements = int(min(tensor.numel() - source_offset, partition_size - dest_offset)) + + if dest_offset > current_size: + pad = flat_partition.narrow(0, current_size, dest_offset - current_size) + pad._zero_padding = True + grad_buffers.append(pad) + current_size = dest_offset + + if num_elements > 0: + grad_buffer = flat_partition.narrow(0, dest_offset, num_elements) + if source_offset == 0 and num_elements == tensor.numel(): + grad_buffer = grad_buffer.view(tensor.shape) + grad_buffers.append(grad_buffer) + current_size = dest_offset + num_elements + + if current_size < partition_size: + pad = flat_partition.narrow(0, current_size, partition_size - current_size) + pad._zero_padding = True + grad_buffers.append(pad) + + return GradPartitionBufferGroup(grad_buffers, flat_partition) + + def _get_compact_grad_partition_buffer(self, group_idx): + grad_group = self.averaged_gradients.get(group_idx) + if grad_group is not None and getattr(grad_group, "flat_partition", None) is not None: + return grad_group + + grad_group = self._build_compact_grad_partition_buffer(group_idx) + self.averaged_gradients[group_idx] = grad_group + return grad_group + + def _get_compact_grad_partition_slice(self, group_idx, partition_id, param_id, num_elements): + if not self._use_compact_grad_partition_buffer(group_idx): + return None + local_partition_id = dist.get_rank(group=self.real_dp_process_group[group_idx]) + if partition_id != local_partition_id: + return None + + dest_offset = int(self.grad_partition_insertion_offset[group_idx][partition_id][param_id]) + grad_group = self._get_compact_grad_partition_buffer(group_idx) + return grad_group.flat_partition.narrow(0, dest_offset, int(num_elements)) + + @staticmethod + def _add_to_grad_partition_buffer(output, source): + if output is not None: + output.add_(source.to(output.dtype)) + + def _add_full_reduction_to_compact_grad_partition(self, reduced_tensor, communication_data_type): + bucket = self.ipg_buckets[communication_data_type] + bucket_offset = 0 + for group_idx, param_idx_in_group, param_id in bucket.params: + param = self.bit16_groups[group_idx][param_idx_in_group] + if self._use_compact_grad_partition_buffer(group_idx): + partition_id = dist.get_rank(group=self.real_dp_process_group[group_idx]) + source_offset = int(self.grad_start_offset[group_idx][partition_id][param_id]) + dest_offset = int(self.grad_partition_insertion_offset[group_idx][partition_id][param_id]) + num_elements = int(min(param.numel() - source_offset, self.partition_size[group_idx] - dest_offset)) + if num_elements > 0: + source = reduced_tensor.narrow(0, bucket_offset + source_offset, num_elements) + output = self._get_compact_grad_partition_slice(group_idx, partition_id, param_id, num_elements) + self._add_to_grad_partition_buffer(output, source) + bucket_offset += param.numel() + # Clear the tensor the reduction gradient attribute is pointing to def clear_grad_attribute(self, param): if self.use_grad_accum_attribute: @@ -1392,7 +1498,7 @@ def gradient_reduction_w_predivide(self, tensor, communication_data_type: torch. if communication_data_type != tensor.dtype and tensor is not tensor_to_allreduce: tensor.copy_(tensor_to_allreduce) - return tensor + return tensor_to_allreduce def allreduce_and_copy_with_multiple_ranks(self, small_bucket, @@ -1400,7 +1506,8 @@ def allreduce_and_copy_with_multiple_ranks(self, log=None, divide=True, process_group=None, - bucket_ranks=None): + bucket_ranks=None, + output_bucket=None): process_group = self.dp_process_group if process_group is None else process_group allreduced = self.allreduce_bucket(small_bucket, communication_data_type, @@ -1409,11 +1516,15 @@ def allreduce_and_copy_with_multiple_ranks(self, process_group=process_group) if self.overlap_comm and not get_accelerator().resolves_data_dependency(): allreduced.record_stream(self.reduction_stream) + synced_tensors = self.unflatten(allreduced, small_bucket) + if output_bucket is None: + output_bucket = [None] * len(small_bucket) local_rank = dist.get_rank(group=process_group) - for buf, synced, bucket_rank in zip(small_bucket, self.unflatten(allreduced, small_bucket), bucket_ranks): + for buf, synced, bucket_rank, output in zip(small_bucket, synced_tensors, bucket_ranks, output_bucket): copy_to_local_rank = local_rank in bucket_rank if isinstance(bucket_rank, frozenset) else local_rank == bucket_rank if copy_to_local_rank: + self._add_to_grad_partition_buffer(output, synced) buf.copy_(synced) if self.overlap_comm and not get_accelerator().resolves_data_dependency(): buf.record_stream(self.reduction_stream) @@ -1427,13 +1538,19 @@ def allreduce_and_scatter(self, process_group=None): small_bucket = [] small_bucket_ranks = [] + small_bucket_outputs = [] numel = 0 allreduce_sizes = [] for i, bucket_elem in enumerate(bucket): - rank, tensor = bucket_elem + if len(bucket_elem) == 2: + rank, tensor = bucket_elem + output = None + else: + rank, tensor, output = bucket_elem small_bucket.append(tensor) small_bucket_ranks.append(rank) + small_bucket_outputs.append(output) numel = numel + tensor.numel() if numel > numel_per_bucket: self.allreduce_and_copy_with_multiple_ranks(small_bucket, @@ -1441,9 +1558,11 @@ def allreduce_and_scatter(self, log=None, divide=divide, process_group=process_group, - bucket_ranks=small_bucket_ranks) + bucket_ranks=small_bucket_ranks, + output_bucket=small_bucket_outputs) small_bucket = [] small_bucket_ranks = [] + small_bucket_outputs = [] numel = 0 if len(small_bucket) > 0: @@ -1452,7 +1571,8 @@ def allreduce_and_scatter(self, log=None, divide=divide, process_group=process_group, - bucket_ranks=small_bucket_ranks) + bucket_ranks=small_bucket_ranks, + output_bucket=small_bucket_outputs) def average_tensor(self, tensor: torch.Tensor, communication_data_type: torch.dtype): if self._offload_gradient_safety_enabled: @@ -1481,7 +1601,8 @@ def average_tensor(self, tensor: torch.Tensor, communication_data_type: torch.dt with get_accelerator().stream(stream): self._record_gradient_stream(tensor, stream) if not self.reduce_scatter: - self.gradient_reduction_w_predivide(tensor, communication_data_type) + reduced_tensor = self.gradient_reduction_w_predivide(tensor, communication_data_type) + self._add_full_reduction_to_compact_grad_partition(reduced_tensor, communication_data_type) return # Accumulate destination ranks and bucket offsets for each gradient slice. @@ -1535,14 +1656,15 @@ def average_tensor(self, tensor: torch.Tensor, communication_data_type: torch.dt numel = partition_ids_w_offsets[idx + 1][1] - offset copy_ranks = muon_copy_ranks or frozenset((partition_id, )) + grad_partition_output = self._get_compact_grad_partition_slice(i, partition_id, param_id, numel) # Merge bucket ranges if they share the same reduction and consumers. - if (partition_id == prev_id and process_group == prev_process_group - and copy_ranks == prev_copy_ranks): - prev_pid, prev_size, prev_numel, _ = rank_and_offsets[-1] - rank_and_offsets[-1] = (prev_pid, prev_size, prev_numel + numel, copy_ranks) + if (grad_partition_output is None and partition_id == prev_id + and process_group == prev_process_group and copy_ranks == prev_copy_ranks): + prev_pid, prev_size, prev_numel, _, _ = rank_and_offsets[-1] + rank_and_offsets[-1] = (prev_pid, prev_size, prev_numel + numel, copy_ranks, None) else: - rank_and_offsets.append((partition_id, curr_size, numel, copy_ranks)) + rank_and_offsets.append((partition_id, curr_size, numel, copy_ranks, grad_partition_output)) real_dp_process_group.append(process_group) curr_size += numel prev_id, prev_process_group, prev_copy_ranks = partition_id, process_group, copy_ranks @@ -1551,7 +1673,7 @@ def average_tensor(self, tensor: torch.Tensor, communication_data_type: torch.dt tensor.div_(dist.get_world_size(group=self.dp_process_group) / float(self.sequence_parallel_size)) buckets = {} - for i, (dst, bucket_offset, numel, copy_ranks) in enumerate(rank_and_offsets): + for i, (dst, bucket_offset, numel, copy_ranks, grad_partition_output) in enumerate(rank_and_offsets): grad_slice = tensor.narrow(0, int(bucket_offset), int(numel)) process_group = real_dp_process_group[i] # Split Muon matrices require all-reduce even when multi-rank bucket all-reduce is disabled. @@ -1562,9 +1684,9 @@ def average_tensor(self, tensor: torch.Tensor, communication_data_type: torch.dt if bucket_key not in buckets: buckets[bucket_key] = [] if bucket_key[0] == "allreduce": - buckets[bucket_key].append((copy_ranks, grad_slice)) + buckets[bucket_key].append((copy_ranks, grad_slice, grad_partition_output)) else: - buckets[bucket_key].append(grad_slice) + buckets[bucket_key].append((grad_slice, grad_partition_output)) for bucket_key in buckets: if bucket_key[0] == "allreduce": @@ -1575,12 +1697,15 @@ def average_tensor(self, tensor: torch.Tensor, communication_data_type: torch.dt process_group=bucket_key[1]) else: _, dst, process_group = bucket_key - self.allreduce_no_retain(buckets[bucket_key], + grad_slices = [item[0] for item in buckets[bucket_key]] + grad_partition_outputs = [item[1] for item in buckets[bucket_key]] + self.allreduce_no_retain(grad_slices, communication_data_type, numel_per_bucket=self.reduce_bucket_size, rank=dst, divide=False, - process_group=process_group) + process_group=process_group, + output_bucket=grad_partition_outputs) ############################################################################## ############################# CPU Offload Methods############################# @@ -1983,7 +2108,21 @@ def _apply_muon_updates_cpu_offload(self): del full_grad, full_momentum, grad, param_momentum, update ############################################################################################ - def copy_grads_in_partition(self, param): + def _get_param_group_idx(self, param): + param_idx_in_group = getattr(param, "param_idx_in_group", None) + if param_idx_in_group is not None: + for group_idx, group in enumerate(self.bit16_groups): + if param_idx_in_group < len(group) and group[param_idx_in_group] is param: + return group_idx + + for group_idx, group in enumerate(self.bit16_groups): + for group_param in group: + if group_param is param: + return group_idx + + raise RuntimeError(f"Unable to find ZeRO parameter group for {debug_param2name(param)}") + + def copy_grads_in_partition(self, param, group_idx=None): if self.cpu_offload: # Accumulate when there were prior backwards in this step (restore from # CPU buffer) or more will follow (save to CPU buffer). Skipping only @@ -2002,6 +2141,12 @@ def copy_grads_in_partition(self, param): self.async_inplace_copy_grad_to_fp32_buffer_from_gpu(param) return + if group_idx is None: + group_idx = self._get_param_group_idx(param) + if self._use_compact_grad_partition_buffer(group_idx): + if self.use_grad_accum_attribute: + self.clear_grad_attribute(param) + return #print(f"ID {self.get_param_id(param)} grad norm {param.grad.norm()}") if self.grads_in_partition is None: self.grads_in_partition_offset = 0 @@ -2084,10 +2229,10 @@ def reduce_ipg_grads(self, comm_dtype=None): else: self.clear_grad_attribute(param) elif self.contiguous_gradients: - self.copy_grads_in_partition(param) + self.copy_grads_in_partition(param, group_idx) else: # zero stage 1 - partition only optimizer state if self.contiguous_gradients and self.is_param_in_current_partition[param_id]: - self.copy_grads_in_partition(param) + self.copy_grads_in_partition(param, group_idx) # Empty and oversized reductions do not consume the contiguous buffer. if (self._offload_gradient_safety_enabled and 0 < bucket.elements <= self.reduce_bucket_size and not get_accelerator().resolves_data_dependency()): @@ -2206,7 +2351,7 @@ def allreduce_bucket(self, if rank is None or rank == dist.get_rank(group=process_group): tensor.copy_(tensor_to_allreduce) - return tensor + return tensor_to_allreduce def _clear_previous_reduced_grads(self): for dtype in self.previous_reduced_grads: @@ -2221,7 +2366,8 @@ def allreduce_and_copy(self, rank=None, log=None, divide=True, - process_group=None): + process_group=None, + output_bucket=None): process_group = self.dp_process_group if process_group is None else process_group if self.overlap_comm: if not get_accelerator().resolves_data_dependency(): @@ -2243,8 +2389,12 @@ def allreduce_and_copy(self, ) if self.overlap_comm and not get_accelerator().resolves_data_dependency(): allreduced.record_stream(stream) - if rank is None or rank == dist.get_rank(group=self.dp_process_group): - for buf, synced in zip(small_bucket, self.unflatten(allreduced, small_bucket)): + if rank is None or rank == dist.get_rank(group=process_group): + synced_tensors = self.unflatten(allreduced, small_bucket) + if output_bucket is None: + output_bucket = [None] * len(small_bucket) + for buf, synced, output in zip(small_bucket, synced_tensors, output_bucket): + self._add_to_grad_partition_buffer(output, synced) buf.copy_(synced) if self.overlap_comm and not get_accelerator().resolves_data_dependency(): buf.record_stream(stream) @@ -2258,11 +2408,16 @@ def allreduce_no_retain( log=None, divide=True, process_group=None, + output_bucket=None, ): small_bucket = [] + small_output_bucket = [] numel = 0 - for tensor in bucket: + if output_bucket is None: + output_bucket = [None] * len(bucket) + for tensor, output in zip(bucket, output_bucket): small_bucket.append(tensor) + small_output_bucket.append(output) numel = numel + tensor.numel() if numel > numel_per_bucket: self.allreduce_and_copy(small_bucket, @@ -2270,8 +2425,10 @@ def allreduce_no_retain( rank=rank, log=None, divide=divide, - process_group=process_group) + process_group=process_group, + output_bucket=small_output_bucket) small_bucket = [] + small_output_bucket = [] numel = 0 if len(small_bucket) > 0: @@ -2280,7 +2437,8 @@ def allreduce_no_retain( rank=rank, log=log, divide=divide, - process_group=process_group) + process_group=process_group, + output_bucket=small_output_bucket) # allows using reduction of gradients instead of using all_reduce @@ -3007,9 +3165,10 @@ def setup_buckets(self): for _, bucket in self.ipg_buckets.items(): bucket.buffer.clear() - # Buffer's dtype is the same as the dtype of optimizer, not dtype for autocast + # ZeRO-1 stages the fp32 accumulator; ZeRO-2 keeps low-precision staging + # and stores the fp32 communication result in the owned partition. buf_0 = torch.empty(int(self.reduce_bucket_size), - dtype=self.dtype, + dtype=self._gradient_reduction_buffer_dtype(), device=get_accelerator().current_device_name()) bucket.buffer.append(buf_0) bucket.index = 0 @@ -3018,7 +3177,7 @@ def setup_buckets(self): if self.overlap_comm: for _, bucket in self.ipg_buckets.items(): buf_1 = torch.empty(int(self.reduce_bucket_size), - dtype=self.dtype, + dtype=self._gradient_reduction_buffer_dtype(), device=get_accelerator().current_device_name()) bucket.buffer.append(buf_1) diff --git a/tests/unit/runtime/zero/test_zero_tensor_fragment.py b/tests/unit/runtime/zero/test_zero_tensor_fragment.py index 1214a8247acb..a0262baca901 100644 --- a/tests/unit/runtime/zero/test_zero_tensor_fragment.py +++ b/tests/unit/runtime/zero/test_zero_tensor_fragment.py @@ -99,6 +99,16 @@ def forward(self, x, y): return self.cel(x, y) +class ScaleModel(torch.nn.Module): + + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.ones(4)) + + def forward(self, x): + return (self.weight * x).sum() + + def run_fragmented_model(model, config_dict, hidden_dim, dtype, validate_after_bwd, validate_after_step): model, _, _, _ = deepspeed.initialize(model=model, model_parameters=model.parameters(), config=config_dict) data_loader = random_dataloader(model=model, @@ -229,6 +239,62 @@ def test_bf16_optimizer_fragments(self, frozen_weights): run_fragmented_model(model, config_dict, hidden_dim, torch.bfloat16, validate_after_bwd, validate_after_step) +class TestZeroBf16Fp32GradAccum(DistributedTest): + world_size = 2 + + @pytest.mark.parametrize('zero_stage', [2, 3]) + def test_safe_get_full_grad_preserves_fp32_accumulated_average(self, zero_stage): + # Catches storing or accumulating fp32 communication results in a bf16 partition. + if not get_accelerator().is_bf16_supported(): + pytest.skip("requires bf16") + + config_dict = { + "train_micro_batch_size_per_gpu": 1, + "gradient_accumulation_steps": 4, + "steps_per_print": 1, + "communication_data_type": "fp32", + "optimizer": { + "type": "Adam", + "params": { + "lr": 1e-6, + "torch_adam": True + } + }, + "bf16": { + "enabled": True + }, + "data_types": { + "grad_accum_dtype": "fp32" + }, + "zero_optimization": { + "stage": zero_stage, + } + } + + model = ScaleModel() + engine, _, _, _ = deepspeed.initialize(model=model, model_parameters=model.parameters(), config=config_dict) + + rank = dist.get_rank() + local_grad = 1.0 if rank == 0 else 1.0078125 + x = torch.full((4, ), local_grad, dtype=torch.bfloat16, device=engine.device) + for micro_step in range(engine.gradient_accumulation_steps()): + engine.backward(engine(x)) + if not engine.is_gradient_accumulation_boundary(): + engine.step() + + if zero_stage == 2: + stored_grad = next(grad for grad in engine.optimizer.averaged_gradients[0] + if not getattr(grad, "_zero_padding", False)) + assert stored_grad.dtype == torch.float32 + + grad = safe_get_full_grad(engine.module.weight) + expected = torch.full((4, ), 1.00390625, dtype=torch.float32, device=engine.device) + assert grad.dtype == torch.float32 + torch.testing.assert_close(grad, expected, rtol=0, atol=0) + + engine.destroy() + + def create_random_values(model, key_list, group, grad_dtype): param_values = {} for n, lp in model.named_parameters(): diff --git a/tests/unit/v1/zero/test_zero2_gradient_safety.py b/tests/unit/v1/zero/test_zero2_gradient_safety.py index c832949cabc4..a6a83101e0da 100644 --- a/tests/unit/v1/zero/test_zero2_gradient_safety.py +++ b/tests/unit/v1/zero/test_zero2_gradient_safety.py @@ -87,6 +87,34 @@ def backward_values(opt, param, values, boundaries): opt._wait_for_offload_copies() +def test_allreduce_and_scatter_accepts_legacy_rank_tensor_buckets(): + opt = zero.DeepSpeedZeroOptimizer.__new__(zero.DeepSpeedZeroOptimizer) + captured = {} + + def capture(small_bucket, + communication_data_type, + log=None, + divide=True, + process_group=None, + bucket_ranks=None, + output_bucket=None): + captured["small_bucket"] = small_bucket + captured["bucket_ranks"] = bucket_ranks + captured["output_bucket"] = output_bucket + + opt.allreduce_and_copy_with_multiple_ranks = capture + tensor = torch.ones(2) + output = torch.empty_like(tensor) + + opt.allreduce_and_scatter([(0, tensor), (1, tensor, output)], torch.float32) + + assert captured["small_bucket"][0] is tensor + assert captured["small_bucket"][1] is tensor + assert captured["bucket_ranks"] == [0, 1] + assert captured["output_bucket"][0] is None + assert captured["output_bucket"][1] is output + + @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32]) @pytest.mark.parametrize("cpu_offload", [False, True]) @pytest.mark.parametrize("bucket_size", [4, 8, 16]) @@ -224,7 +252,7 @@ def test_only_buffer_consumers_replace_reuse_event(monkeypatch, bucket_size, ele opt.reduce_bucket_size = bucket_size opt.is_param_in_current_partition = {0: True} opt.average_tensor = lambda *args: None - opt.copy_grads_in_partition = lambda param: None + opt.copy_grads_in_partition = lambda param, group_idx=None: None bucket = opt.ipg_buckets[torch.float32] bucket.elements = elements previous, consumer = FakeEvent(), FakeStream()