Repository navigation
Keep FlashAttention inside captured CUDA graphs - #28
Open
BruceSheng1202 wants to merge 2 commits into
Open
BruceSheng1202 wants to merge 2 commits into
BruceSheng1202 wants to merge 2 commits into
Conversation
transformers counts CUDA stream capture as tracing. While RowGraphs captured a row it kept a packed-sequence mask (it cannot run its data-dependent single-sequence check) and, before transformers 5.18, also declined SDPA's is_causal shortcut. Captured attention therefore read a materialised causal mask and used the memory-efficient kernel instead of the FlashAttention kernel eager inference picks for the same row. Capture now makes the eager mask decisions for an unpadded single row: the two masking_utils functions are wrapped only for calls made while the current stream is capturing, and are restored afterwards. Rows whose positions do not rise by one run eagerly, because eager inference would mask such a jump as a packed-sequence boundary. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01EpJAgfodeE1m4Fk2tDiLfB
Explain why capture now keeps FlashAttention and record the A100 JevAny-Qwen3.5-4B measurements, including why the 2,048-token default limit stays. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01EpJAgfodeE1m4Fk2tDiLfB
Human Review
Human Review: see the findings and review evidence. AI semantic review has not completed.
Reviewed head |
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.
Summary
transformers counts CUDA stream capture as tracing. While
RowGraphscaptures a row, transformers cannot run its data-dependent check that the position ids form one sequence, so it keeps a packed-sequence mask; before 5.18 it also declines SDPA'sis_causalshortcut for a missing padding mask. Every captured attention layer therefore reads a materialised causal mask and SDPA runs its memory-efficient kernel, while eager inference of the same row uses FlashAttention._eager_mask_decisions()wrapsmasking_utils.find_packed_sequence_indicesand_ignore_causal_mask_sdpaduring capture so they make the eager decisions for an unpadded single row. The wrappers act only while the calling thread's stream is capturing, and the originals are restored on exit.RowGraphs._bucketsends rows whose positions do not rise by one to the eager path. Row mode does not build such rows today, but eager inference would mask a position jump as a packed-sequence boundary, which a graph captured without that mask would not reproduce.docs/DEPLOYMENT.mddescribes the behaviour.Measurements
A100-SXM4-40GB, one node: JevAny-Qwen3.5-4B-LoRA, BF16 merged LoRA, SDPA, FLA 0.5.2, causal-conv1d 1.7.0, torch 2.8.0, transformers 5.17.0. Latency from
python -m scripts.benchmark_latency --records 0 --repeats 3 --serving-kernels --cuda-graphson public JevBench (231 requests x 3), median forward ms:--cuda-graph-max-tokens 4096, 2,049-3,072 (84)--cuda-graph-max-tokens 4096, 3,073-4,096 (24)The 2,048-token default stays right for this model on its own. With the fused Qwen3.5 kernels (
perf/qwen35-fused-kernels) as well, a 4,096-token limit beats eager inference (2,049-3,072 tokens: 102.4 vs 113.5 ms) and lowers JevBench p90 from 112.7 to 102.4 ms.Quality with graphs and serving kernels (
jevany.benchmark,LocalPredictor(exact_kernels=False)): Transfer-v9 dev 823 -> 824 / 1,046 correct, NLL 0.5869 -> 0.5870, ECE 0.035 -> 0.032; JevBench 180 -> 180 / 231. 3 of 1,264 Transfer and 2 of 231 JevBench argmax decisions changed; largest probability change 0.029.Tests
create_causal_maskreturns a mask without the context andNoneinside it (transformers 5.17 and 5.18); rows with a position jump run eagerly.attn_mask=None, is_causal=Trueto SDPA; existing replay-parity tests pass.pytest tests -m "not server"on the A100 node: the same 30 failures asmainin that environment, none new.🤖 Generated with Claude Code
https://claude.ai/code/session_01EpJAgfodeE1m4Fk2tDiLfB