[DCP] Fix sparse indexer local context metadata for CUDA graphs - #2102
Conversation
Signed-off-by: zhuyuhua-v <yuhzhu@amd.com>
🏷️ CI GuideRuns automatically on every eligible PR before approval:
Heavy model tests:
|
There was a problem hiding this comment.
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_lensinto 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.
|
Reviewed at 1. The win does not land on the default
|
Signed-off-by: zhuyuhua-v <yuhzhu@amd.com>
updated |
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().
Motivation
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

after

7 elementwise kernel are removed.
Test Plan
gsm8k 5shot & 20shot for DCP4+TP4
server
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=1acc 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=1Test 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