@@ -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
543665class LeafModuleModel (torch .nn .Module ):
544666 """Model with ModuleList that uses all parameters - for testing leaf module compatibility"""
0 commit comments