Skip to content

Add an opt-in attention-output stash for Hugging Face reentrant checkpointing - #8759

Draft
yh0903 wants to merge 1 commit into
deepspeedai:masterfrom
yh0903:yh0903/autoep-attention-output-stash
Draft

yh0903 wants to merge 1 commit into
deepspeedai:masterfrom
yh0903:yh0903/autoep-attention-output-stash

Conversation

@yh0903

@yh0903 yh0903 commented Oct 6, 2026

Copy link
Copy Markdown
Contributor

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 last num_layers decoder 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

  • New 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 under expert_parallel, and it works with or without AutoEP.
  • The installer wraps Hugging Face's "sdpa" entry in ALL_ATTENTION_FUNCTIONS. Attention modules without a stash call the original function.
  • First forward, inside the reentrant checkpoint (grad disabled, the decoder layer's input requires grad):
    • calls aten._scaled_dot_product_cudnn_attention, the op F.scaled_dot_product_attention dispatches to, with the log-sum-exp;
    • keeps the output.
  • Recompute (grad enabled inside a backward pass):
    • projections, norms and RoPE are still recomputed;
    • the attention node is an autograd function that returns the kept output;
    • its backward is aten._scaled_dot_product_cudnn_attention_backward, called with exactly the arguments of PyTorch's derivative for the forward op.
  • Each kept output is keyed by its decoder layer's input: data pointer, version and shape, recorded by a forward pre-hook. The reentrant recompute runs on a detached alias of that same tensor, so a backward can only use its own forward's output. An output whose forward's graph is freed without a backward is released with that input, through weakref.finalize.
  • Fails instead of falling back, for:
    • a model without attn_implementation="sdpa";
    • a training forward of a stashed layer outside a Hugging Face reentrant checkpoint, including non-reentrant checkpointing;
    • an attention mask (padding) or attention dropout;
    • a PyTorch SDPA backend choice other than cuDNN;
    • a recompute with no kept output;
    • a second install on the same layer, a non-integer or out-of-range num_layers, or missing PyTorch entry points.
  • Docs: a section in autoep.rst after the fused RoPE one.

Numerics

  • The forward output equals Hugging Face's SDPA call element for element, and so do dK and dV.
  • cuDNN accumulates dQ non-deterministically: PyTorch does not select cuDNN SDPA under torch.use_deterministic_algorithms(True), and six autograd runs on the same inputs gave six different dQ results.
  • The replayed dQ stays within that spread, one replay run matched an autograd run bitwise, and both have identical error against an FP32 reference.

Tests

tests/unit/runtime/activation_checkpointing/test_attention_stash.py:

  • CPU, on a tiny Hugging Face Mixtral, with the cuDNN calls swapped for the CPU flash-attention kernels:
    • loss and gradients match the normal recompute for 1 and all 3 layers stashed, with nothing left kept;
    • two forwards, then their backwards in reverse order: each recompute takes its own forward's output;
    • no-grad, eval and dropped forwards retain nothing;
    • every fail-fast path above.
    • Three mutations are each caught by a test: a recompute takes any kept output instead of its own; dropped forwards are not released; training without a checkpoint does not fail.
  • GPU (cuDNN):
    • the replayed node against autograd's: bitwise output, dK and dV; dQ as accurate as autograd's against FP32;
    • a training step of an AutoEP engine on a two-layer Mixtral, with and without the stash: equal loss, attention gradients within 1e-2 relative L2.
  • H100 run of the new tests together with the existing activation-checkpointing and AutoEP unit suites: 240 passed, 1 skipped. Both cuDNN tests of the new file ran and passed. The skip is the MiniMax-M3 router-logit test, which needs Transformers >= 5.15.0; the image has 5.12.0.

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:

Stashed layers Run medians Median Peak allocated
0 18,733.9 / 18,739.1 ms 18,736.5 ms 58.69 GB
30 18,202.0 / 18,215.5 ms 18,208.7 ms (-2.82%) 62.77 GB
48 17,911.1 / 17,910.0 ms 17,910.6 ms (-4.41%) 65.23 GB
  • The peak is in the micro-batch forward/backward, not the optimizer step, so every stashed layer adds its 136 MB to it.
  • Losses agree within same-arm repeat noise over 30 steps.
  • Before the timing, a one-H100 gate ran on the real attention shapes. It checked the op-level numerics above, and that a two-layer model is as far from the baseline as the baseline is from its own rerun. It measured 0.93 ms saved per stashed layer and micro-batch.

…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

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant