From 92302850b658b58800f9250930c9187e2978e74a Mon Sep 17 00:00:00 2001 From: Masahiro Tanaka Date: Mon, 10 Aug 2026 17:37:45 -0700 Subject: [PATCH 1/3] Fix shared loss gradient accumulation Signed-off-by: Masahiro Tanaka --- deepspeed/runtime/engine.py | 72 +++++++------- .../zero/test_zero_shared_loss_gradient.py | 97 +++++++++++++++++++ 2 files changed, 135 insertions(+), 34 deletions(-) create mode 100644 tests/unit/runtime/zero/test_zero_shared_loss_gradient.py diff --git a/deepspeed/runtime/engine.py b/deepspeed/runtime/engine.py index 5db39d237967..7594c45422dd 100755 --- a/deepspeed/runtime/engine.py +++ b/deepspeed/runtime/engine.py @@ -492,6 +492,7 @@ def __init__(self, self._support_torch_style_backward = False # Flag to control whether gradients should be scaled by gradient accumulation steps self._scale_wrt_gas = True + self._is_engine_backward_loss_scaled = False 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)`, @@ -2952,7 +2953,7 @@ 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: + if grad is not None and self._scale_wrt_gas and not self._is_engine_backward_loss_scaled: return grad / self.gradient_accumulation_steps() return grad @@ -3178,42 +3179,45 @@ def backward(self, loss, retain_graph=False, scale_wrt_gas=True): self._running_engine_backward = True # Store scale_wrt_gas so the hook can respect it self._scale_wrt_gas = scale_wrt_gas + previous_engine_backward_loss_scaled = self._is_engine_backward_loss_scaled + self._is_engine_backward_loss_scaled = scale_wrt_gas + 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 - # 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) - - # 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) + + 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: + self._is_engine_backward_loss_scaled = previous_engine_backward_loss_scaled + self._running_engine_backward = False def is_gradient_accumulation_boundary(self): """ 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..de80178fe4a7 --- /dev/null +++ b/tests/unit/runtime/zero/test_zero_shared_loss_gradient.py @@ -0,0 +1,97 @@ +# Copyright (c) DeepSpeed Team. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +import pytest +import torch + +import deepspeed +import deepspeed.comm as dist +from deepspeed.accelerator import get_accelerator +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) + + +def _make_model(device): + model = SharedLinear().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) + + +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()) + + model = _make_model(device) + 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) + 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() From 66038d40af4c97a5f4c21ba0e5925ef49061323a Mon Sep 17 00:00:00 2001 From: Masahiro Tanaka Date: Tue, 11 Aug 2026 15:54:10 -0700 Subject: [PATCH 2/3] Fix multi-engine managed backward GAS scaling Signed-off-by: Masahiro Tanaka --- deepspeed/runtime/engine.py | 116 ++++++++++-- .../zero/test_zero_shared_loss_gradient.py | 168 ++++++++++++++++-- 2 files changed, 253 insertions(+), 31 deletions(-) diff --git a/deepspeed/runtime/engine.py b/deepspeed/runtime/engine.py index 7594c45422dd..5f1a8cd61073 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,12 +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 - self._is_engine_backward_loss_scaled = False 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)`, @@ -2952,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 and not self._is_engine_backward_loss_scaled: + 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 @@ -3176,11 +3259,11 @@ 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 - previous_engine_backward_loss_scaled = self._is_engine_backward_loss_scaled - self._is_engine_backward_loss_scaled = scale_wrt_gas + 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(): @@ -3201,6 +3284,11 @@ def backward(self, loss, retain_graph=False, scale_wrt_gas=True): 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) @@ -3216,8 +3304,12 @@ def backward(self, loss, retain_graph=False, scale_wrt_gas=True): return gas_scaled_loss finally: - self._is_engine_backward_loss_scaled = previous_engine_backward_loss_scaled - self._running_engine_backward = False + 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): """ diff --git a/tests/unit/runtime/zero/test_zero_shared_loss_gradient.py b/tests/unit/runtime/zero/test_zero_shared_loss_gradient.py index de80178fe4a7..c6e47b4bac4e 100644 --- a/tests/unit/runtime/zero/test_zero_shared_loss_gradient.py +++ b/tests/unit/runtime/zero/test_zero_shared_loss_gradient.py @@ -5,10 +5,12 @@ 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 @@ -24,8 +26,25 @@ def forward(self, inputs): return self.shared(inputs) -def _make_model(device): - model = SharedLinear().to(device=device, dtype=torch.bfloat16) +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)) @@ -37,6 +56,41 @@ def _inputs(device, rank, micro_step): 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 @@ -64,23 +118,7 @@ def test_engine_and_module_branches_match_manual_reference(self): microbatch_grad.div_(world_size) reference_grad.add_(microbatch_grad.float()) - model = _make_model(device) - 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) + engine = _make_engine(_make_model(device), gradient_accumulation_steps) try: for micro_step in range(gradient_accumulation_steps): inputs = _inputs(device, rank, micro_step) @@ -95,3 +133,95 @@ def test_engine_and_module_branches_match_manual_reference(self): 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() From 02736449963e139397a1eae1a46319fc6ee1b60a Mon Sep 17 00:00:00 2001 From: Masahiro Tanaka Date: Wed, 12 Aug 2026 09:15:56 -0700 Subject: [PATCH 3/3] Recover or reject aborted ZeRO backward state Discard unfinished ZeRO hook, bucket, and unreduced gradient state at the next root forward when an aborted backward has not crossed a reduction boundary. Track actual bucket drains and fail closed before forward or step once an accumulation window may contain an irreversible contribution. Preserve step guards across runtime optimizer replacement and restore ZeRO-1/2 micro-step state for safe CPU-offload retry. Reject incomplete backward state at the public engine step entry, including non-boundary accumulation steps. Add lifecycle coverage for clean-control parity, oversized ZeRO-2 gradients, BF16 immediate updates, ZeRO-1 accumulation, coalesced ZeRO-2/3 gradients, multi-bucket ZeRO-2/3 failures, ZeRO-3 leaf reductions, step-before-forward rejection, CPU offload, and GAS>1 non-boundary step rejection. Signed-off-by: Masahiro Tanaka --- deepspeed/runtime/base_optimizer.py | 175 ++++++++- deepspeed/runtime/engine.py | 3 + ...st_zero_activation_checkpoint_lifecycle.py | 349 ++++++++++++++++++ 3 files changed, 524 insertions(+), 3 deletions(-) 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 5f1a8cd61073..15c6ee54f496 100755 --- a/deepspeed/runtime/engine.py +++ b/deepspeed/runtime/engine.py @@ -3465,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/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."""