Skip to content

[DCP] Fix sparse indexer local context metadata for CUDA graphs - #2102

Merged
valarLip merged 3 commits into
mainfrom
yuhua/dcp-fusion
Sep 1, 2026
Merged

valarLip merged 3 commits into
mainfrom
yuhua/dcp-fusion

Conversation

@zhuyuhua-v

@zhuyuhua-v zhuyuhua-v commented Sep 1, 2026 •

Copy link
Copy Markdown
Collaborator

Motivation

  • Attach precomputed DCP local context lengths to CUDA Graph capture metadata.
  • Share the same metadata wiring across eager, CUDA Graph, and TBO paths.
  • Avoid capturing redundant elementwise kernels in every full-index layer.
  • Restrict the change to sparse DSA with DCP enabled.

Root Cause

CUDA Graph capture metadata omitted dcp_local_context_lens, causing the sparse indexer to capture the eager elementwise fallback. Runtime metadata updates could not change the already-captured graph.

Compatibility

Pure TP, dense MLA, and non-sparse models remain unchanged. LMCache compatibility is preserved because local context lengths are derived per decode step and are not persisted as KV-cache state.

performance

before
image

after
image
7 elementwise kernel are removed.

Test Plan

gsm8k 5shot & 20shot for DCP4+TP4
server

export AITER_QUICK_REDUCE_QUANTIZATION=INT4
export AITER_USE_FLYDSL_MOE_SORTING=1
TP=${TP:-4}
CONC=${CONC:-64}
export ATOM_ENABLE_DETAILED_ANNOTATION=1
export PYTHONPATH=/home/qichu_qle/yuhua/work/pr/ATOM
export CUDA_VISIBLE_DEVICES=4,5,6,7
 
python -m atom.entrypoints.openai_server \
  --model "$model_path" \
  --server-port 8013 \
  --kv_cache_dtype fp8 \
  --online_quant_config '{"global_quant_config":"ptpc_fp8","exclude_layer":["lm_head","model.embed_tokens","*.mlp.gate","*expert*"]}' \
  -tp $TP \
  -dcp 4 \
  --max-num-seqs $((CONC * 2)) \
  --max-num-batched-tokens 16384 \
  2>&1 | tee "./server-glm5.2-fp4.log"

acc test 20 shot

python3 -m lm_eval --model local-chat-completions --apply_chat_template --tasks gsm8k --output_path ./eval_out-tta1J8 --log_samples --num_fewshot 20 --model_args 'model=/shared/data/amd_int/models/GLM-5.2-MXFP4,base_url=http://0.0.0.0:8013/v1/chat/completions,api_key=EMPTY,eos_string=</s>,max_retries=5,num_concurrent=64,timeout=1800,tokenized_requests=False,max_length=1048576' --gen_kwargs max_tokens=16384,temperature=0,top_p=1

acc test 5 shot

python3 -m lm_eval --model local-chat-completions --apply_chat_template --tasks gsm8k --output_path ./eval_out-tta1J8 --log_samples --num_fewshot 5 --model_args 'model=/shared/data/amd_int/models/GLM-5.2-MXFP4,base_url=http://0.0.0.0:8013/v1/chat/completions,api_key=EMPTY,eos_string=</s>,max_retries=5,num_concurrent=64,timeout=1800,tokenized_requests=False,max_length=1048576' --gen_kwargs max_tokens=16384,temperature=0,top_p=1

Test Result

20 shot

local-chat-completions ({'model': '/shared/data/amd_int/models/GLM-5.2-MXFP4', 'base_url': 'http://0.0.0.0:8013/v1/chat/completions', 'api_key': 'EMPTY', 'eos_string': '</s>', 'max_retries': 5, 'num_concurrent': 64, 'timeout': 1800, 'tokenized_requests': False, 'max_length': 1048576}), gen_kwargs: ({'max_tokens': 16384, 'temperature': 0, 'top_p': 1}), limit: None, num_fewshot: 20, batch_size: 1
|Tasks|Version|     Filter     |n-shot|  Metric   |   |Value |   |Stderr|
|-----|------:|----------------|-----:|-----------|---|-----:|---|-----:|
|gsm8k|      3|flexible-extract|    20|exact_match|↑  |0.9636|±  |0.0052|
|     |       |strict-match    |    20|exact_match|↑  |0.9644|±  |0.0051|

5 shot

local-chat-completions ({'model': '/shared/data/amd_int/models/GLM-5.2-MXFP4', 'base_url': 'http://0.0.0.0:8013/v1/chat/completions', 'api_key': 'EMPTY', 'eos_string': '</s>', 'max_retries': 5, 'num_concurrent': 64, 'timeout': 1800, 'tokenized_requests': False, 'max_length': 1048576}), gen_kwargs: ({'max_tokens': 16384, 'temperature': 0, 'top_p': 1}), limit: None, num_fewshot: 5, batch_size: 1
|Tasks|Version|     Filter     |n-shot|  Metric   |   |Value |   |Stderr|
|-----|------:|----------------|-----:|-----------|---|-----:|---|-----:|
|gsm8k|      3|flexible-extract|     5|exact_match|↑  |0.9675|±  |0.0056|
|     |       |strict-match    |     5|exact_match|↑  |0.9683|±  |0.0055|

Submission Checklist

Signed-off-by: zhuyuhua-v <yuhzhu@amd.com>
@github-actions

github-actions Bot commented Sep 1, 2026

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every eligible PR before approval:

  • ✅ Pre Checkin: Black, Ruff, catalog schema validation, non-GPU unit tests

Heavy model tests:

  • ✅ Run after the PR is approved and Pre Checkin passes
  • ✅ Run immediately when an approval review is submitted
  • ✅ Can be requested before approval with labels
Label Tests
ci:full Run all heavy PR model tests: native ATOM, vLLM, and SGLang
ci:atom Run native ATOM model accuracy tests
ci:vllm Run ATOM vLLM OOT model accuracy tests
ci:sglang Run ATOM SGLang model accuracy tests

Heavy jobs are skipped when the PR is not approved and no matching ci:* label is present.
Add labels via the sidebar or gh pr edit 2102 --add-label <label>

Signed-off-by: zhuyuhua-v <yuhzhu@amd.com>
@zhuyuhua-v
zhuyuhua-v marked this pull request as ready for review September 1, 2026 06:34
Copilot AI lite review requested due to automatic review settings September 1, 2026 06:34

Copilot AI 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.

Pull request overview

This PR fixes CUDA Graph capture behavior for DCP-enabled sparse attention by ensuring the precomputed per-request DCP local context lengths are present in capture-time metadata, so the sparse indexer doesn’t capture the elementwise fallback path. It also unifies the wiring so eager, CUDA Graph, and TBO (ubatch) metadata construction share the same attachment logic.

Changes:

  • Add a shared helper (_attach_dcp_local_context_lens) to attach DCP local context lengths consistently across decode metadata paths.
  • Populate/copy dcp_local_context_lens into CUDA Graph capture metadata for sparse + DCP>1 so capture records the correct branch.
  • Extend tests to validate capture-time publication and ubatch-prefix attachment behavior.

Reviewed changes

Copilot reviewed 2 out of 2 changed files in this pull request and generated no comments.

File Description
tests/test_mla_index_cache.py Adds unit tests validating CUDA Graph capture metadata includes dcp_local_context_lens and that ubatch-prefixed attachment works.
atom/model_ops/attentions/aiter_mla.py Wires dcp_local_context_lens into CUDA Graph capture + TBO ubatch paths via a shared helper, and allocates per-ubatch buffer when needed.

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

yitingw1
yitingw1 previously approved these changes Sep 1, 2026
@valarLip

valarLip commented Sep 1, 2026

Copy link
Copy Markdown
Collaborator

Reviewed at 54115f034. The core mechanism is sound and correctly wired: pinning the fast branch at capture instead of letting the graph record the elementwise fallback is the right fix, and the guards, row counts and buffer sizes on the eager / capture / TBO paths all check out (copy_to_gpu is in-place, so the captured views stay valid). No crash-class defect. Two things worth acting on before merge, then a list.


1. The win does not land on the default --level 3 piecewise path

The consumer accepts the published buffer only on exact equality:

# atom/model_ops/dcp_ops.py:960-961
local_ctx = getattr(attn_metadata, "dcp_local_context_lens", None)
if local_ctx is not None and local_ctx.shape[0] == num_rows:
    return local_ctx

num_rows is the caller's num_decode_tokens, which is scheduled-width:

# atom/models/deepseek_v2.py:1444-1448
num_decode_tokens = (
    context.scheduled_bs * attn_metadata.max_seqlen_q
    if not context.is_prefill
    else 0
)

while the publisher is running-width:

# atom/model_ops/attentions/aiter_mla.py:2034
self._attach_dcp_local_context_lens(attn_metadata, running_bs, copy_to_gpu=True)

running_bs is the capture-ladder rung and is set independently of use_cudagraph (forward_context.py:301-305). Under PIECEWISE — the default (arg_utils.py:249, config.py:1898-1904) — the indexer runs eagerly every step with the real Context, so any batch size off the ladder fails the equality and the fallback's elementwise kernels run on all 21 full-index layers. Only --cudagraph-mode FULL replay and TBO ubatches actually get the win.

Worth confirming which mode the reported benchmark used.

The fix: read the width the engine already stored, rather than re-deriving it

Context already carries this number, and its own docstring says why the product form is the wrong way to get it:

# atom/utils/forward_context.py:485-490
# The step's DP-unified padded shape. `running_bs` counts SEQUENCES (graph
# identity, the draft's pad width), `running_tokens` the hidden_states rows
# MoE pads to. Both stored, because the ratio is not always max_seqlen_q --
# a DSpark ragged step runs a packed width no rectangular bs*q recovers.
running_bs: int = 0
running_tokens: int = 0

So:

num_decode_tokens = context.running_tokens if not context.is_prefill else 0

Three reasons this is better than either widening the guard or spelling running_bs * max_seqlen_q:

  1. It is exactly the height hidden_states has. The piecewise branch calls the model with self.forward_vars["input_ids"].gpu[:running_tokens] (model_runner.py:3030), and forward_context.py:426-427 states it directly — "the cudagraph branch re-slices the buffer to running_tokens itself". So num_tokens = hidden_states.shape[0] at deepseek_v2.py:1550 is running_tokens; using it for num_decode_tokens makes q_fp8[:num_decode_tokens] cover the whole real buffer and leaves q_fp8[num_decode_tokens:num_tokens] correctly empty on a pure-decode step.

  2. The guard then passes exactly, with no change to dcp_ops.py. dcp_ops.py:1050 asserts attn_metadata.max_seqlen_q == 1 on this path, so running_tokens == running_bs * 1 == running_bs — and the published buffer is running_bs rows. Both sides are the same number.

  3. It removes a rectangular assumption rather than replacing it. Both scheduled_bs * max_seqlen_q and running_bs * max_seqlen_q re-derive a quantity the engine already computed, using a ratio the Context docstring above explicitly warns does not always hold.

On the cost side there is nothing new to justify: decode is padded inside the graph already, so this only makes eager decode behave the way graph decode always has. The extra rows are inert either way — slot_mapping[:running_tokens] = -1 (aiter_mla.py:1829) so nothing is written to cache, and context_lens[scheduled_bs:running_bs] = 0 (:1834).

(If you would rather not touch num_decode_tokens, relaxing dcp_ops.py:961 to >= num_rows plus local_ctx[:num_rows] also works — the producer fills real values first and zero-pads the tail, so the prefix is exactly the real rows. But that leaves the scheduled/running mismatch in place for the next buffer that hits it.)

2. The publish guard was narrowed; the host-side producer was not

# aiter_mla.py:1835   producer
if self.dcp_world_size > 1:                       # no is_sparse term

# aiter_mla.py:1760   consumer-side helper
if not self.is_sparse or self.dcp_world_size <= 1:
    attn_metadata.dcp_local_context_lens = None
    return

So on dense-MLA DCP the two numpy slice-assigns at :1847-1848 still run every decode step for a buffer that is then published as None, plus a permanently resident pinned+device CpuGpuBuffer. And the comment above them still reads "Publish it: the sparse indexer used to re-derive this on device with 8 elementwise kernels per full-index layer" — describing a publish that no longer happens on that configuration. A reader on a dense-DCP config concludes the attribute is set and writes a consumer that dereferences None.

Same axis one file up: the declaration comment at :332-337 still gives DCP as the only publish condition, while the new per-ubatch twin at :578 documents the correct is_sparse and dcp > 1 gate — two declarations of one quantity now describing two different contracts. The buffer at :337 is also allocated unconditionally, for every model including dcp=1, which is a third predicate for the same object.


Smaller items

The predicate is spelled at five places that must agree, and each failure mode is different and silent. Sites: :581 (ubatch allocation), :1762 (helper guard), :2194 (vars_used append), :2308 (capture fill), plus :2135+:2149 as a nested pair. Missing the allocation gives a KeyError in the helper (it indexes forward_vars with [], not .get); missing the helper silently falls back to the kernels this PR exists to remove; missing the vars_used append publishes a GPU view of a buffer that was never copied — stale data, correct shape, wrong answers. One cached attribute in __init__ (e.g. self._publishes_dcp_local_lens) would name the decision once.

Declare the field instead of setting it dynamically. dcp_local_context_lens is neither a dataclass field nor an __init__ parameter of AttentionMetaData, so fields(self) / asdict_zerocopy (forward_context.py:668) cannot see it and the sole consumer must use getattr(..., None). AttentionMetaData_DSV4 declares dspark_ragged_lens_gpu: torch.Tensor | None = None as a real field (deepseek_v4_attn.py:214-215), and the sibling g_kv_indptr gets a hard assert at its consumer rather than a silent fallback. One line turns "a new backend forgot to publish" from an unmeasurable perf cliff into something a reader can see — and it deletes the helper's dead = None assignment.

The per-ubatch buffer and its copies are avoidable. Unlike kv_indptr / g_kv_indptr / sparse_kv_indptr, dcp_local_context_lens needs no rebasing, so build_ubatch_metadata could hand out var["dcp_local_context_lens"].gpu[req_start : req_start + running_bs]. prepare_decode already fills [:scheduled_bs] real and [scheduled_bs:running_bs] = 0, so ub0's [0:half) and ub1's [half:bs) land on exactly the values the copies reproduce, and the offset is fixed per captured bs so the view is graph-pointer-stable. As written, a change whose purpose is removing per-step work adds 2 pinned buffers, 4 numpy slice-assigns and 2 cudaMemcpyAsync per TBO decode step to reproduce data already resident on device at a known offset.

One operation split across 34 lines, with a parameter that is derivable. The guard+fill+copy at :2308-2311 and the attach at :2342 repeat the same predicate and must move together, and neither references the other — hoisting the attach into the AttentionMetaData constructor kwargs (next to g_kv_indptr at :2339, where it naturally belongs) silently strands the copy. Collapsing to self._attach_dcp_local_context_lens(attn_matadata, bs, copy_to_gpu=True) leaves one guard and one copy; and then copy_to_gpu=True holds exactly when prefix == "" (:2034, :2342) and False exactly for prefix=p (:2430, already copied via vars_used), so the parameter can be deleted.

Three comments that no longer match the code.

  • The helper's docstring (:1750) says it covers "every decode metadata path" and exists "to prevent any path from silently recording the elementwise fallback". Neither holds: prepare_mtp_decode (:821) is a fourth decode path that does not route through it; the helper is opt-in at three call sites, so a fifth path simply would not call it; and its negative arm assigns None, i.e. manufactures the fallback state it claims to prevent.
  • :2309 says "Capture uses one synthetic local KV token per request", implying the value is load-bearing. The branch it influences is a presence/shape test, not a value test (dcp_ops.py:960-961), and no launch grid on this path derives from local_ctx's values. Attaching the buffer is the real fix; np[:bs] = 1 only changes what the throwaway warmup forward reads. As written the comment invites a future reader to replicate a value that does nothing — and it already did: three review angles independently flagged the ubatch path for "missing" the same seed, which the ubatch zero-length convention does not need.
  • Three sites (dcp_ops.py:956, aiter_mla.py:336, :1846) say "8 elementwise kernels" while the PR description says 7. Counting the fallback body at dcp_ops.py:963-968: floordiv, mul, mul, sub, sub, clamp, add = 7; the trailing .to(torch.int32) is a no-op because context_lens is already int32. Worth picking one number — a reader who budgets 8 and measures 7 goes looking for a missing kernel, or makes context_lens int64 somewhere and silently turns the .to() into a real eighth that nobody re-counts.

Tests: none of the four run in CI, and they would not catch a revert of the TBO half.

.github/scripts/run_unit_tests.sh states its scope as "pure-Python unit tests that run on a plain runner (CPU torch, no GPU, no aiter/MoRIIO native libs)". aiter_mla.py:11 imports aiter at module scope, and this PR's new from atom.model_ops.dcp_ops import dcp_local_context_lens adds a top-level import triton (dcp_ops.py:18-19) — widening the skip predicate, so even a CI image that gains aiter keeps skipping. The tested logic (_attach_dcp_local_context_lens' branching, the capture fill) is pure Python + numpy and needs neither.

And on a revert run: deleting the PR's actual TBO code — the _prepare_ubatch_decode slice-copy at :2151-2154 and the vars_used append at :2194 — leaves all four tests green. test_dcp_local_context_attachment_supports_ubatch_prefix pre-populates var["ub0_dcp_local_context_lens"].gpu itself and then asserts the helper slices it, exercising buffer.gpu[:rows], which was never in doubt; no test calls _prepare_ubatch_decode or build_ubatch_metadata. The code that can actually be wrong is the numpy block — the req_start slicing, the ub_real_reqs zero-fill bound, the nesting of the is_sparse guard inside the dcp>1 guard, and the H2D registration.

Two smaller test notes: the pair at :409 / :417 have identical bodies differing only in fixture arguments and should be one @pytest.mark.parametrize over [(1, True), (4, False)]; and assert "dcp_local_context_lens" not in var is false in production, since :337 allocates it unconditionally — so if someone relaxes the helper's guard to just dcp>1, the tests fail with KeyError while production would not: it would publish an allocated-but-never-written all-zeros buffer that the consumer's shape check accepts, and the natural "fix" is to add the key to the fixture and go green.

Also _FakeMetadataBuffer (:356) backs .gpu with a numpy array, which forces block_size=1 and no mtp_k — the one configuration where build_for_cudagraph_capture's branches all no-op, including the kv_indices.zero_() / kv_last_page_lens.fill_(1) safety block the function exists for. The fixture leaves kv_last_page_lens at zeros where production fills 1, which is the precondition for the underflow the function's own 30-line comment documents. tests/test_compress_plan.py:39's _FakeBuf already backs .gpu with torch.from_numpy — this is a third copy that diverged in the load-bearing direction.


Checked and refuted

Two alarms that did not survive verification, recorded so they are not re-raised: "TBO capture bakes zeros" — the captured values are bit-identical pre/post PR, and zero is the self-consistent ubatch convention; and "capture-time = 1 reads out of bounds" — block 0 is always allocated, so the read is in-bounds.

Signed-off-by: zhuyuhua-v <yuhzhu@amd.com>

Copilot AI 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.

Pull request overview

Copilot reviewed 5 out of 5 changed files in this pull request and generated no new comments.

@zhuyuhua-v

Copy link
Copy Markdown
Collaborator Author

Reviewed at 54115f034. The core mechanism is sound and correctly wired: pinning the fast branch at capture instead of letting the graph record the elementwise fallback is the right fix, and the guards, row counts and buffer sizes on the eager / capture / TBO paths all check out (copy_to_gpu is in-place, so the captured views stay valid). No crash-class defect. Two things worth acting on before merge, then a list.

1. The win does not land on the default --level 3 piecewise path

The consumer accepts the published buffer only on exact equality:

# atom/model_ops/dcp_ops.py:960-961
local_ctx = getattr(attn_metadata, "dcp_local_context_lens", None)
if local_ctx is not None and local_ctx.shape[0] == num_rows:
    return local_ctx

num_rows is the caller's num_decode_tokens, which is scheduled-width:

# atom/models/deepseek_v2.py:1444-1448
num_decode_tokens = (
    context.scheduled_bs * attn_metadata.max_seqlen_q
    if not context.is_prefill
    else 0
)

while the publisher is running-width:

# atom/model_ops/attentions/aiter_mla.py:2034
self._attach_dcp_local_context_lens(attn_metadata, running_bs, copy_to_gpu=True)

running_bs is the capture-ladder rung and is set independently of use_cudagraph (forward_context.py:301-305). Under PIECEWISE — the default (arg_utils.py:249, config.py:1898-1904) — the indexer runs eagerly every step with the real Context, so any batch size off the ladder fails the equality and the fallback's elementwise kernels run on all 21 full-index layers. Only --cudagraph-mode FULL replay and TBO ubatches actually get the win.

Worth confirming which mode the reported benchmark used.

The fix: read the width the engine already stored, rather than re-deriving it

Context already carries this number, and its own docstring says why the product form is the wrong way to get it:

# atom/utils/forward_context.py:485-490
# The step's DP-unified padded shape. `running_bs` counts SEQUENCES (graph
# identity, the draft's pad width), `running_tokens` the hidden_states rows
# MoE pads to. Both stored, because the ratio is not always max_seqlen_q --
# a DSpark ragged step runs a packed width no rectangular bs*q recovers.
running_bs: int = 0
running_tokens: int = 0

So:

num_decode_tokens = context.running_tokens if not context.is_prefill else 0

Three reasons this is better than either widening the guard or spelling running_bs * max_seqlen_q:

  1. It is exactly the height hidden_states has. The piecewise branch calls the model with self.forward_vars["input_ids"].gpu[:running_tokens] (model_runner.py:3030), and forward_context.py:426-427 states it directly — "the cudagraph branch re-slices the buffer to running_tokens itself". So num_tokens = hidden_states.shape[0] at deepseek_v2.py:1550 is running_tokens; using it for num_decode_tokens makes q_fp8[:num_decode_tokens] cover the whole real buffer and leaves q_fp8[num_decode_tokens:num_tokens] correctly empty on a pure-decode step.
  2. The guard then passes exactly, with no change to dcp_ops.py. dcp_ops.py:1050 asserts attn_metadata.max_seqlen_q == 1 on this path, so running_tokens == running_bs * 1 == running_bs — and the published buffer is running_bs rows. Both sides are the same number.
  3. It removes a rectangular assumption rather than replacing it. Both scheduled_bs * max_seqlen_q and running_bs * max_seqlen_q re-derive a quantity the engine already computed, using a ratio the Context docstring above explicitly warns does not always hold.

On the cost side there is nothing new to justify: decode is padded inside the graph already, so this only makes eager decode behave the way graph decode always has. The extra rows are inert either way — slot_mapping[:running_tokens] = -1 (aiter_mla.py:1829) so nothing is written to cache, and context_lens[scheduled_bs:running_bs] = 0 (:1834).

(If you would rather not touch num_decode_tokens, relaxing dcp_ops.py:961 to >= num_rows plus local_ctx[:num_rows] also works — the producer fills real values first and zero-pads the tail, so the prefix is exactly the real rows. But that leaves the scheduled/running mismatch in place for the next buffer that hits it.)

2. The publish guard was narrowed; the host-side producer was not

# aiter_mla.py:1835   producer
if self.dcp_world_size > 1:                       # no is_sparse term

# aiter_mla.py:1760   consumer-side helper
if not self.is_sparse or self.dcp_world_size <= 1:
    attn_metadata.dcp_local_context_lens = None
    return

So on dense-MLA DCP the two numpy slice-assigns at :1847-1848 still run every decode step for a buffer that is then published as None, plus a permanently resident pinned+device CpuGpuBuffer. And the comment above them still reads "Publish it: the sparse indexer used to re-derive this on device with 8 elementwise kernels per full-index layer" — describing a publish that no longer happens on that configuration. A reader on a dense-DCP config concludes the attribute is set and writes a consumer that dereferences None.

Same axis one file up: the declaration comment at :332-337 still gives DCP as the only publish condition, while the new per-ubatch twin at :578 documents the correct is_sparse and dcp > 1 gate — two declarations of one quantity now describing two different contracts. The buffer at :337 is also allocated unconditionally, for every model including dcp=1, which is a third predicate for the same object.

Smaller items

The predicate is spelled at five places that must agree, and each failure mode is different and silent. Sites: :581 (ubatch allocation), :1762 (helper guard), :2194 (vars_used append), :2308 (capture fill), plus :2135+:2149 as a nested pair. Missing the allocation gives a KeyError in the helper (it indexes forward_vars with [], not .get); missing the helper silently falls back to the kernels this PR exists to remove; missing the vars_used append publishes a GPU view of a buffer that was never copied — stale data, correct shape, wrong answers. One cached attribute in __init__ (e.g. self._publishes_dcp_local_lens) would name the decision once.

Declare the field instead of setting it dynamically. dcp_local_context_lens is neither a dataclass field nor an __init__ parameter of AttentionMetaData, so fields(self) / asdict_zerocopy (forward_context.py:668) cannot see it and the sole consumer must use getattr(..., None). AttentionMetaData_DSV4 declares dspark_ragged_lens_gpu: torch.Tensor | None = None as a real field (deepseek_v4_attn.py:214-215), and the sibling g_kv_indptr gets a hard assert at its consumer rather than a silent fallback. One line turns "a new backend forgot to publish" from an unmeasurable perf cliff into something a reader can see — and it deletes the helper's dead = None assignment.

The per-ubatch buffer and its copies are avoidable. Unlike kv_indptr / g_kv_indptr / sparse_kv_indptr, dcp_local_context_lens needs no rebasing, so build_ubatch_metadata could hand out var["dcp_local_context_lens"].gpu[req_start : req_start + running_bs]. prepare_decode already fills [:scheduled_bs] real and [scheduled_bs:running_bs] = 0, so ub0's [0:half) and ub1's [half:bs) land on exactly the values the copies reproduce, and the offset is fixed per captured bs so the view is graph-pointer-stable. As written, a change whose purpose is removing per-step work adds 2 pinned buffers, 4 numpy slice-assigns and 2 cudaMemcpyAsync per TBO decode step to reproduce data already resident on device at a known offset.

One operation split across 34 lines, with a parameter that is derivable. The guard+fill+copy at :2308-2311 and the attach at :2342 repeat the same predicate and must move together, and neither references the other — hoisting the attach into the AttentionMetaData constructor kwargs (next to g_kv_indptr at :2339, where it naturally belongs) silently strands the copy. Collapsing to self._attach_dcp_local_context_lens(attn_matadata, bs, copy_to_gpu=True) leaves one guard and one copy; and then copy_to_gpu=True holds exactly when prefix == "" (:2034, :2342) and False exactly for prefix=p (:2430, already copied via vars_used), so the parameter can be deleted.

Three comments that no longer match the code.

  • The helper's docstring (:1750) says it covers "every decode metadata path" and exists "to prevent any path from silently recording the elementwise fallback". Neither holds: prepare_mtp_decode (:821) is a fourth decode path that does not route through it; the helper is opt-in at three call sites, so a fifth path simply would not call it; and its negative arm assigns None, i.e. manufactures the fallback state it claims to prevent.
  • :2309 says "Capture uses one synthetic local KV token per request", implying the value is load-bearing. The branch it influences is a presence/shape test, not a value test (dcp_ops.py:960-961), and no launch grid on this path derives from local_ctx's values. Attaching the buffer is the real fix; np[:bs] = 1 only changes what the throwaway warmup forward reads. As written the comment invites a future reader to replicate a value that does nothing — and it already did: three review angles independently flagged the ubatch path for "missing" the same seed, which the ubatch zero-length convention does not need.
  • Three sites (dcp_ops.py:956, aiter_mla.py:336, :1846) say "8 elementwise kernels" while the PR description says 7. Counting the fallback body at dcp_ops.py:963-968: floordiv, mul, mul, sub, sub, clamp, add = 7; the trailing .to(torch.int32) is a no-op because context_lens is already int32. Worth picking one number — a reader who budgets 8 and measures 7 goes looking for a missing kernel, or makes context_lens int64 somewhere and silently turns the .to() into a real eighth that nobody re-counts.

Tests: none of the four run in CI, and they would not catch a revert of the TBO half.

.github/scripts/run_unit_tests.sh states its scope as "pure-Python unit tests that run on a plain runner (CPU torch, no GPU, no aiter/MoRIIO native libs)". aiter_mla.py:11 imports aiter at module scope, and this PR's new from atom.model_ops.dcp_ops import dcp_local_context_lens adds a top-level import triton (dcp_ops.py:18-19) — widening the skip predicate, so even a CI image that gains aiter keeps skipping. The tested logic (_attach_dcp_local_context_lens' branching, the capture fill) is pure Python + numpy and needs neither.

And on a revert run: deleting the PR's actual TBO code — the _prepare_ubatch_decode slice-copy at :2151-2154 and the vars_used append at :2194 — leaves all four tests green. test_dcp_local_context_attachment_supports_ubatch_prefix pre-populates var["ub0_dcp_local_context_lens"].gpu itself and then asserts the helper slices it, exercising buffer.gpu[:rows], which was never in doubt; no test calls _prepare_ubatch_decode or build_ubatch_metadata. The code that can actually be wrong is the numpy block — the req_start slicing, the ub_real_reqs zero-fill bound, the nesting of the is_sparse guard inside the dcp>1 guard, and the H2D registration.

Two smaller test notes: the pair at :409 / :417 have identical bodies differing only in fixture arguments and should be one @pytest.mark.parametrize over [(1, True), (4, False)]; and assert "dcp_local_context_lens" not in var is false in production, since :337 allocates it unconditionally — so if someone relaxes the helper's guard to just dcp>1, the tests fail with KeyError while production would not: it would publish an allocated-but-never-written all-zeros buffer that the consumer's shape check accepts, and the natural "fix" is to add the key to the fixture and go green.

Also _FakeMetadataBuffer (:356) backs .gpu with a numpy array, which forces block_size=1 and no mtp_k — the one configuration where build_for_cudagraph_capture's branches all no-op, including the kv_indices.zero_() / kv_last_page_lens.fill_(1) safety block the function exists for. The fixture leaves kv_last_page_lens at zeros where production fills 1, which is the precondition for the underflow the function's own 30-line comment documents. tests/test_compress_plan.py:39's _FakeBuf already backs .gpu with torch.from_numpy — this is a third copy that diverged in the load-bearing direction.

Checked and refuted

Two alarms that did not survive verification, recorded so they are not re-raised: "TBO capture bakes zeros" — the captured values are bit-identical pre/post PR, and zero is the self-consistent ubatch convention; and "capture-time = 1 reads out of bounds" — block 0 is always allocated, so the read is in-bounds.

updated

@valarLip
valarLip merged commit 6540cfd into main Sep 1, 2026
28 of 79 checks passed
@valarLip
valarLip deleted the yuhua/dcp-fusion branch September 1, 2026 11:59
Jasen2201 added a commit that referenced this pull request Sep 1, 2026
Three files conflicted, all where main's #2102 (published DCP local context
lens) landed next to this branch's replicated index cache:

- forward_context.py / aiter_mla.py: both sides declare a new optional field
  and a new __init__ line in the same spot. Kept both; they are independent.
- test_mla_index_cache.py: both sides append tests. Kept both, and dropped the
  function-local `import torch` now that main imports it at module scope.

The two features do not overlap at runtime: dcp_local_context_lens localizes
context for scoring a *sharded* index cache, and the replicated path skips that
consumer entirely via _dcp_index_comm_required().
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.

4 participants