Skip to content

Keep FlashAttention inside captured CUDA graphs - #28

Open
BruceSheng1202 wants to merge 2 commits into
mainfrom
perf/cuda-graph-flash-attention
Open

BruceSheng1202 wants to merge 2 commits into
mainfrom
perf/cuda-graph-flash-attention

Conversation

@BruceSheng1202

Copy link
Copy Markdown
Collaborator

Summary

transformers counts CUDA stream capture as tracing. While RowGraphs captures 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's is_causal shortcut 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() wraps masking_utils.find_packed_sequence_indices and _ignore_causal_mask_sdpa during 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._bucket sends 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.
  • No option or default changes; docs/DEPLOYMENT.md describes 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-graphs on public JevBench (231 requests x 3), median forward ms:

input tokens (requests) main this PR
up to 512 (513) 16.3 / 22.5 16.3 / 22.5
513-1,024 (57) 41.9 40.7
1,025-1,536 (12) 68.4 66.2
1,537-2,048 (3) 114.6 108.5
whole panel: median / mean / p90 16.43 / 41.79 / 130.14 16.33 / 41.54 / 130.25
--cuda-graph-max-tokens 4096, 2,049-3,072 (84) 138.9 (eager 131.2) 130.2
--cuda-graph-max-tokens 4096, 3,073-4,096 (24) 231.8 (eager 192.8) 208.4

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

  • CPU: under simulated capture, create_causal_mask returns a mask without the context and None inside it (transformers 5.17 and 5.18); rows with a position jump run eagerly.
  • CUDA: capture passes attn_mask=None, is_causal=True to SDPA; existing replay-parity tests pass.
  • pytest tests -m "not server" on the A100 node: the same 30 failures as main in that environment, none new.

🤖 Generated with Claude Code

https://claude.ai/code/session_01EpJAgfodeE1m4Fk2tDiLfB

BruceSheng1202 and others added 2 commits October 5, 2026 02:30
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
@github-actions

github-actions Bot commented Oct 5, 2026 •

Copy link
Copy Markdown

Human Review

  • Applicable CI failed or cannot be verified for this PR revision; see the CI evidence below.

Human Review: see the findings and review evidence.

AI semantic review has not completed.

  • CI ci.yml: incomplete (run) — CI cannot be bound to the current PR/base. Update the branch from main and run CI again.
  • CI pages.yml: incomplete (run) — CI cannot be bound to the current PR/base. Update the branch from main and run CI again.

Reviewed head 3b8fca2008c2342ad14cb725541cfc084d190a32; base 2b9ed4c20f0e3d86db28bf7ddc1c11ac80c6dfe0.
CPU CI does not verify production-model quality, CUDA or multi-GPU performance.

@github-actions github-actions Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Human-Review. See the maintained PR review summary for evidence.

@github-actions github-actions Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Human-Review. See the maintained PR review summary for evidence.

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant