diff --git a/deepspeed/runtime/base_optimizer.py b/deepspeed/runtime/base_optimizer.py index f04c5d35f5bd..fe5e508f7d17 100644 --- a/deepspeed/runtime/base_optimizer.py +++ b/deepspeed/runtime/base_optimizer.py @@ -5,7 +5,9 @@ import os import torch +from functools import wraps from typing import Any +from weakref import ref from deepspeed.utils import logger from deepspeed.utils.tensor_fragment import map_to_flat_opt_states @@ -18,6 +20,21 @@ class DeepSpeedOptimizer(object): pass +class _IPGBucketParameterList(list): + """Record when a non-empty IPG bucket is drained during an active backward.""" + + def __init__(self, values, optimizer): + super().__init__(values) + self._optimizer_ref = ref(optimizer) + + def clear(self): + optimizer = self._optimizer_ref() + if (self and optimizer is not None and optimizer._backward_active_depth > 0 + and not optimizer._clearing_aborted_backward): + optimizer._backward_reduction_observed = True + super().clear() + + def _get_universal_checkpoint_ep_info() -> tuple[int, int]: # Universal checkpoints use EP slicing only when an expert group exists. try: @@ -146,13 +163,21 @@ def exit_backward(self): self.backward_active_depth -= 1 def reset_for_new_step(self): - """Reset state at the start of each forward/backward step.""" + """Reset state at the start of each forward/backward step. + + Returns whether the previous backward was still active. This can happen + when autograd exits through an exception before the backward epilogue. + """ + incomplete_backward = self.backward_active_depth > 0 + self.remaining_grad_acc_hooks = 0 + self.backward_active_depth = 0 self.backward_seen_this_step = False self.hooks_fired_this_backward = 0 self.max_expected_hooks_seen = 0 self.epilogue_ran_this_backward = False self.post_backward_callback_queued = False self.post_backward_callback_graph_task_id = None + return incomplete_backward def should_refresh_expected_hook_count(self): """Return True when count_used_parameters_in_backward() should be re-evaluated. @@ -250,8 +275,46 @@ def update_hook_state_and_maybe_run_epilogue(self, current_expected_count): class ZeROOptimizer(DeepSpeedOptimizer): """Base class for ZeRO optimizer implementations (stages 1, 2, and 3).""" + def __init_subclass__(cls, **kwargs): + """Require every ZeRO step implementation to reject an incomplete backward.""" + super().__init_subclass__(**kwargs) + step = cls.__dict__.get("step") + if step is None or getattr(step, "_deepspeed_aborted_backward_guard", False): + return + + @wraps(step) + def guarded_step(self, *args, **kwargs): + self._ensure_backward_complete_before_step() + result = step(self, *args, **kwargs) + self._backward_completed_since_step = False + return result + + guarded_step._deepspeed_aborted_backward_guard = True + cls.step = guarded_step + + def __setattr__(self, name, value): + """Preserve the aborted-backward guard across runtime step replacement.""" + if (name == "step" and callable(value) and "_backward_hook_state" in self.__dict__ + and not getattr(value, "_deepspeed_aborted_backward_guard", False)): + step = value + + @wraps(step) + def guarded_step(*args, **kwargs): + self._ensure_backward_complete_before_step() + result = step(*args, **kwargs) + self._backward_completed_since_step = False + return result + + guarded_step._deepspeed_aborted_backward_guard = True + value = guarded_step + object.__setattr__(self, name, value) + def __init__(self): self._backward_hook_state = BackwardHookStateManager() + self._aborted_backward_requires_restart = False + self._backward_completed_since_step = False + self._backward_reduction_observed = False + self._clearing_aborted_backward = False # Mirrored copy of the engine GAS boundary for managed reduce/offload paths. # Engine owns the source of truth (micro-step / step() / set_*); ZeRO reads this # during backward. Prefer get/set methods over touching the private field. @@ -460,10 +523,116 @@ def enter_backward(self): def exit_backward(self): """Exit backward context. Call at the end of backward pass.""" self._backward_hook_state.exit_backward() + if self._backward_active_depth == 0: + self._backward_completed_since_step = True + + @staticmethod + def _aborted_backward_error(): + return RuntimeError("An aborted backward already contributed gradients; discard this engine and restart the " + "accumulation window") + + def _ensure_backward_complete_before_step(self): + if self._aborted_backward_requires_restart or self._backward_active_depth > 0: + self._aborted_backward_requires_restart = True + raise self._aborted_backward_error() + + def _track_ipg_bucket_reductions(self, buckets): + for bucket in buckets.values(): + if not isinstance(bucket.params, _IPGBucketParameterList): + bucket.params = _IPGBucketParameterList(bucket.params, self) + + def _clear_unreduced_parameter_gradients(self): + groups = getattr(self, "bit16_groups", None) + if groups is None: + groups = getattr(self, "fp16_groups", None) + if groups is None: + groups = getattr(self, "bf16_groups", ()) + for group in groups: + for param in group: + param.grad = None + if hasattr(param, "grad_accum"): + param.grad_accum = None def clear_backward_seen_flag(self): - """Clear the backward seen flag and reset hook counters at the start of each step.""" - self._backward_hook_state.reset_for_new_step() + """Reset hook state and discard partial reduction state from an aborted backward.""" + buckets = getattr(self, "ipg_buckets", {}) + self._track_ipg_bucket_reductions(buckets) + if self._aborted_backward_requires_restart: + raise self._aborted_backward_error() + + hooks_fired = self._backward_hook_state.hooks_fired_this_backward + params_already_reduced = getattr(self, "params_already_reduced", None) + if isinstance(params_already_reduced, dict): + any_param_reduced = any(params_already_reduced.values()) + elif params_already_reduced is not None: + any_param_reduced = any(params_already_reduced) + else: + any_param_reduced = False + bf16_grad_contributed = any(any(group) for group in getattr(self, "fp32_groups_has_gradients", ())) + zero1_grad_contributed = (hasattr(self, "partition_gradients") and not self.partition_gradients + and self._backward_completed_since_step) + coalesced_grad_contributed = (getattr(self, "_coalesce_grad_reduction", False) + and self._backward_completed_since_step) + stage12_buckets_were_ready = getattr(self, "ready_for_gradients", False) + + incomplete_backward = self._backward_hook_state.reset_for_new_step() + if not incomplete_backward: + self._backward_reduction_observed = False + return + + # Bucket parameter lists are instrumented before every backward. Clearing + # a non-empty list while backward is active means a reduction crossed into + # persistent accumulation state. ZeRO-1's separate grad_accum path is also + # irreversible once any hook has contributed to the shared accumulator. + partial_reduction = (self._backward_reduction_observed or any_param_reduced or zero1_grad_contributed + or coalesced_grad_contributed + or (getattr(self, "use_grad_accum_attribute", False) and hooks_fired > 0) + or bf16_grad_contributed) + + # Autograd exceptions bypass the backward epilogue, so a partially filled + # IPG bucket and its reduced flags must not flow into the next root forward. + self._clearing_aborted_backward = True + try: + for bucket in buckets.values(): + if hasattr(bucket, "clear_params"): + bucket.clear_params() + else: + bucket.clear() + finally: + self._clearing_aborted_backward = False + + self._clear_unreduced_parameter_gradients() + + extra_large_params = getattr(self, "extra_large_param_to_reduce", None) + if extra_large_params is not None: + extra_large_params.clear() + + if isinstance(params_already_reduced, dict): + for param_id in params_already_reduced: + params_already_reduced[param_id] = False + elif params_already_reduced is not None: + for param_id in range(len(params_already_reduced)): + params_already_reduced[param_id] = False + + if hasattr(self, "reset_partition_gradient_structures"): + self.reset_partition_gradient_structures() + if hasattr(self, "ready_for_gradients"): + self.ready_for_gradients = False + if hasattr(self, "grads_in_partition_offset"): + self.grads_in_partition_offset = 0 + + if partial_reduction: + self._aborted_backward_requires_restart = True + raise self._aborted_backward_error() + + # Stage 1/2 advances micro_step_id when its first gradient hook sets up + # the IPG buckets. A safe abort before any bucket drain must undo that + # advance so CPU-offload retry does not import the previous step's + # accumulated_grads_in_cpu as a current-window contribution. + if stage12_buckets_were_ready and hasattr(self, "micro_step_id"): + self.micro_step_id -= 1 + + self._backward_reduction_observed = False def should_refresh_expected_hook_count(self): """Return True when count_used_parameters_in_backward() should be re-evaluated.""" diff --git a/deepspeed/runtime/engine.py b/deepspeed/runtime/engine.py index 5db39d237967..15c6ee54f496 100755 --- a/deepspeed/runtime/engine.py +++ b/deepspeed/runtime/engine.py @@ -18,6 +18,9 @@ from torch.optim.lr_scheduler import _LRScheduler from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors from contextlib import contextmanager +from contextvars import ContextVar +from threading import Lock +from weakref import ref from typing import Callable, Dict, Union, Iterable, Container, List @@ -246,6 +249,85 @@ def _checkpoint_parallel_metadata(mpu): } +class _EngineBackwardGraphState: + + def __init__(self): + self.active = False + self.graph_task_ids = [] + + +_ENGINE_BACKWARD_GRAPH_CONTEXT = ContextVar("deepspeed_engine_backward_graph_context", default=()) + + +class _EngineBackwardGraphTracker: + + def __init__(self): + self._lock = Lock() + self._graph_task_refcounts = {} + self._active_states = set() + + @staticmethod + def supported(): + return hasattr(torch._C, "_current_graph_task_id") + + def new_state(self): + return _EngineBackwardGraphState() + + @staticmethod + def _clear_current_context(state): + context = _ENGINE_BACKWARD_GRAPH_CONTEXT.get() + _ENGINE_BACKWARD_GRAPH_CONTEXT.set( + tuple(state_ref for state_ref in context if state_ref() is not None and state_ref() is not state)) + + def register_current_graph(self, grad, state): + graph_task_id = torch._C._current_graph_task_id() if self.supported() else -1 + with self._lock: + if not state.active: + state.active = True + self._active_states.add(state) + if graph_task_id != -1 and graph_task_id not in state.graph_task_ids: + self._graph_task_refcounts[graph_task_id] = self._graph_task_refcounts.get(graph_task_id, 0) + 1 + state.graph_task_ids.append(graph_task_id) + + context = _ENGINE_BACKWARD_GRAPH_CONTEXT.get() + with self._lock: + context = tuple(state_ref for state_ref in context if state_ref() in self._active_states) + _ENGINE_BACKWARD_GRAPH_CONTEXT.set((*context, ref(state))) + torch.autograd.Variable._execution_engine.queue_callback(lambda: self._clear_current_context(state)) + return grad + + def unregister(self, state): + with self._lock: + if not state.active: + return + for graph_task_id in state.graph_task_ids: + refcount = self._graph_task_refcounts[graph_task_id] - 1 + if refcount: + self._graph_task_refcounts[graph_task_id] = refcount + else: + del self._graph_task_refcounts[graph_task_id] + state.active = False + self._active_states.remove(state) + + def is_current_graph_registered(self): + graph_task_id = torch._C._current_graph_task_id() if self.supported() else -1 + context = _ENGINE_BACKWARD_GRAPH_CONTEXT.get() + with self._lock: + if graph_task_id in self._graph_task_refcounts: + return True + active_context = tuple(state_ref for state_ref in context if state_ref() in self._active_states) + if active_context != context: + _ENGINE_BACKWARD_GRAPH_CONTEXT.set(active_context) + return bool(active_context) + + def active_registration_count(self): + with self._lock: + return len(self._active_states) + + +_ENGINE_BACKWARD_GRAPH_TRACKER = _EngineBackwardGraphTracker() + + class DeepSpeedEngine(Module): r"""DeepSpeed engine for training.""" @@ -487,11 +569,11 @@ def __init__(self, # Otherwise, we fallback to DeepSpeed style backward only. # See `count_used_parameters_in_backward` for more details. self._running_engine_backward = False + self._running_engine_backward_count = 0 + self._running_engine_backward_lock = Lock() # True only while step() runs; the unmanaged-mode accumulation boundary. self._running_engine_step = False self._support_torch_style_backward = False - # Flag to control whether gradients should be scaled by gradient accumulation steps - self._scale_wrt_gas = True if isinstance(self.optimizer, ZeROOptimizer) and check_internal_apis_for_count_used_parameters(): self._support_torch_style_backward = True # These hooks are used for non-scalar backward support, such as `out.backward(out_grad)`, @@ -2951,8 +3033,10 @@ def _backward_epilogue(self): def _backward_prologue_per_tensor(self, grad): if is_functorch_transforming(): return grad - # Only scale gradients if scale_wrt_gas is True, consistent with backward() parameter - if grad is not None and self._scale_wrt_gas: + graph_loss_scaled = _ENGINE_BACKWARD_GRAPH_TRACKER.is_current_graph_registered() + if graph_loss_scaled: + return grad + if grad is not None: return grad / self.gradient_accumulation_steps() return grad @@ -3175,45 +3259,57 @@ def backward(self, loss, retain_graph=False, scale_wrt_gas=True): assert maybe_loss_for_backward( loss), "loss must be a scalar tensor. If you need to pass output gradients, backward() of output tensors" - self._running_engine_backward = True - # Store scale_wrt_gas so the hook can respect it - self._scale_wrt_gas = scale_wrt_gas - - # Unmanaged mode: count this backward so step() can advance global_samples by the actual micro-batch count. - if not self.managed_gradient_accumulation(): - self._unmanaged_backward_count += 1 - - # Set flag to prevent hooks from firing (we'll manually call prologue/epilogue) - backward_kwargs = {"retain_graph": retain_graph} - if self.eigenvalue_enabled(): - backward_kwargs["create_graph"] = True - backward_kwargs["retain_graph"] = True - - # Used only for return value - gas_scaled_loss = loss / self.gradient_accumulation_steps() if scale_wrt_gas else loss - - # TODO: handle these scaling with direct calls to loss.backward() - if isinstance(self.optimizer, ZeROOptimizer): - loss = self.optimizer.scale_if_loss(loss) - elif self.torch_autocast_z0_gradscaler: - loss = self.torch_autocast_z0_gradscaler.scale(loss) - - with compiled_autograd(self._is_compiled_autograd_enabled, self._compile_kwargs): - if self.zero_optimization() or not self.amp_enabled(): - loss.backward(**backward_kwargs) - elif self.amp_enabled(): - # AMP requires delaying unscale when inside gradient accumulation boundaries - # https://nvidia.github.io/apex/advanced.html#gradient-accumulation-across-iterations - delay_unscale = not self.is_gradient_accumulation_boundary() - with amp.scale_loss(loss, self.optimizer, delay_unscale=delay_unscale) as scaled_loss: - scaled_loss.backward(**backward_kwargs) + with self._running_engine_backward_lock: + self._running_engine_backward_count += 1 + self._running_engine_backward = True + engine_backward_graph_state = _ENGINE_BACKWARD_GRAPH_TRACKER.new_state() + engine_backward_graph_hook = None + try: + # Unmanaged mode: count this backward so step() can advance global_samples by the actual micro-batch count. + if not self.managed_gradient_accumulation(): + self._unmanaged_backward_count += 1 - # backward_epilogue is not called in a hook when self._support_torch_style_backward is False - self._backward_epilogue() + # Set flag to prevent hooks from firing (we'll manually call prologue/epilogue) + backward_kwargs = {"retain_graph": retain_graph} + if self.eigenvalue_enabled(): + backward_kwargs["create_graph"] = True + backward_kwargs["retain_graph"] = True - self._running_engine_backward = False + loss = loss / self.gradient_accumulation_steps() if scale_wrt_gas else loss + gas_scaled_loss = loss - return gas_scaled_loss + # TODO: handle these scaling with direct calls to loss.backward() + if isinstance(self.optimizer, ZeROOptimizer): + loss = self.optimizer.scale_if_loss(loss) + elif self.torch_autocast_z0_gradscaler: + loss = self.torch_autocast_z0_gradscaler.scale(loss) + + if loss.requires_grad: + engine_backward_graph_hook = loss.register_hook( + lambda grad: _ENGINE_BACKWARD_GRAPH_TRACKER.register_current_graph( + grad, engine_backward_graph_state)) + + with compiled_autograd(self._is_compiled_autograd_enabled, self._compile_kwargs): + if self.zero_optimization() or not self.amp_enabled(): + loss.backward(**backward_kwargs) + elif self.amp_enabled(): + # AMP requires delaying unscale when inside gradient accumulation boundaries + # https://nvidia.github.io/apex/advanced.html#gradient-accumulation-across-iterations + delay_unscale = not self.is_gradient_accumulation_boundary() + with amp.scale_loss(loss, self.optimizer, delay_unscale=delay_unscale) as scaled_loss: + scaled_loss.backward(**backward_kwargs) + + # backward_epilogue is not called in a hook when self._support_torch_style_backward is False + self._backward_epilogue() + + return gas_scaled_loss + finally: + if engine_backward_graph_hook is not None: + engine_backward_graph_hook.remove() + _ENGINE_BACKWARD_GRAPH_TRACKER.unregister(engine_backward_graph_state) + with self._running_engine_backward_lock: + self._running_engine_backward_count -= 1 + self._running_engine_backward = self._running_engine_backward_count > 0 def is_gradient_accumulation_boundary(self): """ @@ -3369,6 +3465,9 @@ def step(self, lr_kwargs=None): assert not self.inside_no_sync_ctxt, \ "It is illegal to call Engine.step() inside no_sync context manager" + if isinstance(self.optimizer, ZeROOptimizer): + self.optimizer._ensure_backward_complete_before_step() + see_memory_usage("Engine before step", force=self.memory_breakdown()) # Check early because self.global_steps is incremented at some point here. diff --git a/tests/unit/runtime/zero/test_zero_shared_loss_gradient.py b/tests/unit/runtime/zero/test_zero_shared_loss_gradient.py new file mode 100644 index 000000000000..c6e47b4bac4e --- /dev/null +++ b/tests/unit/runtime/zero/test_zero_shared_loss_gradient.py @@ -0,0 +1,227 @@ +# Copyright (c) DeepSpeed Team. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +import pytest +import torch +from torch.utils.checkpoint import checkpoint + +import deepspeed +import deepspeed.comm as dist +from deepspeed.accelerator import get_accelerator +from deepspeed.runtime.engine import _ENGINE_BACKWARD_GRAPH_TRACKER +from deepspeed.utils import safe_get_full_grad +from unit.common import DistributedTest +from unit.util import bf16_required_version_check + + +class SharedLinear(torch.nn.Module): + + def __init__(self): + super().__init__() + self.shared = torch.nn.Linear(4, 4, bias=False) + + def forward(self, inputs): + return self.shared(inputs) + + +class _RaiseInBackward(torch.autograd.Function): + + @staticmethod + def forward(ctx, inputs): + return inputs.clone() + + @staticmethod + def backward(ctx, grad): + raise RuntimeError("intentional backward failure") + + +class FailingSharedLinear(SharedLinear): + + def forward(self, inputs): + return _RaiseInBackward.apply(super().forward(inputs)) + + +def _make_model(device, model_class=SharedLinear): + model = model_class().to(device=device, dtype=torch.bfloat16) + with torch.no_grad(): + values = torch.arange(16, device=device, dtype=torch.float32).reshape(4, 4) / 16 + model.shared.weight.copy_(values.to(torch.bfloat16)) + return model + + +def _inputs(device, rank, micro_step): + values = torch.arange(8, device=device, dtype=torch.float32).reshape(2, 4) + return (values + rank * 3 + micro_step).to(torch.bfloat16) + + +def _make_engine(model, gradient_accumulation_steps): + optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3) + config = { + "train_micro_batch_size_per_gpu": 2, + "gradient_accumulation_steps": gradient_accumulation_steps, + "bf16": { + "enabled": True, + }, + "zero_allow_untested_optimizer": True, + "zero_optimization": { + "stage": 2, + "overlap_comm": True, + "contiguous_gradients": True, + "reduce_scatter": True, + }, + } + engine, *_ = deepspeed.initialize(model=model, optimizer=optimizer, config=config) + return engine + + +def _reference_grad(model, inputs): + model.zero_grad(set_to_none=True) + model(inputs).float().square().mean().backward() + grad = model.shared.weight.grad.detach().clone() + dist.all_reduce(grad) + grad.div_(dist.get_world_size()) + return grad.float() + + +def _advance_to_accumulation_boundary(engine, inputs, gradient_accumulation_steps): + for _ in range(gradient_accumulation_steps - 1): + engine.backward(engine(inputs).float().sum() * 0) + engine.step() + + +class TestZero2SharedLossGradient(DistributedTest): + world_size = 2 + + def test_engine_and_module_branches_match_manual_reference(self): + if not bf16_required_version_check(): + pytest.skip("BF16 ZeRO-2 test requires BF16 accelerator support.") + + gradient_accumulation_steps = 8 + device = get_accelerator().current_device_name() + rank = dist.get_rank() + world_size = dist.get_world_size() + + reference_model = _make_model(device) + reference_grad = torch.zeros_like(reference_model.shared.weight, dtype=torch.float32) + for micro_step in range(gradient_accumulation_steps): + reference_model.zero_grad(set_to_none=True) + inputs = _inputs(device, rank, micro_step) + secondary_inputs = inputs * 0.5 + output = reference_model(inputs) + reference_model(secondary_inputs) + loss = output.float().square().mean() / gradient_accumulation_steps + loss.backward() + + microbatch_grad = reference_model.shared.weight.grad.detach().clone() + dist.all_reduce(microbatch_grad) + microbatch_grad.div_(world_size) + reference_grad.add_(microbatch_grad.float()) + + engine = _make_engine(_make_model(device), gradient_accumulation_steps) + try: + for micro_step in range(gradient_accumulation_steps): + inputs = _inputs(device, rank, micro_step) + secondary_inputs = inputs * 0.5 + output = engine(inputs) + engine.module(secondary_inputs) + engine.backward(output.float().square().mean()) + if micro_step + 1 < gradient_accumulation_steps: + engine.step() + + actual_grad = safe_get_full_grad(engine.module.shared.weight) + assert actual_grad is not None + torch.testing.assert_close(actual_grad.float(), reference_grad, rtol=5e-3, atol=2.0) + finally: + engine.destroy() + + @pytest.mark.parametrize("graph_task_ids_available", [True, False], ids=["graph-task-id", "fallback"]) + def test_two_engine_managed_backward_scales_each_branch_once(self, monkeypatch, graph_task_ids_available): + if not bf16_required_version_check(): + pytest.skip("BF16 ZeRO-2 test requires BF16 accelerator support.") + + gradient_accumulation_steps = 8 + device = get_accelerator().current_device_name() + rank = dist.get_rank() + inputs = [_inputs(device, rank, micro_step) for micro_step in range(2)] + reference_grads = [_reference_grad(_make_model(device), value) for value in inputs] + engines = [_make_engine(_make_model(device), gradient_accumulation_steps) for _ in range(2)] + try: + monkeypatch.setattr(type(_ENGINE_BACKWARD_GRAPH_TRACKER), "supported", + staticmethod(lambda: graph_task_ids_available)) + for engine, value in zip(engines, inputs): + _advance_to_accumulation_boundary(engine, value, gradient_accumulation_steps) + losses = [engine(value).float().square().mean() for engine, value in zip(engines, inputs)] + engines[0].backward(sum(losses)) + + for index, (engine, reference_grad) in enumerate(zip(engines, reference_grads)): + actual_grad = safe_get_full_grad(engine.module.shared.weight) + assert actual_grad is not None + torch.testing.assert_close(actual_grad.float(), + reference_grad / gradient_accumulation_steps, + rtol=5e-3, + atol=2.0, + msg=f"engine {index} gradient was not scaled exactly once") + finally: + for engine in engines: + engine.destroy() + + def test_reentrant_checkpointed_engine_branches_are_scaled_once(self): + if not bf16_required_version_check(): + pytest.skip("BF16 ZeRO-2 test requires BF16 accelerator support.") + + gradient_accumulation_steps = 8 + device = get_accelerator().current_device_name() + rank = dist.get_rank() + inputs = [_inputs(device, rank, micro_step) for micro_step in range(2)] + reference_grads = [_reference_grad(_make_model(device), value) for value in inputs] + engines = [_make_engine(_make_model(device), gradient_accumulation_steps) for _ in range(2)] + try: + for engine, value in zip(engines, inputs): + _advance_to_accumulation_boundary(engine, value, gradient_accumulation_steps) + losses = [ + checkpoint(engine, value.detach().requires_grad_(True), use_reentrant=True).float().square().mean() + for engine, value in zip(engines, inputs) + ] + engines[0].backward(sum(losses)) + + for index, (engine, reference_grad) in enumerate(zip(engines, reference_grads)): + actual_grad = safe_get_full_grad(engine.module.shared.weight) + assert actual_grad is not None + torch.testing.assert_close(actual_grad.float(), + reference_grad / gradient_accumulation_steps, + rtol=5e-3, + atol=2.0, + msg=f"reentrant engine {index} gradient was not scaled exactly once") + finally: + for engine in engines: + engine.destroy() + + def test_managed_backward_context_restored_after_exception(self): + if not bf16_required_version_check(): + pytest.skip("BF16 ZeRO-2 test requires BF16 accelerator support.") + + gradient_accumulation_steps = 8 + device = get_accelerator().current_device_name() + rank = dist.get_rank() + inputs = _inputs(device, rank, 0) + reference_grad = _reference_grad(_make_model(device), inputs) + failing_engine = _make_engine(_make_model(device, FailingSharedLinear), gradient_accumulation_steps) + direct_backward_engine = _make_engine(_make_model(device), gradient_accumulation_steps) + try: + _advance_to_accumulation_boundary(direct_backward_engine, inputs, gradient_accumulation_steps) + failing_loss = failing_engine(inputs).float().square().mean() + with pytest.raises(RuntimeError, match="intentional backward failure"): + failing_engine.backward(failing_loss) + assert failing_engine._running_engine_backward is False + assert _ENGINE_BACKWARD_GRAPH_TRACKER.active_registration_count() == 0 + + direct_backward_engine(inputs).float().square().mean().backward() + actual_grad = safe_get_full_grad(direct_backward_engine.module.shared.weight) + assert actual_grad is not None + torch.testing.assert_close(actual_grad.float(), + reference_grad / gradient_accumulation_steps, + rtol=5e-3, + atol=2.0) + finally: + failing_engine.destroy() + direct_backward_engine.destroy() diff --git a/tests/unit/v1/zero/test_zero_activation_checkpoint_lifecycle.py b/tests/unit/v1/zero/test_zero_activation_checkpoint_lifecycle.py index 29d07bafa8fe..237ff4340718 100644 --- a/tests/unit/v1/zero/test_zero_activation_checkpoint_lifecycle.py +++ b/tests/unit/v1/zero/test_zero_activation_checkpoint_lifecycle.py @@ -119,6 +119,32 @@ def _assert_checkpoint_state_clean(engine, *, require_partitioned=True): f"status={parameter.ds_status}") +def _assert_backward_state_clean(engine): + """Assert that the next root forward discarded an incomplete backward.""" + optimizer = engine.optimizer + assert optimizer._backward_active_depth == 0 + assert optimizer._remaining_grad_acc_hooks == 0 + for bucket in optimizer.ipg_buckets.values(): + assert not bucket.params + assert bucket.elements == 0 + reduced_flags = optimizer.params_already_reduced + if isinstance(reduced_flags, dict): + reduced_flags = reduced_flags.values() + assert not any(reduced_flags) + assert not getattr(optimizer, "extra_large_param_to_reduce", {}) + groups = getattr(optimizer, "bit16_groups", getattr(optimizer, "fp16_groups", ())) + for group in groups: + for parameter in group: + assert parameter.grad is None + assert getattr(parameter, "grad_accum", None) is None + + +def _snapshot_trainable_parameters(engine): + parameters = [parameter for parameter in engine.module.parameters() if parameter.requires_grad] + with deepspeed.zero.GatheredParameters(parameters): + return [parameter.detach().float().cpu().clone() for parameter in parameters] + + class _RecursiveFrozenBlock(torch.nn.Module): """Recursively invoke one module instance so its ZeRO ds_id overlaps with itself.""" @@ -442,6 +468,7 @@ def test_reset_step_drains_incomplete_backward_state(self): def _observe_after_deepspeed_reset(module, unused_inputs): if observe_reset["enabled"]: _assert_checkpoint_state_clean(engine) + _assert_backward_state_clean(engine) observe_reset["enabled"] = False handle = engine.module.register_forward_pre_hook(_observe_after_deepspeed_reset) @@ -454,6 +481,328 @@ def _observe_after_deepspeed_reset(module, unused_inputs): _assert_checkpoint_state_clean(engine) engine.destroy() + @pytest.mark.parametrize("zero_stage", [2, 3], ids=["zero2", "zero3"]) + def test_partial_reduction_before_aborted_backward_requires_restart(self, zero_stage): + """A retry must fail closed once an aborted backward has reduced a partial bucket.""" + device, _, _ = initialize_distributed() + model = _IncompleteBackwardModel(hidden_dim=8) + model.backward_control["raise"] = False + config = get_config_dict(zero_stage, gradient_accumulation_steps=2, force_fp32=True) + config["zero_optimization"]["reduce_bucket_size"] = 8 + if zero_stage == 3: + config["zero_optimization"]["stage3_prefetch_bucket_size"] = 0 + config["zero_optimization"]["stage3_max_reuse_distance"] = 0 + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + + first = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + engine.backward(engine(first).sum()) + engine.step() + + model.backward_control["raise"] = True + failing = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + engine.backward(engine(failing).sum()) + + retry = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with pytest.raises(RuntimeError, match="restart the accumulation window"): + engine(retry) + with pytest.raises(RuntimeError, match="restart the accumulation window"): + engine(retry) + + engine.destroy() + + def test_zero3_leaf_partial_reduction_before_abort_requires_restart(self): + """A multi-parameter ZeRO-3 leaf hook must record every bucket drain.""" + from deepspeed.utils import set_z3_leaf_modules + + device, _, _ = initialize_distributed() + model = _IncompleteBackwardModel(hidden_dim=8) + model.backward_control["raise"] = False + set_z3_leaf_modules(model, [torch.nn.Linear]) + config = get_config_dict(3, gradient_accumulation_steps=2, force_fp32=True) + config["zero_optimization"]["reduce_bucket_size"] = 8 + config["zero_optimization"]["stage3_prefetch_bucket_size"] = 0 + config["zero_optimization"]["stage3_max_reuse_distance"] = 0 + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + + first = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + engine.backward(engine(first).sum()) + engine.step() + + model.backward_control["raise"] = True + failing = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + engine.backward(engine(failing).sum()) + + retry = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with pytest.raises(RuntimeError, match="restart the accumulation window"): + engine(retry) + engine.destroy() + + @pytest.mark.parametrize("zero_stage", [2, 3], ids=["zero2", "zero3"]) + def test_step_after_aborted_backward_requires_restart(self, zero_stage): + """An optimizer step cannot bypass aborted-backward cleanup at the next forward.""" + device, _, _ = initialize_distributed() + model = _IncompleteBackwardModel(hidden_dim=8) + config = get_config_dict(zero_stage, gradient_accumulation_steps=1, force_fp32=True) + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + + failing = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + engine.backward(engine(failing).sum()) + with pytest.raises(RuntimeError, match="restart the accumulation window"): + engine.step() + + retry = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with pytest.raises(RuntimeError, match="restart the accumulation window"): + engine(retry) + engine.destroy() + + @pytest.mark.parametrize("zero_stage", [2, 3], ids=["zero2", "zero3"]) + def test_non_boundary_step_after_aborted_backward_requires_restart(self, zero_stage): + """A non-boundary GAS step must reject an incomplete backward without advancing.""" + device, _, _ = initialize_distributed() + model = _IncompleteBackwardModel(hidden_dim=8) + config = get_config_dict(zero_stage, gradient_accumulation_steps=2, force_fp32=True) + config["zero_optimization"]["reduce_bucket_size"] = 1_000_000 + if zero_stage == 3: + config["zero_optimization"]["stage3_prefetch_bucket_size"] = 0 + config["zero_optimization"]["stage3_max_reuse_distance"] = 0 + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + + failing = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + engine.backward(engine(failing).sum()) + + initial_micro_steps = engine.micro_steps + with pytest.raises(RuntimeError, match="restart the accumulation window"): + engine.step() + assert engine.micro_steps == initial_micro_steps + + retry = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with pytest.raises(RuntimeError, match="restart the accumulation window"): + engine(retry) + engine.destroy() + + def test_runtime_step_replacement_keeps_aborted_backward_guard(self): + """Instance-level ZeRO step replacement must retain the incomplete-backward guard.""" + device, _, _ = initialize_distributed() + model = _IncompleteBackwardModel(hidden_dim=8) + config = get_config_dict(3, gradient_accumulation_steps=1, force_fp32=True) + config["zero_optimization"]["stage3_prefetch_bucket_size"] = 0 + config["zero_optimization"]["stage3_max_reuse_distance"] = 0 + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + replacement_calls = [] + engine.optimizer.step = lambda closure=None: replacement_calls.append(closure) + + failing = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + engine.backward(engine(failing).sum()) + with pytest.raises(RuntimeError, match="restart the accumulation window"): + engine.step() + + assert not replacement_calls + engine.destroy() + + def test_zero2_oversized_unreduced_gradient_is_cleared_for_retry(self): + """Discard an oversized ZeRO-2 bucket side-table entry after a safe abort.""" + device, _, _ = initialize_distributed() + torch.manual_seed(1234) + model = _IncompleteBackwardModel(hidden_dim=8) + model.head.bias.requires_grad_(False) + config = get_config_dict(2, gradient_accumulation_steps=1, force_fp32=True) + config["zero_optimization"]["reduce_bucket_size"] = 1 + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + + failing = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + engine.backward(engine(failing).sum()) + assert engine.optimizer.extra_large_param_to_reduce + + retry = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + engine.backward(engine(retry).sum()) + engine.step() + + _assert_backward_state_clean(engine) + engine.destroy() + + def test_bf16_immediate_update_aborted_backward_requires_restart(self): + """BF16 hooks that reached the persistent FP32 buffer make retry unsafe.""" + if not get_accelerator().is_bf16_supported(): + pytest.skip("bfloat16 is not supported on this accelerator") + + device, _, _ = initialize_distributed() + model = _IncompleteBackwardModel(hidden_dim=8) + config = get_config_dict(1, gradient_accumulation_steps=1) + config["bf16"] = {"enabled": True, "immediate_grad_update": True} + config["data_types"] = {"grad_accum_dtype": "fp32"} + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + + failing = torch.randn(2, 8, device=device, dtype=torch.bfloat16, requires_grad=True) + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + engine.backward(engine(failing).sum()) + assert any(any(group) for group in engine.optimizer.fp32_groups_has_gradients) + + retry = torch.randn(2, 8, device=device, dtype=torch.bfloat16, requires_grad=True) + with pytest.raises(RuntimeError, match="restart the accumulation window"): + engine(retry) + engine.destroy() + + def test_zero1_prior_accumulation_aborted_backward_requires_restart(self): + """Keep a valid earlier ZeRO-1 microbatch from being silently discarded.""" + device, _, _ = initialize_distributed() + engines = [] + + for _ in range(2): + torch.manual_seed(1234) + model = _IncompleteBackwardModel(hidden_dim=8) + model.backward_control["raise"] = False + config = get_config_dict(1, gradient_accumulation_steps=2, force_fp32=True) + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + engines.append(engine) + + control, aborted = engines + first = torch.randn(2, 8, device=device, dtype=torch.float32) + second = torch.randn(2, 8, device=device, dtype=torch.float32) + for engine in engines: + engine.backward(engine(first.clone().requires_grad_(True)).sum()) + engine.step() + + control.backward(control(second.clone().requires_grad_(True)).sum()) + control.step() + + aborted.module.backward_control["raise"] = True + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + aborted.backward(aborted(second.clone().requires_grad_(True)).sum()) + with pytest.raises(RuntimeError, match="restart the accumulation window"): + aborted(second.clone().requires_grad_(True)) + + control.destroy() + aborted.destroy() + + @pytest.mark.parametrize("zero_stage", [2, 3], ids=["zero2", "zero3"]) + def test_coalesced_prior_backward_abort_requires_restart(self, zero_stage): + """Keep an earlier coalesced backward from being silently discarded.""" + device, _, _ = initialize_distributed() + model = _IncompleteBackwardModel(hidden_dim=8) + model.backward_control["raise"] = False + config = get_config_dict(zero_stage, gradient_accumulation_steps=1, force_fp32=True) + if zero_stage == 3: + config["zero_optimization"]["stage3_prefetch_bucket_size"] = 0 + config["zero_optimization"]["stage3_max_reuse_distance"] = 0 + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + + first = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + failing = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + retry = torch.randn(2, 8, device=device, dtype=torch.float32, requires_grad=True) + with engine.coalesce_grad_reduction(): + engine.backward(engine(first).sum()) + model.backward_control["raise"] = True + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + engine.backward(engine(failing).sum()) + with pytest.raises(RuntimeError, match="restart the accumulation window"): + engine(retry) + + engine.destroy() + + @pytest.mark.parametrize("zero_stage", [2, 3], ids=["zero2", "zero3"]) + def test_pre_reduction_abort_retry_matches_clean_control(self, zero_stage): + """Discard unreduced failed gradients without losing earlier GAS contributions.""" + device, _, _ = initialize_distributed() + first = torch.randn(2, 8, device=device, dtype=torch.float32) + second = torch.randn(2, 8, device=device, dtype=torch.float32) + engines = [] + + for _ in range(2): + torch.manual_seed(1234) + model = _IncompleteBackwardModel(hidden_dim=8) + model.backward_control["raise"] = False + config = get_config_dict(zero_stage, gradient_accumulation_steps=2, force_fp32=True) + config["zero_optimization"]["reduce_bucket_size"] = 1_000_000 + if zero_stage == 3: + config["zero_optimization"]["stage3_prefetch_bucket_size"] = 0 + config["zero_optimization"]["stage3_max_reuse_distance"] = 0 + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + engines.append(engine) + + control, recovered = engines + for engine in engines: + engine.backward(engine(first.clone().requires_grad_(True)).sum()) + engine.step() + + control.backward(control(second.clone().requires_grad_(True)).sum()) + control.step() + + recovered.module.backward_control["raise"] = True + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + recovered.backward(recovered(second.clone().requires_grad_(True)).sum()) + recovered.backward(recovered(second.clone().requires_grad_(True)).sum()) + recovered.step() + + for control_parameter, recovered_parameter in zip(_snapshot_trainable_parameters(control), + _snapshot_trainable_parameters(recovered)): + torch.testing.assert_close(recovered_parameter, control_parameter, rtol=0, atol=0) + + control.destroy() + recovered.destroy() + + def test_zero2_cpu_offload_pre_reduction_abort_retry_matches_clean_control(self): + """A safe retry must not import CPU gradients from the previous optimizer step.""" + device, _, dtype = initialize_distributed() + first_window = [torch.randn(2, 8, device=device, dtype=dtype) for _ in range(2)] + second_window = [torch.randn(2, 8, device=device, dtype=dtype) for _ in range(2)] + engines = [] + + for _ in range(2): + torch.manual_seed(1234) + model = _IncompleteBackwardModel(hidden_dim=8) + model.backward_control["raise"] = False + config = get_config_dict(2, gradient_accumulation_steps=2) + config["zero_optimization"]["reduce_bucket_size"] = 1_000_000 + config["zero_optimization"]["offload_optimizer"] = {"device": "cpu"} + config["zero_force_ds_cpu_optimizer"] = False + trainable_parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] + engine, _, _, _ = deepspeed.initialize(config=config, model=model, model_parameters=trainable_parameters) + engines.append(engine) + + control, recovered = engines + for microbatch in first_window: + for engine in engines: + engine.backward(engine(microbatch.clone().requires_grad_(True)).sum()) + engine.step() + for engine in engines: + assert engine.optimizer.accumulated_grads_in_cpu + + control.backward(control(second_window[0].clone().requires_grad_(True)).sum()) + control.step() + + recovered.module.backward_control["raise"] = True + with pytest.raises(RuntimeError, match="injected incomplete checkpoint backward"): + recovered.backward(recovered(second_window[0].clone().requires_grad_(True)).sum()) + recovered.backward(recovered(second_window[0].clone().requires_grad_(True)).sum()) + recovered.step() + + for engine in engines: + engine.backward(engine(second_window[1].clone().requires_grad_(True)).sum()) + engine.step() + + for control_parameter, recovered_parameter in zip(_snapshot_trainable_parameters(control), + _snapshot_trainable_parameters(recovered)): + torch.testing.assert_close(recovered_parameter, control_parameter, rtol=0, atol=0) + + control.destroy() + recovered.destroy() + @_RECOMPUTE_RELEASE_TIMING def test_frozen_parameter_without_backward_consumer_releases_at_last_use(self): """Frozen param with no backward consumer should release at its last use, not at epilogue."""