Skip to content

Refactor Deepseek v4 to use Attention Gym's Compressed Sparse Attention - #4452

Closed
floatingtrees wants to merge 12 commits into
pytorch:mainfrom
floatingtrees:deepseek-v4-performance
Closed

floatingtrees wants to merge 12 commits into
pytorch:mainfrom
floatingtrees:deepseek-v4-performance

Conversation

@floatingtrees

@floatingtrees floatingtrees commented Sep 3, 2026 •

Copy link
Copy Markdown
Contributor

Compressed Sparse Attention (CSA) is a data dependent sparse attention mechanism. It does not work well with Flex Attention; Flex Attention uses a block sparsity mask to determine which tiles to compute, and it recomputes the block mask every time the mask changes. Since CSA is data dependent, this means Flex Attention recomputes the block mask on almost every forward pass.

For the debug model, the Attention Gym operator is around 3x faster at 1024 sequence length and 11x faster at 2048 sequence length. For an end to end run of the debug model, model performance rises from 36 TFLOPS/GPU to 105 TFLOPS/GPU.

Tests were performed on 2xGB200.

This PR currently contains:

  1. A refactor to torchtitan/models/deepseek_v4/attention.py that calls the Attention Gym operator instead of Flex Attention.
  2. Correctness tests.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Meta Open Source bot. label Sep 3, 2026
@floatingtrees
floatingtrees marked this pull request as draft September 3, 2026 21:15
@tianyu-l
tianyu-l requested a review from drisspg September 4, 2026 02:10
)


class SlidingWindowAttention(DSV4FlexAttention):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Is this the commonly seen sliding window attention? If so, why do we need to build ad hoc mask per forward, instead of only once per iteration?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Thanks for pointing this out, I'll add a cache for the mask.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

This is still inefficient, right? Is there a plan for it?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Yes, there's a PR in Attention Gym that adds a fused kernel for the indexer here, but it hasn't been merged yet. The current implementation is a workaround for now.

@@ -322,6 +309,20 @@


class CompressedSparseAttention(DSV4FlexAttention):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

could you show some perf comparison?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Hi Tianyu, what kind of perf comparison were you thinking of? In terms of end to end performance, the old model logs looked like this:

step:  6  loss:  4.07696  grad_norm:  2.3680  memory:  6.39GiB(3.47%)  tps: 11,274  tflops: 36.96  mfu: 1.48%
step:  7  loss:  3.94543  grad_norm:  1.7404  memory:  6.39GiB(3.47%)  tps: 11,293  tflops: 37.03  mfu: 1.48%
step:  8  loss:  3.84256  grad_norm:  1.9111  memory:  6.39GiB(3.47%)  tps: 11,209  tflops: 36.75  mfu: 1.47%

After after the changes, it's roughly 3x faster:

step:  6  loss:  4.07704  grad_norm:  2.3684  memory:  5.94GiB(3.22%)  tps: 33,094  tflops: 108.50  mfu: 4.34%
step:  7  loss:  3.94542  grad_norm:  1.7406  memory:  5.94GiB(3.22%)  tps: 32,667  tflops: 107.10  mfu: 4.28%
step:  8  loss:  3.84257  grad_norm:  1.9117  memory:  5.94GiB(3.22%)  tps: 32,430  tflops: 106.32  mfu: 4.25%

@tianyu-l tianyu-l left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

one remaining comment

selected_indices.append(sink_indices)
selected_indices = torch.cat(selected_indices, dim=-1)
block_mask = attention_masks
if block_mask is None:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Should never be None after this change?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

+1

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Thanks for catching that, I fixed it.


seqlen = q.size(0)

with spmd.no_typecheck():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

TODO: We should open a PR in Attention Gym to add a keyword-only scale argument to selected_attention,

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

That makes sense, I added a pr for this: meta-pytorch/attention-gym#540

"CompressedSparseAttention does not accept attention_masks; "
"top-k selection is computed internally."
)
if attn_sink is None:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

ill take another look at the fa4 change

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Thanks, that would be helpful.

Comment on lines +415 to 425
q_bhtk,
swa_k_b1tk,
cmp_k_b1tk,
cmp_topk,
attention_sink=attn_sink,
doc_ids=None,
sliding_window_size=self.window_size,
backend=backend,
)
return out_bhtk.squeeze(0).transpose(0, 1)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

I thought our Triton implementation didn't support head_dim=512 for training, and the FA4 path didn't have sinks enabled yet? How do we handle the production configs here?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Currently, production configs aren't handled; I was thinking of waiting until the flash attention change got merged.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Thanks for fixing up the flash attention change! I updated the code to route head_dim=512 to the cute backend.

Comment thread torchtitan/models/deepseek_v4/model.py
backend = "triton" if q.device.type == "cuda" else "eager"
out_bhtk = selected_attention(
q_bhtk,
swa_k_b1tk,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Does this transpose result in a query copy inside Attention Gym's query.contiguous()? we shoudl do a path to make sure we swap contigousu to the min alignment needed

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

I think this is only a problem for the triton backend; FA4 only cares about the last dimension being contiguous and aligned. Do you think this is a major issue if it only affects the debug models?

cached = self.mask_cache.get(cache_key)
if cached is None:
cached = inner.build_block_mask(seqlen=seqlen, device=device)
self.mask_cache[cache_key] = cached

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

can we add cached-versus-uncached output and gradient comparisons for SWA/HCA, including cache reuse and a sequence-length change id ont think we are going through cache path in tests right?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Ok, I'll add this. We aren't going through the cache path for normal tests.

@floatingtrees
floatingtrees marked this pull request as ready for review September 10, 2026 22:32
@floatingtrees

Copy link
Copy Markdown
Contributor Author

I got it to use the newest version of attention gym, so the integration tests should pass now. I also fixed the lints that were failing due to some attention mask changes. Some lints still fail, but those are unrelated to this PR; they also fail on commit 1584a18 on main.

@tianyu-l
tianyu-l requested a review from drisspg September 11, 2026 21:20
@drisspg

drisspg commented Sep 11, 2026

Copy link
Copy Markdown
Contributor

doing 1 more pass; lets way till we publish 1 more relase for kda-cp branch before integrating this

Comment on lines +400 to +403
backend = (
"cute"
if head_dim == 512
and torch.cuda.get_device_capability(q.device) == (10, 0)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

We need Dao-AILab/flash-attention#2882 to land for V4 Flash's 64 query heads and TP-sharded head counts. We should also cut a new Attention Gym release that relaxes the CuTe backend's hard-coded h == 128 check and includes scale, then pin that release consistently in pyproject.toml, requirements.txt, and .ci/docker/requirements.txt. Landing FA4 alone isn't enough: both the current Gym pin and 499404b still reject non-128 head counts before calling FA4. Let's also ensure the installed flash-attn-4 version includes #2882.

sdmyzlp added a commit to sdmyzlp/torchtitan that referenced this pull request Sep 14, 2026
Replace the hand-written chunked sparse core with Attention Gym's
selected_attention, pinned to the commit the torchtitan DeepSeek-V4 PR (pytorch#4452)
uses (b16d6d3). The operator fuses the sliding window, the selected compressed
entries and the attention sink into one softmax, and returns the per-head LSE.

Packed documents are isolated with doc_ids alone, the only varlen metadata:
- selected_attention applies doc_ids[q] == doc_ids[k] to its window branch;
- the indexer applies the same equality one axis over, entry j covering tokens
  [j*ratio, (j+1)*ratio), so entry j belongs to doc_ids[j*ratio] and is
  selectable by t iff doc_ids[t] == doc_ids[j*ratio] and j < (t+1)//ratio
  (the operator leaves the sparse branch to the caller);
- the distillation loss needs no metadata: the LSE it consumes is already
  document-masked and its support is the already-filtered top-k.

The teacher is now exp(logit - lse) summed over heads and L1-normalized, which
is algebraically the DeepSeek-V4/Megatron teacher (NVIDIA/Megatron-LM#5776) with
the LSE supplied by the kernel instead of a second softmax pass; the fp32
underflow shift is preserved.

Both the indexer and the sparse core are single-pass now; the query chunking was
a kernel concern and is gone. Packing alignment is a new prerequisite, so
PackingAlignedDatasetConfig pads every packed segment to a multiple of the
compressor ratio before the packer runs (the packer's own cuts are multiples).

Tests: packed-document isolation (perturbing one segment leaves the other
bit-identical), no cross-document indexer selection, the LSE teacher against a
float64 oracle, and the packing-alignment helper.
sdmyzlp added a commit to sdmyzlp/torchtitan that referenced this pull request Sep 14, 2026
The DeepSeek V4.1-Flash text backbone, written from the paper and the reference
inference implementation but following torchtitan's model conventions: no
FlexAttention, no shared-state object, and nothing reused from the deepseek_v4
implementation.

- mHC: Single-Pass mHC. HcPre predicts the mixing coefficients and collapses
  the branches with the coefficients of the previous sublayer; HcPost applies
  the residual update. No learned output head: the stack collapses with the
  last block's coefficients.
- CSA2 attention: Attention Gym's selected_attention fuses the sliding window,
  the selected compressed entries and the per-head attention sink into one
  masked softmax and returns the per-head LSE. Compressed main KV is produced
  by kv_source_layers and reused by the rest.
- Lightning indexer: single-pass, graph-free top-k selection over all
  candidates, then a second differentiable pass over the selected entries only,
  so the [T, heads, entries] score tensor never enters the autograd graph.
- Indexer distillation (IndexerKLLoss): the teacher is exp(logit - lse) summed
  over heads and L1-normalized, i.e. the DeepSeek-V4/Megatron teacher
  (NVIDIA/Megatron-LM#5776) with the LSE supplied by the kernel instead of a
  second softmax pass. The sliding window and the sink therefore enter each
  head's denominator, so a head sitting on its window or the sink cannot
  outvote a head that actually uses the compressed entries. The teacher stays
  detached and the attention output is unchanged.
- Packed documents are isolated with doc_ids alone, the only varlen metadata:
  selected_attention applies doc_ids[q] == doc_ids[k] to its window branch, and
  the indexer applies the same equality one axis over, entry j covering tokens
  [j*ratio, (j+1)*ratio), so entry j belongs to doc_ids[j*ratio] and is
  selectable by t iff doc_ids[t] == doc_ids[j*ratio] and j < (t+1)//ratio. The
  distillation loss needs no metadata: its LSE is already document-masked and
  its support is the already-filtered top-k.
- The metadata travels in DeepSeekV41Metadata, built in get_attention_masks;
  preprocess_inputs is overridden because Decoder.preprocess_inputs only builds
  masks for the FlexAttention and VarlenAttention inner backends. forward
  requires positions and the metadata, both of which the framework always
  supplies.
- Hierarchical sparse indexer behind a use_candidates switch, off by default:
  the candidate pool boundary receives no gradient, so it stays a training and
  inference consistency device rather than a learned component.
- Cross-layer tensors (compressed KV, index keys, top-k, student logits,
  candidate pool) are threaded as individual layer inputs.
- Configs: the deepseek_v4_1 debugmodel with use_candidates on and off. Packed
  segments are padded to a multiple of the largest compress_ratio, so the
  compressor's pooling and the entry->document mapping are segment-exact.
- The expert FFN uses the 10.0 SwiGLU branch clamp of the report.

Pins attention-gym to the commit torchtitan's DeepSeek-V4 PR (pytorch#4452) uses
(b16d6d3), registers the flavor in the CLI-freeze guard and the debug-config
defaults, and lets codespell ignore the ``thn`` einsum subscript.

Not included yet: Engram, the vision encoder, DSpark/MTP, sharding.py,
state_dict_adapter.py, TP/PP/CP support, and indexer quantization.
@tianyu-l

Copy link
Copy Markdown
Contributor

what's the status of this PR

@drisspg

drisspg commented Sep 21, 2026 •

Copy link
Copy Markdown
Contributor

I'll take over @tianyu-l i was waiting to land my fix for smaller headda in fa4 but I'll do some routing to triton before that lands

@yinli-systems

yinli-systems commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

Hi — I took a pass at the current integration gap without opening a competing PR.

I ported this PR's intent onto current main (0e4a2803b) here: yinli-systems/torchtitan@469f9bb0b (the implementation is in its parent, 8214c85b2; the head updates the tests for current main).

The port addresses the drift since the last update:

  • uses the attn-gym==0.0.13 gather_attn API already pinned on main, including scale;
  • leaves fused backend selection to Attention Gym's current CuTe/Triton routing, so it no longer depends on the closed flash-attention#2882 path for non-128 or TP-sharded head counts;
  • restores model-level SWA/HCA mask caching on the current DSV4FlexInnerAttention / model-config structure;
  • constructs static SWA/HCA block metadata analytically, rather than materializing the cached implementation's full [L, KV] count and Boolean tensors;
  • retains an uncached fallback for direct inner-attention callers.

Coverage includes cache reuse and sequence-length invalidation, incompatible-layer validation before deduplication, cached-vs-fresh output/gradient parity, CSA BlockMask-vs-gather_attn output/gradient parity, asymmetric block sizes, and a 32K assertion that mask_mod captures no token-pair tensor.

Current validation:

  • CPU/reference: 5 passed, plus 8 subtests.
  • Focused formatting/lint/doc hooks pass; targeted Pyrefly 0.45.1 reports 0 errors.
  • A separate RTX 5090 benchmark of the same analytic static-mask builder at 32K reduced construction peak allocation from 8.21 GB to 3.21 MB and median construction time from 19.89 ms to 1.07 ms. These are mask-construction measurements, not end-to-end training speedups.
  • Current-main CUDA suite on an RTX 5090: 5 passed, plus 10 subtests, in 126.87 seconds. This run used a dependency-only diagnostic workaround described below; it is not a claim that the stock dependency passes.

The CUDA validation also exposed an independent compatibility blocker in the pinned attn-gym==0.0.13: its Triton backward kernel uses tl.cat(..., dim=0), while the installed Triton 3.6 API rejects dim and only supports one-dimensional tl.cat. The stock dependency therefore compiled the forward path but failed the backward test on both RTX 4090 and RTX 5090. For diagnosis only, I kept the original dependency untouched and used a local copy where the concatenated matrix product is expressed as the sum of its two component matrix products (probability * grad_output and scaled grad_score * query). With that algebra-preserving workaround, the complete focused CUDA suite above passed. I did not vendor or push this dependency change.

So the TorchTitan port and its cache/forward/backward parity tests are now validated under the intended gradient algebra, while an upstream Attention Gym/Triton compatibility fix is still required for the stock pinned environment.

If this direction matches the takeover plan, please feel free to cherry-pick the commits or reuse the tests/analytic builder while reshaping the final backend routing.

AI disclosure: the implementation and wording were AI-assisted; I reviewed the code, ran the checks above, and kept only claims supported by the recorded evidence.

@drisspg

drisspg commented Sep 24, 2026

Copy link
Copy Markdown
Contributor

Closing for : #4855

@yinli-systems thanks for the comment we use nightly for alto fo features which is 3.8+ where tl.cat is supported

@drisspg drisspg closed this Sep 24, 2026
drisspg added a commit that referenced this pull request Sep 29, 2026
## Human notes

Step 1 we port over CSA to use the gather_attn impl; supports triton and
routes to FA4 cute on B2/300 machines.

full e2e speedups in top of stack 

## Agent Notes

CompressedSparseAttention (compress_ratio == 4) no longer concatenates
[sliding-window KV, compressed KV, sink] and builds a FlexAttention
BlockMask
from the indexer top-k. It calls attn_gym.sparse.selected_attention,
which
fuses the sliding-window branch, the per-query top-k branch over the
compressed KV, and the sink softmax denominator. The fused impl picks
CuTe
(FA4) for supported eager SM100/SM103 shapes and Triton otherwise; CPU
uses
the reference impl. The CSA branch and indexer arguments leave
DSV4FlexInnerAttention._forward_impl, which now serves sliding-window
and HCA
layers only.

The GPU test keeps an inline copy of the previous BlockMask formulation
and
checks outputs and gradients match it at the default and a non-default
softmax_scale. A second test runs DSV4-shaped CSA (64 heads, head_dim
512,
bf16) through the FA4 CuTe kernels against the fp32 reference
implementation
and fails if CuTe is unavailable. The B200 lane runs this file, since
the GPU
unit-test lane (A10G) cannot reach the CuTe path.

Based on #4452 by @floatingtrees.

Co-authored-by: Jonathan Zhou
<96200235+floatingtrees@users.noreply.github.com>

---

<sub>Stack created with <a
href="https://github.com/github/gh-stack">GitHub Stacks CLI</a> • <a
href="https://gh.io/stacks-feedback">Give Feedback 💬</a></sub>

Co-authored-by: Jonathan Zhou <96200235+floatingtrees@users.noreply.github.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

CLA Signed This label is managed by the Meta Open Source bot.

Projects

Status: Done

Development

Successfully merging this pull request may close these issues.

5 participants