Repository navigation
Refactor Deepseek v4 to use Attention Gym's Compressed Sparse Attention - #4452
floatingtrees wants to merge 12 commits into
Conversation
| ) | ||
|
|
||
|
|
||
| class SlidingWindowAttention(DSV4FlexAttention): |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
Thanks for pointing this out, I'll add a cache for the mask.
There was a problem hiding this comment.
This is still inefficient, right? Is there a plan for it?
There was a problem hiding this comment.
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): | |||
There was a problem hiding this comment.
could you show some perf comparison?
There was a problem hiding this comment.
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%
| selected_indices.append(sink_indices) | ||
| selected_indices = torch.cat(selected_indices, dim=-1) | ||
| block_mask = attention_masks | ||
| if block_mask is None: |
There was a problem hiding this comment.
Should never be None after this change?
There was a problem hiding this comment.
Thanks for catching that, I fixed it.
|
|
||
| seqlen = q.size(0) | ||
|
|
||
| with spmd.no_typecheck(): |
There was a problem hiding this comment.
TODO: We should open a PR in Attention Gym to add a keyword-only scale argument to selected_attention,
There was a problem hiding this comment.
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: |
There was a problem hiding this comment.
ill take another look at the fa4 change
There was a problem hiding this comment.
Thanks, that would be helpful.
| 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) | ||
|
|
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
Currently, production configs aren't handled; I was thinking of waiting until the flash attention change got merged.
There was a problem hiding this comment.
Thanks for fixing up the flash attention change! I updated the code to route head_dim=512 to the cute backend.
| backend = "triton" if q.device.type == "cuda" else "eager" | ||
| out_bhtk = selected_attention( | ||
| q_bhtk, | ||
| swa_k_b1tk, |
There was a problem hiding this comment.
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
There was a problem hiding this comment.
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 |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
Ok, I'll add this. We aren't going through the cache path for normal tests.
|
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. |
|
doing 1 more pass; lets way till we publish 1 more relase for kda-cp branch before integrating this |
| backend = ( | ||
| "cute" | ||
| if head_dim == 512 | ||
| and torch.cuda.get_device_capability(q.device) == (10, 0) |
There was a problem hiding this comment.
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.
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.
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.
|
what's the status of this PR |
|
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 |
|
Hi — I took a pass at the current integration gap without opening a competing PR. I ported this PR's intent onto current The port addresses the drift since the last update:
Coverage includes cache reuse and sequence-length invalidation, incompatible-layer validation before deduplication, cached-vs-fresh output/gradient parity, CSA Current validation:
The CUDA validation also exposed an independent compatibility blocker in the pinned 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. |
|
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 |
## 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>
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:
torchtitan/models/deepseek_v4/attention.pythat calls the Attention Gym operator instead of Flex Attention.