Skip to content

Commit 5bc41ca

Browse files
committed
test: Add regression test for multi-loss separate backward (#7352)
Signed-off-by: nathon-lee <leejianwoo@gmail.com> fix: add ZeRO-3 second backward after retain_graph=True fails with tensor size mismatch Signed-off-by: nathon-lee <leejianwoo@gmail.com> fix: Stage 3 Temporarily change the exemption from xfail to skip (for this test case only) Signed-off-by: nathon-lee <leejianwoo@gmail.com> fix: Fix ZeRO-3 behavior for two separate backward passes on the same forward graph. Signed-off-by: nathon-lee <leejianwoo@gmail.com>
1 parent 683bd0b commit 5bc41ca

5 files changed

Lines changed: 94 additions & 18 deletions

File tree

deepspeed/runtime/base_optimizer.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -237,6 +237,7 @@ class ZeROOptimizer(DeepSpeedOptimizer):
237237

238238
def __init__(self):
239239
self._backward_hook_state = BackwardHookStateManager()
240+
self.retain_graph_on_current_backward = False
240241

241242
# Delegate backward hook state management to the manager.
242243
# These properties provide backward compatibility with code that accesses
@@ -399,10 +400,14 @@ def backward(self, loss, **kwargs):
399400

400401
scaled_loss = self.backward_prologue(loss)
401402
retain_graph = kwargs.pop('retain_graph', False)
403+
self.retain_graph_on_current_backward = retain_graph
402404
self.enter_backward()
403-
scaled_loss.backward(retain_graph=retain_graph)
404-
self.backward_epilogue()
405-
self.exit_backward()
405+
try:
406+
scaled_loss.backward(retain_graph=retain_graph)
407+
self.backward_epilogue()
408+
finally:
409+
self.exit_backward()
410+
self.retain_graph_on_current_backward = False
406411

407412
def register_grad_acc_post_hook(self, hook):
408413
"""Register a callback to run when all gradient hooks have fired."""

deepspeed/runtime/engine.py

Lines changed: 18 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -2745,23 +2745,27 @@ def backward(self, loss, retain_graph=False, scale_wrt_gas=True):
27452745
# TODO: handle these scaling with direct calls to loss.backward()
27462746
if isinstance(self.optimizer, ZeROOptimizer):
27472747
loss = self.optimizer.scale_if_loss(loss)
2748+
self.optimizer.retain_graph_on_current_backward = retain_graph
27482749
elif self.torch_autocast_z0_gradscaler:
27492750
loss = self.torch_autocast_z0_gradscaler.scale(loss)
27502751

2751-
with compiled_autograd(self._is_compiled_autograd_enabled, self._compile_kwargs):
2752-
if self.zero_optimization() or not self.amp_enabled():
2753-
loss.backward(**backward_kwargs)
2754-
elif self.amp_enabled():
2755-
# AMP requires delaying unscale when inside gradient accumulation boundaries
2756-
# https://nvidia.github.io/apex/advanced.html#gradient-accumulation-across-iterations
2757-
delay_unscale = not self.is_gradient_accumulation_boundary()
2758-
with amp.scale_loss(loss, self.optimizer, delay_unscale=delay_unscale) as scaled_loss:
2759-
scaled_loss.backward(**backward_kwargs)
2760-
2761-
# backward_epilogue is not called in a hook when self._support_torch_style_backward is False
2762-
self._backward_epilogue()
2763-
2764-
self._running_engine_backward = False
2752+
try:
2753+
with compiled_autograd(self._is_compiled_autograd_enabled, self._compile_kwargs):
2754+
if self.zero_optimization() or not self.amp_enabled():
2755+
loss.backward(**backward_kwargs)
2756+
elif self.amp_enabled():
2757+
# AMP requires delaying unscale when inside gradient accumulation boundaries
2758+
# https://nvidia.github.io/apex/advanced.html#gradient-accumulation-across-iterations
2759+
delay_unscale = not self.is_gradient_accumulation_boundary()
2760+
with amp.scale_loss(loss, self.optimizer, delay_unscale=delay_unscale) as scaled_loss:
2761+
scaled_loss.backward(**backward_kwargs)
2762+
2763+
# backward_epilogue is not called in a hook when self._support_torch_style_backward is False
2764+
self._backward_epilogue()
2765+
finally:
2766+
self._running_engine_backward = False
2767+
if isinstance(self.optimizer, ZeROOptimizer):
2768+
self.optimizer.retain_graph_on_current_backward = False
27652769

27662770
return gas_scaled_loss
27672771

deepspeed/runtime/zero/parameter_offload.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -552,7 +552,13 @@ def post_sub_module_backward_function(self, sub_module):
552552
for param in params_to_fetch:
553553
param.data = param.data.t() if len(param.ds_shape) != 1 else param.data
554554

555-
self.get_param_coordinator().release_sub_module(sub_module, forward=False)
555+
# Keep gathered params alive when the current backward retains the graph,
556+
# so a second backward over the same forward can reuse valid saved tensors.
557+
zero_optimizer = getattr(self, "zero_optimizer", None)
558+
retain_graph_backward = bool(zero_optimizer is not None
559+
and getattr(zero_optimizer, "retain_graph_on_current_backward", False))
560+
if not retain_graph_backward:
561+
self.get_param_coordinator().release_sub_module(sub_module, forward=False)
556562

557563
see_memory_usage(
558564
f"After sub module backward function {sub_module.__class__.__name__} {sub_module.ds_id} after release",

deepspeed/runtime/zero/stage3.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -274,6 +274,7 @@ def __init__(
274274
zero_module_granularity_threshold=zero_module_granularity_threshold,
275275
log_trace_cache_warnings=log_trace_cache_warnings,
276276
)
277+
self.parameter_offload.zero_optimizer = self
277278

278279
self.persistent_parameters = self.parameter_offload.persistent_parameters
279280
self._configure_offloading(offload_optimizer_config, offload_param_config)

tests/unit/v1/zero/test_zero_user_backward.py

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -539,6 +539,66 @@ def test_separate_loss_function(self, zero_stage):
539539

540540
model_engine.destroy()
541541

542+
def test_two_losses_separate_backward_gas1(self, zero_stage):
543+
"""Regression test for https://github.com/deepspeedai/DeepSpeed/issues/7352
544+
545+
A single forward followed by two separate backward passes, with
546+
zero_grad() in between, must produce independent gradients for each
547+
loss when gradient_accumulation_steps == 1. Previously the first
548+
backward was treated as the accumulation boundary, which froze loss1's
549+
gradients so that zero_grad() had no effect and loss2's gradients were
550+
doubled (grad2 == grad1 + grad2). Each DeepSpeed gradient set is
551+
compared against an equivalent PyTorch DDP baseline.
552+
"""
553+
hidden_dim = 4
554+
batch_size = 2
555+
556+
# Default gradient_accumulation_steps=1, which is the failing case.
557+
model_ddp, optimizer_ddp, model_engine, device, dtype = setup_models_and_engines(
558+
model_class=SimpleOutputModel, zero_stage=zero_stage, hidden_dim=hidden_dim)
559+
560+
loss_fn = torch.nn.CrossEntropyLoss()
561+
562+
# Two different targets so loss1 and loss2 yield distinct gradients;
563+
# this makes accidental accumulation (grad1 + grad2) detectable.
564+
torch.manual_seed(456)
565+
x = torch.randn(batch_size, hidden_dim, device=device, dtype=dtype)
566+
y1 = torch.randint(0, hidden_dim, (batch_size, ), device=device)
567+
y2 = torch.randint(0, hidden_dim, (batch_size, ), device=device)
568+
569+
# DDP baseline: separate backward for each loss with zero_grad in between.
570+
output_ddp = model_ddp(x)
571+
loss1_ddp = loss_fn(output_ddp, y1)
572+
loss2_ddp = loss_fn(output_ddp, y2)
573+
574+
optimizer_ddp.zero_grad()
575+
loss1_ddp.backward(retain_graph=True)
576+
grads1_ddp = collect_ddp_gradients(model_ddp)
577+
578+
optimizer_ddp.zero_grad()
579+
loss2_ddp.backward()
580+
grads2_ddp = collect_ddp_gradients(model_ddp)
581+
582+
# DeepSpeed: identical sequence.
583+
output_ds = model_engine(x)
584+
loss1_ds = loss_fn(output_ds, y1)
585+
loss2_ds = loss_fn(output_ds, y2)
586+
587+
model_engine.zero_grad()
588+
model_engine.backward(loss1_ds, retain_graph=True)
589+
grads1_ds = collect_gradients_safe(model_engine)
590+
591+
model_engine.zero_grad()
592+
model_engine.backward(loss2_ds)
593+
grads2_ds = collect_gradients_safe(model_engine)
594+
595+
# The second backward must NOT accumulate loss1's gradients on top of
596+
# loss2; both gradient sets must match their DDP counterparts.
597+
compare_gradients(grads1_ddp, grads1_ds, "loss1")
598+
compare_gradients(grads2_ddp, grads2_ds, "loss2")
599+
600+
model_engine.destroy()
601+
542602

543603
class LeafModuleModel(torch.nn.Module):
544604
"""Model with ModuleList that uses all parameters - for testing leaf module compatibility"""

0 commit comments

Comments
 (0)