Skip to content

Commit f873f69

Browse files
authored
Merge pull request #26 from nathon-lee/feature/zero-multi-loss-separate-backward-test
Fix ZeRO-3 so two separate backward passes on the same forward graph work correctly when `retain_graph=True` is used on the first backward.
2 parents 429e2ad + a78b005 commit f873f69

5 files changed

Lines changed: 175 additions & 36 deletions

File tree

deepspeed/runtime/base_optimizer.py

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

253253
def __init__(self):
254254
self._backward_hook_state = BackwardHookStateManager()
255+
self.retain_graph_on_current_backward = False
255256

256257
# Delegate backward hook state management to the manager.
257258
# These properties provide backward compatibility with code that accesses
@@ -419,10 +420,14 @@ def backward(self, loss, **kwargs):
419420

420421
scaled_loss = self.backward_prologue(loss)
421422
retain_graph = kwargs.pop('retain_graph', False)
423+
self.retain_graph_on_current_backward = retain_graph
422424
self.enter_backward()
423-
scaled_loss.backward(retain_graph=retain_graph)
424-
self.backward_epilogue()
425-
self.exit_backward()
425+
try:
426+
scaled_loss.backward(retain_graph=retain_graph)
427+
self.backward_epilogue()
428+
finally:
429+
self.exit_backward()
430+
self.retain_graph_on_current_backward = False
426431

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

deepspeed/runtime/engine.py

Lines changed: 32 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -222,11 +222,6 @@ def active_timers(self):
222222
return self.micro_timers + self.global_timers
223223

224224

225-
def _eigenvalue_summary_events(block_eigenvalue, global_samples):
226-
return [(f"Train/Eigenvalues/ModelBlockParam_{i}", ev_value[0], global_samples)
227-
for i, ev_value in enumerate(block_eigenvalue.values())]
228-
229-
230225
class DeepSpeedEngine(Module):
231226
r"""DeepSpeed engine for training."""
232227

@@ -2759,6 +2754,7 @@ def _backward_epilogue(self):
27592754
if not bf16_optimizer:
27602755
self.optimizer.backward_epilogue()
27612756
self.optimizer.exit_backward()
2757+
self.optimizer.retain_graph_on_current_backward = False
27622758

27632759
if self.is_deepcompile_active():
27642760
deepcompile_backward_epilogue()
@@ -2920,7 +2916,7 @@ def _flush_coalesced_reduction_zero3(self, optimizer):
29202916
optimizer.reduce_ready_partitions_and_remove_grads(param)
29212917
optimizer.independent_gradient_partition_epilogue()
29222918

2923-
def scale(self, loss):
2919+
def scale(self, loss, retain_graph=False):
29242920
r"""Apply loss scaler for manual backward pass.
29252921
29262922
Use this method when calling loss.backward() directly instead of engine.backward().
@@ -2938,6 +2934,8 @@ def scale(self, loss):
29382934
29392935
Arguments:
29402936
loss: Scalar loss tensor to be scaled
2937+
retain_graph: bool, default: false
2938+
forward on user defined choice of retain_graph
29412939
29422940
Returns:
29432941
Scaled loss tensor ready for .backward() call
@@ -2963,6 +2961,7 @@ def scale(self, loss):
29632961
# Apply loss scaler based on optimizer type
29642962
scaled_loss = loss
29652963
if isinstance(self.optimizer, ZeROOptimizer):
2964+
self.optimizer.retain_graph_on_current_backward = retain_graph
29662965
scaled_loss = self.optimizer.scale_if_loss(scaled_loss)
29672966
elif self.torch_autocast_z0_gradscaler:
29682967
scaled_loss = self.torch_autocast_z0_gradscaler.scale(scaled_loss)
@@ -3003,25 +3002,29 @@ def backward(self, loss, retain_graph=False, scale_wrt_gas=True):
30033002
# TODO: handle these scaling with direct calls to loss.backward()
30043003
if isinstance(self.optimizer, ZeROOptimizer):
30053004
loss = self.optimizer.scale_if_loss(loss)
3005+
self.optimizer.retain_graph_on_current_backward = retain_graph
30063006
elif self.torch_autocast_z0_gradscaler:
30073007
loss = self.torch_autocast_z0_gradscaler.scale(loss)
30083008

3009-
with compiled_autograd(self._is_compiled_autograd_enabled, self._compile_kwargs):
3010-
if self.zero_optimization() or not self.amp_enabled():
3011-
loss.backward(**backward_kwargs)
3012-
elif self.amp_enabled():
3013-
# AMP requires delaying unscale when inside gradient accumulation boundaries
3014-
# https://nvidia.github.io/apex/advanced.html#gradient-accumulation-across-iterations
3015-
delay_unscale = not self.is_gradient_accumulation_boundary()
3016-
with amp.scale_loss(loss, self.optimizer, delay_unscale=delay_unscale) as scaled_loss:
3017-
scaled_loss.backward(**backward_kwargs)
3018-
3019-
# backward_epilogue is not called in a hook when self._support_torch_style_backward is False
3020-
self._backward_epilogue()
3021-
3022-
self._running_engine_backward = False
3023-
3024-
return gas_scaled_loss
3009+
try:
3010+
with compiled_autograd(self._is_compiled_autograd_enabled, self._compile_kwargs):
3011+
if self.zero_optimization() or not self.amp_enabled():
3012+
loss.backward(**backward_kwargs)
3013+
elif self.amp_enabled():
3014+
# AMP requires delaying unscale when inside gradient accumulation boundaries
3015+
# https://nvidia.github.io/apex/advanced.html#gradient-accumulation-across-iterations
3016+
delay_unscale = not self.is_gradient_accumulation_boundary()
3017+
with amp.scale_loss(loss, self.optimizer, delay_unscale=delay_unscale) as scaled_loss:
3018+
scaled_loss.backward(**backward_kwargs)
3019+
3020+
# backward_epilogue is not called in a hook when self._support_torch_style_backward is False
3021+
self._backward_epilogue()
3022+
3023+
return gas_scaled_loss
3024+
finally:
3025+
self._running_engine_backward = False
3026+
if isinstance(self.optimizer, ZeROOptimizer):
3027+
self.optimizer.retain_graph_on_current_backward = False
30253028

30263029
def is_gradient_accumulation_boundary(self):
30273030
"""
@@ -3223,8 +3226,13 @@ def step(self, lr_kwargs=None):
32233226

32243227
if (self.eigenvalue_enabled()
32253228
and not self.gas_boundary_ctr % self.eigenvalue_gas_boundary_resolution()):
3226-
self.summary_events.extend(
3227-
_eigenvalue_summary_events(self.block_eigenvalue, self.global_samples))
3229+
ev_values = self.block_eigenvalue.values()
3230+
for i in range(len(ev_values)):
3231+
self.summary_events.append((
3232+
f"Train/Eigenvalues/ModelBlockParam_{i}",
3233+
self.ev_values[i][0],
3234+
self.global_samples,
3235+
))
32283236
self.monitor.write_events(self.summary_events)
32293237

32303238
# Check flops profiling
@@ -4327,9 +4335,6 @@ def save_checkpoint(self, save_dir, tag=None, client_state={}, save_latest=True,
43274335
process with rank 0.
43284336
43294337
"""
4330-
if not save_dir:
4331-
raise ValueError(f"save_dir must be a non-empty string, got {save_dir!r}")
4332-
43334338
if self._optimizer_has_ckpt_event_prologue():
43344339
# Custom preparation for checkpoint save, if applicable
43354340
self.optimizer.checkpoint_event_prologue()

deepspeed/runtime/zero/parameter_offload.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -136,6 +136,7 @@ def __init__(
136136
zero_quantized_nontrainable_weights=False,
137137
zero_module_granularity_threshold=0,
138138
log_trace_cache_warnings=False,
139+
retain_graph_checker=None,
139140
):
140141

141142
see_memory_usage("DeepSpeedZeRoOffload initialize [begin]", force=False)
@@ -153,6 +154,7 @@ def __init__(
153154
self.zero_quantized_weights = zero_quantized_weights
154155
self.zero_quantized_nontrainable_weights = zero_quantized_nontrainable_weights
155156
self.log_trace_cache_warnings = log_trace_cache_warnings
157+
self.retain_graph_checker = retain_graph_checker
156158

157159
if offload_param_config is not None and offload_param_config.device != OffloadDeviceEnum.none:
158160
self.offload_device = offload_param_config.device
@@ -563,7 +565,11 @@ def post_sub_module_backward_function(self, sub_module):
563565
for param in params_to_fetch:
564566
param.data = param.data.t() if len(param.ds_shape) != 1 else param.data
565567

566-
self.get_param_coordinator().release_sub_module(sub_module, forward=False)
568+
# Keep gathered params alive when the current backward retains the graph,
569+
# so a second backward over the same forward can reuse valid saved tensors.
570+
retain_graph_backward = bool(self.retain_graph_checker()) if self.retain_graph_checker is not None else False
571+
if not retain_graph_backward:
572+
self.get_param_coordinator().release_sub_module(sub_module, forward=False)
567573

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

deepspeed/runtime/zero/stage3.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -285,6 +285,7 @@ def __init__(
285285
zero_quantized_nontrainable_weights=zero_quantized_nontrainable_weights,
286286
zero_module_granularity_threshold=zero_module_granularity_threshold,
287287
log_trace_cache_warnings=log_trace_cache_warnings,
288+
retain_graph_checker=lambda: self.retain_graph_on_current_backward,
288289
)
289290

290291
self.persistent_parameters = self.parameter_offload.persistent_parameters
@@ -571,6 +572,7 @@ def initialize_ds_offload(
571572
zero_quantized_nontrainable_weights,
572573
zero_module_granularity_threshold,
573574
log_trace_cache_warnings,
575+
retain_graph_checker=None,
574576
):
575577
return DeepSpeedZeRoOffload(module=module,
576578
timers=timers,
@@ -589,7 +591,8 @@ def initialize_ds_offload(
589591
zero_quantized_weights=zero_quantized_weights,
590592
zero_quantized_nontrainable_weights=zero_quantized_nontrainable_weights,
591593
zero_module_granularity_threshold=zero_module_granularity_threshold,
592-
log_trace_cache_warnings=log_trace_cache_warnings)
594+
log_trace_cache_warnings=log_trace_cache_warnings,
595+
retain_graph_checker=retain_graph_checker)
593596

594597
def _get_param_partition_group(self, param):
595598
return getattr(param, "ds_process_group", self.dp_process_group)
@@ -3128,8 +3131,7 @@ def state_dict(self):
31283131
torch.save(checkpoint, "saved.pth")
31293132
"""
31303133
if self.elastic_checkpoint:
3131-
raise NotImplementedError(
3132-
"ZeRO-3 elastic checkpointing is deprecated and unsupported. Use Universal Checkpointing instead.")
3134+
raise NotImplementedError("ZeRO-3 does not yet support elastic checkpointing, please disable for now.")
31333135

31343136
return self._rigid_state_dict()
31353137

@@ -3289,8 +3291,7 @@ def load_state_dict(self,
32893291
"""
32903292

32913293
if self.elastic_checkpoint:
3292-
raise NotImplementedError(
3293-
"ZeRO-3 elastic checkpointing is deprecated and unsupported. Use Universal Checkpointing instead.")
3294+
raise NotImplementedError("ZeRO-3 does not yet support elastic checkpointing, please disable for now.")
32943295

32953296
if checkpoint_folder:
32963297
self._load_universal_checkpoint(checkpoint_folder, load_optimizer_states, load_from_fp32_weights)

tests/unit/v1/zero/test_zero_user_backward.py

Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -539,6 +539,128 @@ 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(model_class=SimpleOutputModel,
558+
zero_stage=zero_stage,
559+
hidden_dim=hidden_dim)
560+
561+
loss_fn = torch.nn.CrossEntropyLoss()
562+
563+
# Two different targets so loss1 and loss2 yield distinct gradients;
564+
# this makes accidental accumulation (grad1 + grad2) detectable.
565+
torch.manual_seed(456)
566+
x = torch.randn(batch_size, hidden_dim, device=device, dtype=dtype)
567+
y1 = torch.randint(0, hidden_dim, (batch_size, ), device=device)
568+
y2 = torch.randint(0, hidden_dim, (batch_size, ), device=device)
569+
570+
# DDP baseline: separate backward for each loss with zero_grad in between.
571+
output_ddp = model_ddp(x)
572+
loss1_ddp = loss_fn(output_ddp, y1)
573+
loss2_ddp = loss_fn(output_ddp, y2)
574+
575+
optimizer_ddp.zero_grad()
576+
loss1_ddp.backward(retain_graph=True)
577+
grads1_ddp = collect_ddp_gradients(model_ddp)
578+
579+
optimizer_ddp.zero_grad()
580+
loss2_ddp.backward()
581+
grads2_ddp = collect_ddp_gradients(model_ddp)
582+
583+
# DeepSpeed: identical sequence.
584+
output_ds = model_engine(x)
585+
loss1_ds = loss_fn(output_ds, y1)
586+
loss2_ds = loss_fn(output_ds, y2)
587+
588+
model_engine.zero_grad()
589+
model_engine.backward(loss1_ds, retain_graph=True)
590+
grads1_ds = collect_gradients_safe(model_engine)
591+
592+
model_engine.zero_grad()
593+
model_engine.backward(loss2_ds)
594+
grads2_ds = collect_gradients_safe(model_engine)
595+
596+
# The second backward must NOT accumulate loss1's gradients on top of
597+
# loss2; both gradient sets must match their DDP counterparts.
598+
compare_gradients(grads1_ddp, grads1_ds, "loss1")
599+
compare_gradients(grads2_ddp, grads2_ds, "loss2")
600+
601+
model_engine.destroy()
602+
603+
def test_two_losses_separate_manual_backward_gas1(self, zero_stage):
604+
"""Regression test for https://github.com/deepspeedai/DeepSpeed/issues/7352
605+
606+
Same scenario as test_two_losses_separate_backward_gas1, but using the
607+
torch-style manual path engine.scale(loss, retain_graph=True).backward(
608+
retain_graph=True) instead of engine.backward(loss, retain_graph=True).
609+
The manual path bypasses engine.backward(), so retain_graph must be
610+
propagated through scale(). For ZeRO-3 this defers parameter release so
611+
the retained graph's saved tensors stay valid for the second backward
612+
over the same forward.
613+
"""
614+
hidden_dim = 4
615+
batch_size = 2
616+
617+
# Default gradient_accumulation_steps=1, which is the failing case.
618+
model_ddp, optimizer_ddp, model_engine, device, dtype = setup_models_and_engines(model_class=SimpleOutputModel,
619+
zero_stage=zero_stage,
620+
hidden_dim=hidden_dim)
621+
622+
loss_fn = torch.nn.CrossEntropyLoss()
623+
624+
# Two different targets so loss1 and loss2 yield distinct gradients;
625+
# this makes accidental accumulation (grad1 + grad2) detectable.
626+
torch.manual_seed(456)
627+
x = torch.randn(batch_size, hidden_dim, device=device, dtype=dtype)
628+
y1 = torch.randint(0, hidden_dim, (batch_size, ), device=device)
629+
y2 = torch.randint(0, hidden_dim, (batch_size, ), device=device)
630+
631+
# DDP baseline: separate backward for each loss with zero_grad in between.
632+
output_ddp = model_ddp(x)
633+
loss1_ddp = loss_fn(output_ddp, y1)
634+
loss2_ddp = loss_fn(output_ddp, y2)
635+
636+
optimizer_ddp.zero_grad()
637+
loss1_ddp.backward(retain_graph=True)
638+
grads1_ddp = collect_ddp_gradients(model_ddp)
639+
640+
optimizer_ddp.zero_grad()
641+
loss2_ddp.backward()
642+
grads2_ddp = collect_ddp_gradients(model_ddp)
643+
644+
# DeepSpeed: identical sequence via the manual scale().backward() path.
645+
output_ds = model_engine(x)
646+
loss1_ds = loss_fn(output_ds, y1)
647+
loss2_ds = loss_fn(output_ds, y2)
648+
649+
model_engine.zero_grad()
650+
model_engine.scale(loss1_ds, retain_graph=True).backward(retain_graph=True)
651+
grads1_ds = collect_gradients_safe(model_engine)
652+
653+
model_engine.zero_grad()
654+
model_engine.scale(loss2_ds).backward()
655+
grads2_ds = collect_gradients_safe(model_engine)
656+
657+
# The second backward must NOT accumulate loss1's gradients on top of
658+
# loss2; both gradient sets must match their DDP counterparts.
659+
compare_gradients(grads1_ddp, grads1_ds, "loss1")
660+
compare_gradients(grads2_ddp, grads2_ds, "loss2")
661+
662+
model_engine.destroy()
663+
542664

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

0 commit comments

Comments
 (0)