Skip to content
Merged
Show file tree
Hide file tree
Changes from 36 commits
Commits
Show all changes
39 commits
Select commit Hold shift + click to select a range
8a27283
rename backward prologue method
tohtana Oct 17, 2025
d85cfd9
refactor loss scaling
tohtana Oct 18, 2025
bded5c8
refactor backward
tohtana Oct 18, 2025
1f413d6
fix for bf16 optimizer
tohtana Oct 18, 2025
cc87977
simplify preprocess/postprocess of backward
tohtana Oct 18, 2025
95018a3
fix order of backward postprocess
tohtana Oct 19, 2025
80d0e7d
enable non-scalar backward only for ZeROOptimizer
tohtana Oct 19, 2025
db70476
fix zero+fp16 case
tohtana Oct 19, 2025
50b29d8
add config to enable allow_user_backward
tohtana Oct 19, 2025
076b187
fix flag for error handling
tohtana Oct 19, 2025
f6748d1
resolve conflict
tohtana Nov 3, 2025
280b1fa
add test cases
tohtana Nov 7, 2025
5d5e64e
Merge branch 'master' into tohtana/backward_non_scalar
tohtana Nov 7, 2025
6ce26f3
fix format
tohtana Nov 7, 2025
0c579d5
return scaled loss from engine's backward
tohtana Nov 7, 2025
1615036
Merge branch 'master' into tohtana/backward_non_scalar
tohtana Nov 10, 2025
c8758f7
remove option to enable user backward
tohtana Nov 11, 2025
9962f2c
add hook utility
tohtana Nov 12, 2025
a8f15a0
fix for z2
tohtana Nov 12, 2025
1d0a721
fix scaling
tohtana Nov 12, 2025
39372ac
exclude unused params from counter
tohtana Nov 13, 2025
7eacbc7
set default flag
tohtana Nov 13, 2025
6cce937
handle non-zero optimizer
tohtana Nov 13, 2025
b72b5a7
call epilogue in engine's backward
tohtana Nov 13, 2025
1307a87
prevent hooks from being called from nested backward
tohtana Nov 13, 2025
98cc865
run post hook fo rz3
tohtana Nov 13, 2025
adb6990
Merge branch 'master' into tohtana/backward_non_scalar
tohtana Nov 13, 2025
9328dfa
added comments
tohtana Nov 13, 2025
01b3251
remove hard-coded tolerances
tohtana Nov 13, 2025
78f7ad4
add test for multiple engines
tohtana Nov 13, 2025
73f7ff1
update document
tohtana Nov 14, 2025
26308cd
remove deprecated comment
tohtana Nov 17, 2025
9963546
simplify utility func to count effective grad nodes
tohtana Nov 17, 2025
b730f46
fix combination with leaf module
tohtana Nov 17, 2025
08b1599
refactor tests
tohtana Nov 17, 2025
92d3068
refactor tests
tohtana Nov 18, 2025
ebac40b
fix loss scaling
tohtana Nov 18, 2025
fcf7c8c
Merge branch 'master' into tohtana/backward_non_scalar
tohtana Nov 18, 2025
90e1b7d
Merge branch 'master' into tohtana/backward_non_scalar
tohtana Nov 18, 2025
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
63 changes: 63 additions & 0 deletions deepspeed/runtime/base_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,13 @@

import os
import torch
from typing import Any

from deepspeed.utils import logger
from deepspeed.utils.tensor_fragment import map_to_flat_opt_states
from deepspeed.runtime.utils import bwc_tensor_model_parallel_rank, see_memory_usage
from deepspeed.runtime.torch_autocast import get_comm_dtype, is_autocast_initialized
from deepspeed.runtime.utils import maybe_loss_for_backward


class DeepSpeedOptimizer(object):
Expand All @@ -18,6 +20,11 @@ class DeepSpeedOptimizer(object):

class ZeROOptimizer(DeepSpeedOptimizer):

def __init__(self):
self._remaining_grad_acc_hooks = 0
self._grad_acc_post_hooks = []
self._backward_active_depth = 0

def load_hp_checkpoint_state_from_checkpoint_dir(self, lp_groups_name: str, checkpoint_dir: str) -> None:
checkpoint_dir = os.path.join(checkpoint_dir, "zero")
optim_state_path = os.path.join(checkpoint_dir, "optimizer_state.pt")
Expand Down Expand Up @@ -79,3 +86,59 @@ def get_param_comm_dtype(self, param):
return get_comm_dtype(param)
else:
return self.communication_data_type

def scale_if_loss(self, value: Any) -> Any:
"""
Applies loss scaling to the input value if it is a loss tensor.
"""
if maybe_loss_for_backward(value):
if self.custom_loss_scaler:
return self.external_loss_scale * value
if self.torch_autocast_gradscaler:
return self.torch_autocast_gradscaler.scale(value)
return self.loss_scaler.scale_loss(value)

return value

def backward_prologue(self):
pass

def backward_epilogue(self, **kwargs):
pass

def backward(self, loss, **kwargs):
assert maybe_loss_for_backward(loss), "Optimizer's backward() only accepts a scalar tensor"

scaled_loss = self.backward_prologue(loss)
retain_graph = kwargs.pop('retain_graph', False)
self.enter_backward()
scaled_loss.backward(retain_graph=retain_graph)
self.backward_epilogue()
self.exit_backward()

def register_grad_acc_post_hook(self, hook):
self._grad_acc_post_hooks.append(hook)

def unregister_grad_acc_post_hooks(self):
self._grad_acc_post_hooks = []

def run_grad_acc_post_hooks(self):
# Custom autograd Functions (e.g., TiledFusedLogitsLoss) can invoke
# `torch.autograd.backward()` from their *forward* pass before the user
# ever calls `engine.backward(loss)`. Those early backward calls still
# trigger ZeRO's grad hooks, but we must not run the engine's
# post-backward logic (which reduces/clears grads) until the outer/user
# backward is active. The depth guard filters out only those pre-user
# invocations while still allowing backward calls that happen during
# the real user backward.
if self._backward_active_depth == 0:
return
for hook in self._grad_acc_post_hooks:
hook()

def enter_backward(self):
self._backward_active_depth += 1

def exit_backward(self):
if self._backward_active_depth > 0:
self._backward_active_depth -= 1
12 changes: 2 additions & 10 deletions deepspeed/runtime/bf16_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -316,18 +316,10 @@ def step(self, closure=None):

self.clear_hp_grads()

def backward(self, loss, retain_graph=False, update_hp_grads=True, clear_lp_grads=False, **bwd_kwargs):
"""Perform a backward pass and copy the low-precision gradients to the
high-precision copy.

We copy/accumulate to the high-precision grads now to prevent accumulating in the
bf16 grads after successive backward() calls (i.e., grad accumulation steps > 1)

The low-precision grads are deallocated during this procedure.
"""
def backward_prologue(self):
self.clear_lp_grads()
loss.backward(retain_graph=retain_graph, **bwd_kwargs)

def backward_epilogue(self, update_hp_grads=True, clear_lp_grads=False, **bwd_kwargs):
if update_hp_grads:
self.update_hp_grads(clear_lp_grads=clear_lp_grads)

Expand Down
153 changes: 94 additions & 59 deletions deepspeed/runtime/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,8 +24,9 @@
import deepspeed

from deepspeed import comm as dist
from deepspeed.runtime.utils import see_memory_usage, DummyOptim
from deepspeed.runtime.utils import see_memory_usage, DummyOptim, register_output_backward_hooks, check_internal_apis_for_count_used_parameters
from .zero.offload_config import OffloadDeviceEnum, OffloadStateTypeEnum
from deepspeed.runtime.base_optimizer import ZeROOptimizer
from deepspeed.runtime.zero.stage_1_and_2 import DeepSpeedZeroOptimizer
from deepspeed.runtime.zenflow.zenflow_stage_1_and_2 import ZenFlowZeroOptimizer
from deepspeed.runtime.zero.partition_parameters import ZeroParamStatus
Expand Down Expand Up @@ -85,7 +86,7 @@
from deepspeed.utils.debug import debug_extract_module_and_param_names, debug_clear_module_and_param_names
from deepspeed.monitor.monitor import MonitorMaster
from deepspeed.runtime.progressive_layer_drop import ProgressiveLayerDrop
from deepspeed.runtime.utils import clip_grad_norm_, compare_tensors_in_structures
from deepspeed.runtime.utils import clip_grad_norm_, compare_tensors_in_structures, maybe_loss_for_backward
from deepspeed.runtime.eigenvalue import Eigenvalue
from deepspeed.runtime.data_pipeline.constants import DATA_SAMPLING, \
DATA_ROUTING, DATA_SAMPLING_ENABLED, CURRICULUM_LEARNING, \
Expand Down Expand Up @@ -420,6 +421,28 @@ def __init__(self,
self.register_compile_pass(selective_gather.NAME, selective_gather.selective_gather)
self.register_compile_pass(offload_adam_states.NAME, offload_adam_states.move_opt_states)

# We now support PyTorch style backward, but it relies on the counter in ZeRO optimizers.
# However, we need some internal APIs to count the number of only used parameters.
# So we only enable this feature when those internal APIs are available.
# Otherwise, we fallback to DeepSpeed style backward only.
# See `count_used_parameters_in_backward` for more details.
self._running_engine_backward = False
self._support_torch_style_backward = 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)`,
# not for `engine.backward(loss)`. In this case, we need to ensure that the preprocessing
# and postprocessing around the backward call are handled correctly.
# However, we cannot use `register_full_backward_hook` for post-backward hooks.
# If none of the module inputs require gradients, `register_full_backward_hook` fires
# when the gradients of the module outputs are computed. Our gradient
# accumulation hooks are called later. But we want `_backward_post_hook` to be called
# only after all gradients have been computed.
# To handle this, the optimizer maintains a counter to track the number of gradients
# that have been computed. When all gradients are ready, it calls `_backward_post_hook`.
# See also: https://pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.register_full_backward_hook
self.optimizer.register_grad_acc_post_hook(self._backward_post_hook)

def _optimized_linear_offload_setup(self):
self.optimized_linear_base_weight_sharding = False
self.optimized_linear_lora_enabled = False
Expand Down Expand Up @@ -2184,6 +2207,13 @@ def forward(self, *inputs, **kwargs):
with autocast_if_enabled(self):
loss = self.module(*inputs, **kwargs)

# Register output backward hooks
# preprocess_once_fn is called for preprocessing
# preprocess_per_tensor_fn scales a tensor for gradient accumulation
register_output_backward_hooks(loss,
preprocess_once_fn=self._backward_prologue,
preprocess_per_tensor_fn=self._backward_prologue_per_tensor)

if self.autotuning_profile_model_info():
activation_mem = get_ma_status() - ma
self.autotuning_model_info["activation_mem_per_gpu"] = activation_mem
Expand Down Expand Up @@ -2255,76 +2285,54 @@ def allreduce_gradients(self, bucket_size=MEMORY_OPT_ALLREDUCE_SIZE):
elif self.zenflow:
self.optimizer.reduce_gradients(pipeline_parallel=self.pipeline_parallelism)

def _backward_prologue(self, loss, scale_wrt_gas=True):
see_memory_usage("Engine before backward", force=self.memory_breakdown())
if self.scale_wrt_gas is not None:
scale_wrt_gas = self.scale_wrt_gas
def _backward_prologue(self):
self._start_timers(self.engine_timers.backward_timers)

# scale loss w.r.t. gradient accumulation if reduction is not disabled
do_gradient_reduction = self.enable_backward_allreduce and not self.inside_no_sync_ctxt and not self.is_deepcompile_active(
)
if do_gradient_reduction and self.gradient_accumulation_steps() > 1 and scale_wrt_gas:
loss = self._scale_loss_by_gas(loss.float())
# When necessary internal APIs are not available, we disable direct calls to tensor.backward()
# and limit to engine.backward(loss) only.
if not self._support_torch_style_backward and not self._running_engine_backward:
raise RuntimeError(
"Direct calls to tensor.backward() are not supported with this PyTorch version. Please use engine.backward(loss) instead."
)

# Log training loss
mean_loss = loss.mean().detach()
self.losses = mean_loss if self.losses is None else self.losses + mean_loss
if self.monitor.enabled:
if self.is_gradient_accumulation_boundary():
if self.global_rank == 0:
self.summary_events = [(
"Train/Samples/train_loss",
self.losses.item(),
self.global_samples,
)]
self.monitor.write_events(self.summary_events)
see_memory_usage("Engine before backward", force=self.memory_breakdown())

assert not self.eigenvalue_enabled(), "Eigenvalue is not supported with non-scalar backward"
assert not self.amp_enabled(), "Apex AMP is not supported with non-scalar backward"

if self.is_deepcompile_active():
deepcompile_backward_prologue(self.is_gradient_accumulation_boundary())

if isinstance(self.optimizer, ZeROOptimizer):
self.optimizer.backward_prologue()
self.optimizer.enter_backward()

if self.zenflow and self.auto_update:
self.optimizer.zenflow_state ^= 1

return loss
if self.zero_optimization():
self.optimizer.is_gradient_accumulation_boundary = self.is_gradient_accumulation_boundary()

def _backward_epilogue(self):
self._start_timers(self.engine_timers.backward_reduce_timers)
if self.enable_backward_allreduce and not self.inside_no_sync_ctxt:
# Traditional code path that allreduces the module parameter grads
self.allreduce_gradients()

self._stop_timers(self.engine_timers.backward_reduce_timers)
if isinstance(self.optimizer, ZeROOptimizer):
self.optimizer.backward_epilogue()
self.optimizer.exit_backward()

see_memory_usage("Engine after backward", force=self.memory_breakdown())
self._stop_timers(self.engine_timers.backward_reduce_timers)
self._stop_timers(self.engine_timers.backward_timers)

def _do_optimizer_backward(self, loss, retain_graph):
self._start_timers(self.engine_timers.backward_inner_timers)
if self.zero_optimization():
self.optimizer.is_gradient_accumulation_boundary = self.is_gradient_accumulation_boundary()
self.optimizer.backward(loss, retain_graph=retain_graph)
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(retain_graph=retain_graph)
elif self.fp16_enabled():
if self.eigenvalue_enabled():
self.optimizer.backward(loss, create_graph=True, retain_graph=True)
else:
self.optimizer.backward(loss, retain_graph=retain_graph)
elif self.bfloat16_enabled():
self.optimizer.backward(loss, retain_graph=retain_graph)
else:
if self.torch_autocast_z0_gradscaler:
if self.eigenvalue_enabled():
self.torch_autocast_z0_gradscaler.scale(loss).backward(create_graph=True, retain_graph=True)
else:
self.torch_autocast_z0_gradscaler.scale(loss).backward(retain_graph=retain_graph)
elif self.eigenvalue_enabled():
loss.backward(create_graph=True, retain_graph=True)
else:
loss.backward(retain_graph=retain_graph)
self._stop_timers(self.engine_timers.backward_inner_timers)
def _backward_prologue_per_tensor(self, grad):
return grad / self.gradient_accumulation_steps()

def _backward_post_hook(self):
if not self._running_engine_backward:
self._backward_epilogue()

@contextmanager
def no_sync(self):
Expand Down Expand Up @@ -2356,14 +2364,41 @@ def backward(self, loss, retain_graph=False, scale_wrt_gas=True):
"""
assert self.optimizer is not None and not isinstance(self.optimizer, DummyOptim), \
"must provide optimizer during init in order to use backward"
assert maybe_loss_for_backward(
loss), "loss must be a scalar tensor. If you need to pass output gradients, backward() of output tensors"

self._start_timers(self.engine_timers.backward_timers)
loss = self._backward_prologue(loss, scale_wrt_gas)
self._do_optimizer_backward(loss, retain_graph)
self._running_engine_backward = True

# 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)

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()
self._stop_timers(self.engine_timers.backward_timers)

return loss
self._running_engine_backward = False

return gas_scaled_loss

def is_gradient_accumulation_boundary(self):
"""
Expand Down
8 changes: 7 additions & 1 deletion deepspeed/runtime/fp16/loss_scaler.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,8 +60,14 @@ def scale_gradient(self, module, grad_in, grad_out):
def update_scale(self, overflow):
pass

def scale_loss(self, loss):
""" Scales the loss by the current loss scale.
We need this function to scale loss without calling backward on it.
"""
return loss * self.loss_scale

def backward(self, loss, retain_graph=False):
scaled_loss = loss * self.loss_scale
scaled_loss = self.scale_loss(loss)
scaled_loss.backward(retain_graph=retain_graph)
# print(f'LossScalerBackward: {scaled_loss=}')

Expand Down
Loading
Loading