Repository navigation
Conversation
…pointing With Hugging Face per-layer reentrant gradient checkpointing, backward recomputes every decoder layer, including its attention forward. install_attention_stash(model, num_layers) keeps the cuDNN SDPA output and log-sum-exp of the last num_layers decoder layers from the first forward, and the recompute builds the attention node from the kept output instead of running the attention forward again. Projections, norms and RoPE are still recomputed; the replayed backward is the call PyTorch's derivative for _scaled_dot_product_cudnn_attention makes. Kept outputs are keyed by the decoder layer's input (data pointer, version and shape); the reentrant recompute runs on a detached alias of that tensor, so each backward can only use its own forward's output, and outputs of forwards dropped without a backward are released with their input. Unsupported use fails instead of falling back: non-SDPA models, training outside a reentrant checkpoint, attention masks or dropout, and SDPA backends other than cuDNN. Like the fused RoPE installer, it belongs to the attention layers and is installed rather than configured under expert_parallel. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: yh0903 <helloyu0903@gmail.com>
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
With Hugging Face per-layer reentrant gradient checkpointing, backward recomputes every decoder layer, including its attention forward. On Qwen3-30B-A3B on 8 H100s (sequence 4096, micro-batch 4), that recomputed attention forward was about 0.79 s of an 18.7 s step.
This adds an opt-in installer,
install_attention_stash(model, num_layers). For the lastnum_layersdecoder layers, the first forward keeps its attention output, and the recompute uses it instead of running the attention forward again. It is a memory-for-compute trade: one attention output plus its log-sum-exp per stashed layer, held from the layer's forward to its recompute.Change
deepspeed/runtime/activation_checkpointing/attention_stash.py. Like the fused RoPE installer (Add an opt-in fused Triton RoPE for Hugging Face split-half models #8663), it belongs to the attention layers, so it is installed rather than configured underexpert_parallel, and it works with or without AutoEP."sdpa"entry inALL_ATTENTION_FUNCTIONS. Attention modules without a stash call the original function.aten._scaled_dot_product_cudnn_attention, the opF.scaled_dot_product_attentiondispatches to, with the log-sum-exp;aten._scaled_dot_product_cudnn_attention_backward, called with exactly the arguments of PyTorch's derivative for the forward op.weakref.finalize.attn_implementation="sdpa";num_layers, or missing PyTorch entry points.autoep.rstafter the fused RoPE one.Numerics
torch.use_deterministic_algorithms(True), and six autograd runs on the same inputs gave six different dQ results.Tests
tests/unit/runtime/activation_checkpointing/test_attention_stash.py:End to end
Qwen3-30B-A3B on 8 H100s: expert parallel 8 with DeepEP, sequence 4096, micro-batch 4, 16 accumulation steps, per-layer reentrant checkpointing. Measured with an equivalent benchmark prototype on our benchmark stack, one allocation, order 0 / 30 / 48 stashed layers, then reverse: