diff --git a/.github/benchmark/models_atomesh.yaml b/.github/benchmark/models_atomesh.yaml
index c1be8bc3f0..64d9ca031c 100644
--- a/.github/benchmark/models_atomesh.yaml
+++ b/.github/benchmark/models_atomesh.yaml
@@ -325,6 +325,78 @@ models:
name: glm-52-mxfp4-1p1d-tp4-agentic-1m-c48
concurrency: [48]
+ # CPP4 prefill (PP4 x TP1) paired with TP4+DCP4 decode on one node.
+ - &glm52_agentic_lmcache_cpp4_dcp4
+ <<: *glm52_agentic_tp4_tp4
+ name: glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c48
+ topology: 1p1d_cpp4_dcp4
+ pd_worker_layout: single_node
+ concurrency: [48]
+ server:
+ common_args:
+ kv_cache_dtype: fp8
+ block_size: 16
+ max_num_seqs: 512
+ enable_prefix_caching: true
+ decode_max_num_seqs: 512
+ # Empty means "do not pass the flag", so the model config's own
+ # context length applies. The 16384 the model defaults carry would
+ # truncate this suite's 1M agentic traces.
+ max_model_len: ""
+ max_num_batched_tokens: ""
+ decode_max_num_batched_tokens: ""
+ online_quant_config: '{"global_quant_config":"ptpc_fp8","exclude_layer":["lm_head","model.embed_tokens","*.mlp.gate","*expert*"]}'
+ gpu_memory_utilization: 0.85
+ extra_args: "--level 3 --method mtp --num-speculative-tokens 3 --spec-decode-acceptance-rate 0.6633"
+ prefill:
+ workers: 1
+ tp: 1
+ cudagraph: none
+ extra_args: "--pipeline-parallel-size 4 --enforce-eager --max-num-batched-tokens 8192"
+ decode:
+ workers: 1
+ tp: 4
+ cudagraph: >-
+ [1,2,4,8,16,24,32,40,48,56,64,72,80,88,96,104,112,120,128,136,144,152,160,168,176,184,192,200,208,216,224,232,240,248,256]
+ extra_args: "--decode-context-parallel-size 4 --cudagraph-mode FULL"
+ env:
+ common:
+ ATOM_MLA_PAGE_SIZE: "1"
+ ATOM_ONLINE_QUANT_STREAMING: "0"
+ ATOM_SPARSE_INDEXER_LOGITS_BUDGET_MB: "2047"
+ ATOM_USE_TRITON_MLA: "0"
+ MAX_JOBS: "16"
+ PYTHONHASHSEED: "0"
+ TOPK_FORCE_PATH: "one"
+ prefill:
+ HIP_VISIBLE_DEVICES: "0,1,2,3"
+ VLLM_PP_LAYER_PARTITION: "20,20,20,18"
+ LMCACHE_LOCAL_CPU: "True"
+ LMCACHE_MAX_LOCAL_CPU_SIZE: "256"
+ LMCACHE_CHUNK_SIZE: "256"
+ OFFLOAD_PROFILE: "1"
+ OFFLOAD_MIN_LOAD_TOKENS: "0"
+ PREFILL_KV_TRANSFER_CONFIG: >-
+ {"kv_connector":"multi","connectors":[{"kv_connector":"mooncake","kv_role":"kv_producer","proxy_ip":"${ROLE_IP}","handshake_port":${HANDSHAKE_PORT},"protocol":"rdma"},{"kv_connector":"lmcache_offload","kv_role":"offload"}]}
+ decode:
+ HIP_VISIBLE_DEVICES: "4,5,6,7"
+ # Decode-only: the sharded index cache reaches its global top-k
+ # through a candidate all-gather that only handles qlen=1, so MTP's
+ # 4-token verify aborts on it. Replicating the index cache removes
+ # that exchange. Setting it on prefill instead raises at startup --
+ # prefill runs pp4/dcp1 and the layout rejects both.
+ ATOM_DCP_REPLICATE_INDEX_CACHE: "1"
+ DECODE_KV_TRANSFER_CONFIG: >-
+ {"kv_connector":"mooncake","kv_role":"kv_consumer","proxy_ip":"${ROLE_IP}","handshake_port":${HANDSHAKE_PORT},"protocol":"rdma"}
+
+ - <<: *glm52_agentic_lmcache_cpp4_dcp4
+ name: glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c32
+ concurrency: [32]
+
+ - <<: *glm52_agentic_lmcache_cpp4_dcp4
+ name: glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c40
+ concurrency: [40]
+
- &glm52_agentic_lmcache_tp4_dpa
<<: *glm52_agentic_tp4_tp4
name: glm-52-mxfp4-1p1d-tp4-dpa-agentic-lmcache-1m-c24
diff --git a/.github/scripts/atomesh/pd_matrix.py b/.github/scripts/atomesh/pd_matrix.py
index b80efb095c..e4bfe92610 100644
--- a/.github/scripts/atomesh/pd_matrix.py
+++ b/.github/scripts/atomesh/pd_matrix.py
@@ -206,7 +206,9 @@ def role_env(
model_cfg.get("env", {}).get(role, {}),
suite_cfg.get("env", {}).get(role, {}),
)
- env = resolve_env_refs_in_value(env, preserve_names={"ROLE_IP"})
+ # ROLE_IP and HANDSHAKE_PORT are filled in by pd_server_atom.sh at
+ # launch time (host IP and 6301 + ATOMESH_SERVICE_PORT_OFFSET).
+ env = resolve_env_refs_in_value(env, preserve_names={"ROLE_IP", "HANDSHAKE_PORT"})
return {str(key): str(value) for key, value in env.items()}
diff --git a/.github/scripts/atomesh/pd_server_atom.sh b/.github/scripts/atomesh/pd_server_atom.sh
index f78eca6aac..83ec88984d 100644
--- a/.github/scripts/atomesh/pd_server_atom.sh
+++ b/.github/scripts/atomesh/pd_server_atom.sh
@@ -242,14 +242,44 @@ dump_launch_info() {
apply_prefixed_env() {
local prefix="$1"
local role_ip="$2"
+ local handshake_port="${3:-${HANDSHAKE_PORT}}"
local name raw value
while IFS='=' read -r name raw; do
[[ "${name}" == "${prefix}"* ]] || continue
value="${raw//\$\{ROLE_IP\}/${role_ip}}"
+ value="${value//\$\{HANDSHAKE_PORT\}/${handshake_port}}"
export "${name#${prefix}}=${value}"
done < <(env)
}
+# Names the last apply_role_env() exported for a role.
+ROLE_ENV_NAMES=()
+
+# Both servers are launched from this one shell, so a variable exported for one
+# role stays in the environment the next role inherits. Only names the two roles
+# both define get overwritten; a prefill-only name reaches decode unchanged --
+# VLLM_PP_LAYER_PARTITION from a pp4 prefill aborts a pp1 decode's model build
+# with "len(partitions)=4 does not match pp_size=1". Drop the previous role's
+# names before applying this one's.
+apply_role_env() {
+ local prefix="$1"
+ local role_ip="$2"
+ local handshake_port="${3:-${HANDSHAKE_PORT}}"
+ local name
+ for name in ${ROLE_ENV_NAMES[@]+"${ROLE_ENV_NAMES[@]}"}; do
+ unset "${name}"
+ done
+ ROLE_ENV_NAMES=()
+ while IFS='=' read -r name _; do
+ [[ "${name}" == "${prefix}"* ]] || continue
+ ROLE_ENV_NAMES+=("${name#${prefix}}")
+ done < <(env)
+ # A name the common block also sets was just unset with the previous role's,
+ # so put the common value back before the role overrides it.
+ apply_prefixed_env "ATOMESH_ENV_" "${role_ip}" "${handshake_port}"
+ apply_prefixed_env "${prefix}" "${role_ip}" "${handshake_port}"
+}
+
host_ip="$(echo "${IPADDRS}" | tr ',' '\n' | sed -n "$((NODE_RANK + 1))p")"
if [[ -z "${host_ip}" ]]; then
host_ip="$(hostname -I 2>/dev/null | awk '{print $1}')"
@@ -543,7 +573,7 @@ start_prefill() {
local handshake_port="${3:-${HANDSHAKE_PORT}}"
local dp_master_port="${4:-${PREFILL_DP_MASTER_PORT}}"
local dp_base_port="${5:-${PREFILL_DP_BASE_PORT}}"
- apply_prefixed_env "ATOMESH_PREFILL_ENV_" "${host_ip}"
+ apply_role_env "ATOMESH_PREFILL_ENV_" "${host_ip}" "${handshake_port}"
local -a prefill_cache_env=()
build_server_cache_env "prefill" "${server_port}" prefill_cache_env
local -a prefill_dp_env=()
@@ -580,7 +610,7 @@ start_decode() {
local handshake_port="${3:-${HANDSHAKE_PORT}}"
local dp_master_port="${4:-${DECODE_DP_MASTER_PORT}}"
local dp_base_port="${5:-${DECODE_DP_BASE_PORT}}"
- apply_prefixed_env "ATOMESH_DECODE_ENV_" "${host_ip}"
+ apply_role_env "ATOMESH_DECODE_ENV_" "${host_ip}" "${handshake_port}"
local max_conc
max_conc="$(echo "${BENCH_MAX_CONCURRENCY}" | tr 'x,' '\n' | sort -n | tail -1)"
local decode_max_num_seqs="${MAX_NUM_SEQS}"
diff --git a/.github/scripts/atomesh/process_result.py b/.github/scripts/atomesh/process_result.py
index 9528d216ca..062221e13b 100644
--- a/.github/scripts/atomesh/process_result.py
+++ b/.github/scripts/atomesh/process_result.py
@@ -27,6 +27,9 @@
)
TOPOLOGY_RE = re.compile(r"(?P
\d+)p(?P\d+)d", re.IGNORECASE)
TP_RE = re.compile(r"tp(?P\d+)", re.IGNORECASE)
+DUAL_TP_RE = re.compile(r"tp(?P\d+)-tp(?P\d+)", re.IGNORECASE)
+CPP_PP_RE = re.compile(r"(?:^|[\s_-])(?:cpp|pp)(?P\d+)(?:$|[\s_-])", re.IGNORECASE)
+PP_ARG_RE = re.compile(r"--pipeline-parallel-size(?:=|\s+)(\d+)", re.IGNORECASE)
EVAL_CONC_RE = re.compile(r"(?:^|[_-])c(?P\d+)(?:$|[_-])", re.IGNORECASE)
EVAL_TOPOLOGY_RE = re.compile(
r"(?:^|[_-])(?P\d+p\d+d(?:[_-]dpa)?)(?:$|[_-])",
@@ -101,6 +104,12 @@ def int_value(*values: Any) -> int | None:
return int(parsed) if parsed is not None else None
+def pp_from_server_args(payload: dict[str, Any], key: str) -> int | None:
+ """Pipeline-parallel size as the role's own launch flag set it."""
+ match = PP_ARG_RE.search(string_value(payload.get(key)))
+ return int(match.group(1)) if match else None
+
+
def round_or_none(*values: Any, digits: int = 4) -> float | None:
parsed = number(*values)
return round(parsed, digits) if parsed is not None else None
@@ -202,7 +211,6 @@ def topology_resources(
)
)
topology = TOPOLOGY_RE.search(text)
- tp = TP_RE.search(text)
prefill_workers = int_value(
payload.get("prefill_workers"), payload.get("num_prefill_workers")
)
@@ -219,16 +227,43 @@ def topology_resources(
decode_tp = int_value(
payload.get("decode_tp"), payload.get("decode_tensor_parallel_size")
)
- if tp:
- prefill_tp = prefill_tp or int(tp.group("tp"))
- decode_tp = decode_tp or int(tp.group("tp"))
+ dual_tp = DUAL_TP_RE.search(text)
+ if dual_tp:
+ prefill_tp = prefill_tp or int(dual_tp.group("prefill_tp"))
+ decode_tp = decode_tp or int(dual_tp.group("decode_tp"))
+ else:
+ tp = TP_RE.search(text)
+ if tp:
+ tp_size = int(tp.group("tp"))
+ prefill_tp = prefill_tp or tp_size
+ decode_tp = decode_tp or tp_size
+
+ # The launch flag outranks the topology label: it is what the server ran with.
+ prefill_pp = int_value(
+ payload.get("prefill_pp"), payload.get("prefill_pipeline_parallel_size")
+ )
+ if prefill_pp is None:
+ prefill_pp = pp_from_server_args(payload, "prefill_extra_server_args")
+ if prefill_pp is None:
+ cpp_pp = CPP_PP_RE.search(text)
+ if cpp_pp:
+ prefill_pp = int(cpp_pp.group("pp"))
+ prefill_pp = prefill_pp or 1
+
+ # No topology label encodes a decode-side PP, so the flag is the only source.
+ decode_pp = int_value(
+ payload.get("decode_pp"), payload.get("decode_pipeline_parallel_size")
+ )
+ if decode_pp is None:
+ decode_pp = pp_from_server_args(payload, "decode_extra_server_args")
+ decode_pp = decode_pp or 1
num_prefill_gpu = int_value(payload.get("num_prefill_gpu"))
num_decode_gpu = int_value(payload.get("num_decode_gpu"))
if num_prefill_gpu is None and prefill_workers and prefill_tp:
- num_prefill_gpu = prefill_workers * prefill_tp
+ num_prefill_gpu = prefill_workers * prefill_tp * prefill_pp
if num_decode_gpu is None and decode_workers and decode_tp:
- num_decode_gpu = decode_workers * decode_tp
+ num_decode_gpu = decode_workers * decode_tp * decode_pp
total_gpu = int_value(payload.get("total_gpu"))
if total_gpu is None and num_prefill_gpu is not None and num_decode_gpu is not None:
total_gpu = num_prefill_gpu + num_decode_gpu
@@ -339,6 +374,10 @@ def enrich_payload(
enriched.setdefault("decode_workers", env.get("DECODE_WORKERS"))
enriched.setdefault("prefill_tp", env.get("PREFILL_TP"))
enriched.setdefault("decode_tp", env.get("DECODE_TP"))
+ enriched.setdefault(
+ "prefill_extra_server_args", env.get("PREFILL_EXTRA_SERVER_ARGS")
+ )
+ enriched.setdefault("decode_extra_server_args", env.get("DECODE_EXTRA_SERVER_ARGS"))
runner = env.get("SLURM_SUBMIT_RUNNER", "")
if hardware:
enriched["hardware"] = hardware
diff --git a/.github/workflows/atomesh-benchmark.yaml b/.github/workflows/atomesh-benchmark.yaml
index 0332b83866..0040770322 100644
--- a/.github/workflows/atomesh-benchmark.yaml
+++ b/.github/workflows/atomesh-benchmark.yaml
@@ -60,6 +60,9 @@ on:
glm-52-mxfp4-1p1d-tp4-agentic-1m-c1,glm-52-mxfp4-1p1d-tp4-agentic-1m-c2,
glm-52-mxfp4-1p1d-tp4-agentic-1m-c4,glm-52-mxfp4-1p1d-tp4-agentic-1m-c8,
glm-52-mxfp4-1p1d-tp4-agentic-1m-c48,
+ glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c32,
+ glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c40,
+ glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c48,
glm-52-mxfp4-1p1d-tp4-dpa-agentic-lmcache-1m-c24,
glm-52-mxfp4-1p1d-tp4-dpa-agentic-lmcache-1m-c32,
glm-52-mxfp4-1p1d-tp4-dpa-agentic-lmcache-1m-c48
@@ -124,6 +127,10 @@ on:
- mia1-p02-g42,mia1-p02-g44
- mia1-p02-g42,mia1-p02-g47
- mia1-p02-g44,mia1-p02-g47
+ # single_node cases only use the first entry, so keep orderings that
+ # let any one of the three nodes be selected.
+ - mia1-p02-g47,mia1-p02-g44
+ - mia1-p01-g43,mia1-p01-g36
- mia1-p01-g36,mia1-p01-g43
atomesh_2p1d_nodes:
description: 'ATOMesh 2P1D nodes'
diff --git a/atom/distributed/dcp_utils.py b/atom/distributed/dcp_utils.py
index 26a1e9b77a..6065614b59 100644
--- a/atom/distributed/dcp_utils.py
+++ b/atom/distributed/dcp_utils.py
@@ -18,6 +18,7 @@
"""
from atom.config import get_current_atom_config
+from atom.utils import envs
def get_dcp_world_size() -> int:
@@ -36,6 +37,18 @@ def dcp_is_enabled() -> bool:
return get_dcp_world_size() > 1
+def dcp_replicated_index_cache_enabled(atom_config=None) -> bool:
+ """Whether native GLM-5.2 uses a full replicated index cache under DCP."""
+ if not envs.ATOM_DCP_REPLICATE_INDEX_CACHE:
+ return False
+ config = atom_config or get_current_atom_config()
+ hf_config = config.hf_config
+ return (
+ config.decode_context_parallel_size > 1
+ and getattr(hf_config, "model_type", None) == "glm_moe_dsa"
+ )
+
+
def get_dcp_group():
"""The DCP process group (aiter parallel state). Only valid when DCP is enabled."""
from aiter.dist.parallel_state import get_dcp_group as _get_dcp_group
diff --git a/atom/kv_transfer/disaggregation/mooncake/mooncake_connector.py b/atom/kv_transfer/disaggregation/mooncake/mooncake_connector.py
index d7d23b0177..ab94453bd4 100644
--- a/atom/kv_transfer/disaggregation/mooncake/mooncake_connector.py
+++ b/atom/kv_transfer/disaggregation/mooncake/mooncake_connector.py
@@ -22,11 +22,16 @@
import msgpack
import msgspec
+import numpy as np
import torch
import zmq
from aiter.dist.parallel_state import get_dp_group, get_tp_group
from atom.config import Config
+from atom.distributed.dcp_utils import (
+ dcp_replicated_index_cache_enabled,
+ get_dcp_group,
+)
from atom.kv_transfer.disaggregation.base import (
KVConnectorBase,
KVConnectorSchedulerBase,
@@ -38,6 +43,8 @@
side_channel_port_offset as _port_offset,
)
from atom.kv_transfer.disaggregation.types import (
+ INDEX_CACHE_ROLE,
+ MLA_KV_ROLE,
ConnectorMetadata,
KVTransferRegion,
ReqId,
@@ -143,6 +150,119 @@ def _configure_mooncake_transport(protocol: str) -> None:
os.environ["MC_FORCE_TCP"] = "true"
+def _coalesce(
+ src: np.ndarray, dst: np.ndarray, length: np.ndarray
+) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
+ """Merge runs that are contiguous on both sides into single descriptors."""
+ if src.size == 0:
+ empty = np.empty(0, dtype=np.int64)
+ return empty, empty.copy(), empty.copy()
+ contiguous = (src[1:] == src[:-1] + length[:-1]) & (
+ dst[1:] == dst[:-1] + length[:-1]
+ )
+ starts = np.concatenate(([True], ~contiguous))
+ group = np.cumsum(starts) - 1
+ merged_len = np.bincount(group, weights=length).astype(np.int64)
+ return src[starts], dst[starts], merged_len
+
+
+def plan_sharded(
+ src_block_ids,
+ dst_block_ids,
+ block_size: int,
+ dcp_size: int,
+ dcp_rank: int,
+ interleave_size: int = 1,
+) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
+ """Plan the KV (interleave-sharded) transfer for one DCP rank.
+
+ Walks ``block_size // interleave_size`` runs per destination block. Runs
+ whose source block is past the end of the producer's list are dropped: the
+ block manager sizes every rank's table from rank 0's share, so ranks above
+ it can own a trailing virtual block that has no source tokens at all.
+ """
+ S = int(interleave_size)
+ if not 1 <= S <= block_size or block_size % S:
+ # Runs are S long and start on an S-group boundary, so an interleave
+ # that does not tile a block puts a run astride two of them -- writing
+ # past one destination block into the next and leaving the tokens it
+ # skipped untransferred.
+ raise ValueError(
+ f"DCP interleave_size={interleave_size} must divide block_size="
+ f"{block_size}; the sharded plan cannot express a run that "
+ "straddles two blocks."
+ )
+ src_ids = np.asarray(src_block_ids, dtype=np.int64)
+ dst_ids = np.asarray(dst_block_ids, dtype=np.int64)
+
+ # One run per S-group of this rank's local slots. Starting each run on an
+ # S-group boundary is what lets dcp_global_pos reduce to the group form
+ # below, and keeps the whole run inside one source block.
+ local = np.arange(0, dst_ids.size * block_size, S, dtype=np.int64)
+ dst_block, dst_token = np.divmod(local, block_size)
+ g = ((local // S) * dcp_size + dcp_rank) * S
+ src_block, src_token = np.divmod(g, block_size)
+
+ keep = src_block < src_ids.size
+ return _coalesce(
+ src_ids[src_block[keep]] * block_size + src_token[keep],
+ dst_ids[dst_block[keep]] * block_size + dst_token[keep],
+ np.full(int(keep.sum()), S, dtype=np.int64),
+ )
+
+
+def plan_replicated_index(
+ src_block_ids,
+ dst_block_ids,
+ dcp_size: int,
+ src_page_bytes: int,
+ key_plane_bytes: int,
+ scale_plane_bytes: int,
+) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
+ """Plan the replicated index-cache transfer in bytes; same on every rank.
+
+ ``indexer_k_quant_and_cache(preshuffle=True)`` stores a page as an
+ MFMA-16x16-tiled fp8 key plane followed by an fp32 scale plane, so a page is
+ only copyable whole -- addressing one of its tokens is not a byte range at
+ all. The destination page is ``dcp_size`` source pages wide and interleaves
+ the two planes differently: its keys run ``dcp_size * key_plane_bytes``
+ before its scales start. So each source page crosses as a key run and a
+ scale run, landing at sub-page ``s`` of each plane rather than at one
+ sub-page offset of the page.
+ """
+ src_ids = np.asarray(src_block_ids, dtype=np.int64)
+ dst_ids = np.asarray(dst_block_ids, dtype=np.int64)
+
+ dst_block = np.repeat(np.arange(dst_ids.size, dtype=np.int64), dcp_size)
+ sub = np.tile(np.arange(dcp_size, dtype=np.int64), dst_ids.size)
+ src_block = dst_block * dcp_size + sub
+
+ keep = src_block < src_ids.size
+ src_page = src_ids[src_block[keep]] * src_page_bytes
+ dst_page = dst_ids[dst_block[keep]] * src_page_bytes * dcp_size
+ sub = sub[keep]
+ n = sub.size
+
+ # Two descriptors per source page, with no _coalesce pass: a page's key run
+ # stops short of the next page by that page's scale plane and padding, so
+ # consecutive block ids never make the source side contiguous.
+ return (
+ np.concatenate([src_page, src_page + key_plane_bytes]),
+ np.concatenate(
+ [
+ dst_page + sub * key_plane_bytes,
+ dst_page + dcp_size * key_plane_bytes + sub * scale_plane_bytes,
+ ]
+ ),
+ np.concatenate(
+ [
+ np.full(n, key_plane_bytes, dtype=np.int64),
+ np.full(n, scale_plane_bytes, dtype=np.int64),
+ ]
+ ),
+ )
+
+
# ZMQ side-channel message types
MSG_WRITE_REQUEST = b"write_request"
MSG_WRITE_DONE = b"write_done"
@@ -306,7 +426,8 @@ def __init__(self, config: Config) -> None:
self.dp_rank = config.parallel_config.data_parallel_rank
self.pp_size = config.pipeline_parallel_size
self.block_size = config.kv_cache_block_size
- self.hash_block_size = self.block_size * config.decode_context_parallel_size
+ self.dcp_size = config.decode_context_parallel_size
+ self.hash_block_size = self.block_size * self.dcp_size
self.host_ip = get_ip()
# Pending requests: req_id -> (Sequence, block_table)
@@ -386,15 +507,20 @@ def update_state_after_alloc(self, seq: Sequence) -> None:
# prefix cache. Per-request state (including the SWA ring slot) is
# not covered by a block-only delta, so it takes a full transfer.
num_computed_blocks = 0
+ # The producer turns each consumer block back into `dcp_size` of
+ # its own, so the two PHYSICAL block sizes have to match. The wire
+ # value is the producer's hash size, which is that only while the
+ # producer runs dcp=1, as CPP prefill does.
remote_hash_block_size = params.get("hash_block_size")
- if remote_hash_block_size != self.hash_block_size:
+ if remote_hash_block_size != self.block_size:
logger.warning(
"PD incremental transfer disabled for req %s: producer "
- "hash_block_size=%r, consumer hash_block_size=%d; "
+ "hash_block_size=%r, consumer block_size=%d (dcp=%d); "
"falling back to full transfer",
seq.id,
remote_hash_block_size,
- self.hash_block_size,
+ self.block_size,
+ self.dcp_size,
)
elif not seq.has_per_req_cache and self.hash_block_size > 0:
num_computed_blocks = seq.num_cached_tokens // self.hash_block_size
@@ -478,6 +604,15 @@ def __init__(self, config: Config) -> None:
self.pp_rank = config.parallel_config.pipeline_parallel_rank
self.pp_size = config.pipeline_parallel_size
self.num_hidden_layers = config.hf_config.num_hidden_layers
+ self.block_size = config.kv_cache_block_size
+ # The consumer ships its DCP topology in the write_request; the
+ # producer relayouts on the way out and keeps dcp_size == 1 of its own.
+ self.dcp_size = config.decode_context_parallel_size
+ self.dcp_rank = get_dcp_group().rank_in_group if self.dcp_size > 1 else 0
+ self.dcp_interleave_size = config.dcp_config.interleave_size
+ # Indexer regions then take the replicated plan while the MLA latent
+ # regions stay sharded.
+ self.replicate_index_cache = dcp_replicated_index_cache_enabled(config)
# Global index of this stage's first layer; consumer regions are ordered
# over all layers, so a producer stage writes at this layer offset.
self._start_layer = 0
@@ -595,6 +730,8 @@ def __init__(self, config: Config) -> None:
self.kv_caches: dict[str, Any] | None = None
self.kv_caches_base_addr: list[int] = []
self._per_block_bytes_list: list[int] = []
+ self._block_region_roles: list[str | None] = []
+ self._block_region_planes: list[tuple[int, int] | None] = []
self.kv_cache_shape: tuple[int, ...] | None = None
self.block_len: int = config.kv_cache_block_size
self.num_blocks: int = 0
@@ -769,6 +906,15 @@ def register_kv_caches(
self.kv_caches_base_addr = [r.base_addr for r in tt.block_regions]
self._per_block_bytes_list = [r.unit_bytes for r in tt.block_regions]
+ self._block_region_roles = [r.semantic_role for r in tt.block_regions]
+ self._block_region_planes = [
+ (
+ (r.key_plane_bytes, r.scale_plane_bytes)
+ if r.key_plane_bytes is not None
+ else None
+ )
+ for r in tt.block_regions
+ ]
# Under pipeline parallelism this stage holds only layers
# [start_layer, end_layer); its local regions map onto the consumer's
@@ -956,17 +1102,18 @@ def start_load_kv(self, metadata: ConnectorMetadata) -> None:
self._pending_recv_nonce[req_id] = write_nonce
# PD incremental: slice off locally cached prefix blocks; invalid
- # offset falls back to full transfer.
+ # offset falls back to full transfer. Under DCP the same prefix
+ # costs dcp_size times as many source blocks.
remote_block_ids = meta.remote_block_ids or []
off = meta.num_computed_blocks
if (
off < 0
or off >= len(meta.local_block_ids)
- or off >= len(remote_block_ids)
+ or off * self.dcp_size >= len(remote_block_ids)
):
off = 0
dst_block_ids = meta.local_block_ids[off:]
- src_block_ids = remote_block_ids[off:]
+ src_block_ids = remote_block_ids[off * self.dcp_size :]
# Build the (stage-independent) write_request payload once.
request_body = {
@@ -976,6 +1123,9 @@ def start_load_kv(self, metadata: ConnectorMetadata) -> None:
"consumer_rpc_port": self.rpc_port,
# Consumer's layer count per group for producer stride validation.
"consumer_num_layers": self._num_local_layers,
+ # Role of each of this side's block regions, so the producer can
+ # check its per-region plan lands on the same kind of region.
+ "consumer_region_roles": self._block_region_roles,
"dst_block_ids": dst_block_ids,
# Source block_ids so downstream stages (no scheduler, no
# _completed_prefills) can transfer without a local lookup.
@@ -987,6 +1137,11 @@ def start_load_kv(self, metadata: ConnectorMetadata) -> None:
"notify_port": self._notification_port,
"consumer_tp_size": self.tp_size,
"write_nonce": write_nonce,
+ # DCP relayout: which shard of each block this rank owns.
+ "consumer_dcp_size": self.dcp_size,
+ "consumer_dcp_rank": self.dcp_rank,
+ "consumer_dcp_interleave": self.dcp_interleave_size,
+ "consumer_replicates_index_cache": self.replicate_index_cache,
}
consumer_staging_pool_idx = -1
@@ -1202,9 +1357,25 @@ def _execute_transfer(self, request_data: dict) -> None:
notify_port = request_data["notify_port"]
consumer_tp_size = request_data.get("consumer_tp_size", self.tp_size)
consumers_per_rank = max(1, consumer_tp_size // self.tp_size)
+ consumer_dcp_size = max(1, request_data.get("consumer_dcp_size", 1))
write_nonce = request_data.get("write_nonce", 0)
has_slot_data = request_data.get("has_slot_regions", False)
+ if has_slot_data and consumer_dcp_size > 1:
+ # _execute_block_slot_transfer pairs whole blocks, so it has no
+ # way to express a consumer block that holds every dcp_size-th
+ # token. The block-count check below now admits that ratio, so
+ # refuse here rather than write another rank's KV.
+ logger.error(
+ "[PRODUCER] req %s carries per-request state regions and "
+ "the consumer runs dcp_size=%d; the slot-transfer path has "
+ "no DCP relayout. Aborting instead of writing misaligned "
+ "KV.",
+ req_id,
+ consumer_dcp_size,
+ )
+ return
+
logger.debug(
"[PRODUCER] _execute_transfer: req_id=%s, transfer_id=%s, "
"consumer=%s:%s, dst_blocks=%d, has_slot_data=%s",
@@ -1260,16 +1431,22 @@ def _execute_transfer(self, request_data: dict) -> None:
# PD incremental (TP-TP only): consumer already sliced dst; slice
# producer's src by the same offset. PP src arrives pre-sliced.
if self.pp_size == 1:
- off = request_data.get("num_computed_blocks", 0)
+ # The consumer's offset counts destination blocks; under DCP
+ # each of those spans consumer_dcp_size source blocks.
+ off = request_data.get("num_computed_blocks", 0) * consumer_dcp_size
if 0 < off < len(src_block_ids):
src_block_ids = src_block_ids[off:]
- if len(src_block_ids) != len(dst_block_ids):
+ expected_dst_blocks = -(-len(src_block_ids) // consumer_dcp_size)
+ if len(dst_block_ids) != expected_dst_blocks:
logger.error(
"[PRODUCER] src/dst block count mismatch for req %s "
- "(src=%d, dst=%d); aborting transfer to avoid misaligned KV.",
+ "(src=%d, dst=%d, expected dst=%d at dcp_size=%d); aborting "
+ "transfer to avoid misaligned KV.",
req_id,
len(src_block_ids),
len(dst_block_ids),
+ expected_dst_blocks,
+ consumer_dcp_size,
)
return
target = f"{consumer_host}:{consumer_rpc_port}"
@@ -1430,23 +1607,147 @@ def _execute_block_transfer(
request_data.get("consumer_num_layers"),
self._block_region_consumer_indices,
)
+ # The plan comes from this stage's role list but the bytes land at
+ # cmap[region_idx], and equal region counts do not make the two orders
+ # match. Writing an index region with the latent's plan corrupts it into
+ # plausible text instead of faulting, so check the pairing.
+ consumer_roles = request_data.get("consumer_region_roles")
+ if consumer_roles is not None:
+ for region_idx in range(num_regions):
+ cidx = cmap[region_idx]
+ if cidx >= len(consumer_roles):
+ raise RuntimeError(
+ f"Consumer sent {len(consumer_roles)} region roles for "
+ f"{len(consumer_base_addrs)} regions (req {req_id}); "
+ "region pairing cannot be checked against a peer this "
+ "far out of sync."
+ )
+ remote_role = consumer_roles[cidx]
+ if self._block_region_roles[region_idx] != remote_role:
+ raise RuntimeError(
+ f"Region role mismatch for req {req_id}: local region "
+ f"{region_idx} is "
+ f"{self._block_region_roles[region_idx]!r}, but consumer "
+ f"region {cmap[region_idx]} is {remote_role!r}"
+ )
+ # Under DCP the consumer rank owns only part of each block, so
+ # whole-block descriptors no longer line up and the push becomes the
+ # per-region relayout described at the top of this file. Safe in token
+ # units because DCP only runs on MLA, which stores a token contiguously.
+ dcp_size = max(1, request_data.get("consumer_dcp_size", 1))
+ interleave = request_data.get("consumer_dcp_interleave", 1)
+ replicates_index = request_data.get("consumer_replicates_index_cache", False)
+ sharded_plan = None
+ if dcp_size > 1:
+ if self.dcp_size > 1:
+ raise RuntimeError(
+ f"Producer runs dcp_size={self.dcp_size}; the relayout "
+ "reads its blocks as whole global ones, so a sharded "
+ "producer would be replanned against tokens it does not "
+ "hold. Run the prefill node without DCP."
+ )
+ sharded_plan = plan_sharded(
+ src_block_ids,
+ dst_block_ids,
+ self.block_size,
+ dcp_size,
+ request_data["consumer_dcp_rank"],
+ interleave,
+ )
+ # A preshuffled index page has no token-addressable bytes, so only
+ # a whole-page interleave keeps the sharded plan page-aligned.
+ if (
+ not replicates_index
+ and interleave != self.block_size
+ and INDEX_CACHE_ROLE in self._block_region_roles
+ ):
+ raise RuntimeError(
+ f"A DCP interleave of {interleave} shards the DSA index "
+ "cache below a page, which its preshuffled layout does not "
+ "allow. Run the decode node with "
+ "ATOM_DCP_REPLICATE_INDEX_CACHE=1, or with an interleave of "
+ f"{self.block_size}."
+ )
+
+ # plan_replicated_index reads only the block ids, dcp_size and the page
+ # geometry, and the first two are fixed for the whole request. Every
+ # index layer registers the same geometry, so without this the identical
+ # plan is rebuilt once per layer.
+ replicated_index_plans: dict[tuple[int, int, int], tuple] = {}
+
for region_idx in range(num_regions):
src_base = self.kv_caches_base_addr[region_idx]
dst_base = consumer_base_addrs[cmap[region_idx]]
bpb = self._per_block_bytes_list[region_idx]
- for sb, db in zip(src_block_ids, dst_block_ids):
- src_addrs.append(src_base + sb * bpb)
- dst_addrs.append(dst_base + db * bpb)
- sizes.append(bpb)
-
- logger.debug(
- "[PRODUCER] block RDMA write: req=%s, %d regions × %d blocks, "
- "total_bytes=%d",
- req_id,
- num_regions,
- len(src_block_ids),
- sum(sizes),
- )
+ role = self._block_region_roles[region_idx]
+ # sharded_plan addresses a block in MLA token units, so any other
+ # region would be relaid out under a rule that is not its layout.
+ if dcp_size > 1 and role != MLA_KV_ROLE and role != INDEX_CACHE_ROLE:
+ raise RuntimeError(
+ f"Region {region_idx} has semantic_role {role!r}; under "
+ f"consumer dcp_size={dcp_size} the only relayouts defined "
+ f"are {MLA_KV_ROLE} and {INDEX_CACHE_ROLE}."
+ )
+ plan = sharded_plan
+ unit = 0
+ if dcp_size > 1 and replicates_index and role == INDEX_CACHE_ROLE:
+ planes = self._block_region_planes[region_idx]
+ if planes is None:
+ raise RuntimeError(
+ f"Region {region_idx} is registered as "
+ f"{INDEX_CACHE_ROLE} but reports no key/scale plane "
+ "sizes; a replicated index page cannot be split into "
+ "its two destination runs without them."
+ )
+ key_bytes, scale_bytes = planes
+ geometry = (bpb, key_bytes, scale_bytes)
+ plan = replicated_index_plans.get(geometry)
+ if plan is None:
+ plan = plan_replicated_index(
+ src_block_ids,
+ dst_block_ids,
+ dcp_size,
+ bpb,
+ key_bytes,
+ scale_bytes,
+ )
+ replicated_index_plans[geometry] = plan
+ unit = 1 # already in bytes
+ if plan is None:
+ for sb, db in zip(src_block_ids, dst_block_ids):
+ src_addrs.append(src_base + sb * bpb)
+ dst_addrs.append(dst_base + db * bpb)
+ sizes.append(bpb)
+ continue
+ if not unit:
+ # The destination page is wider only in whole tokens and the
+ # plan already counts in its token space, so both ends scale by
+ # the source's per-token width.
+ if bpb % self.block_size:
+ raise RuntimeError(
+ f"Region {region_idx} stores {bpb} bytes per block, "
+ f"which block_size {self.block_size} does not divide. "
+ "Addressing a single token needs its bytes contiguous, "
+ "which holds for the MLA layout the token-unit relayout "
+ "is built on."
+ )
+ unit = bpb // self.block_size
+ # One descriptor per token on a token-unit plan, so scale in the
+ # array; tolist() boxes to the Python ints the engine takes.
+ plan_src, plan_dst, plan_len = plan
+ src_addrs.extend((src_base + plan_src * unit).tolist())
+ dst_addrs.extend((dst_base + plan_dst * unit).tolist())
+ sizes.extend((plan_len * unit).tolist())
+
+ if logger.isEnabledFor(logging.DEBUG):
+ logger.debug(
+ "[PRODUCER] block RDMA write: req=%s, %d regions × %d blocks, "
+ "total_bytes=%d",
+ req_id,
+ num_regions,
+ len(src_block_ids),
+ sum(sizes),
+ )
if not self._rdma_write_with_retry(
target, src_addrs, dst_addrs, sizes, req_id, "block"
diff --git a/atom/kv_transfer/disaggregation/types.py b/atom/kv_transfer/disaggregation/types.py
index 3f00478e8d..7e487302f0 100644
--- a/atom/kv_transfer/disaggregation/types.py
+++ b/atom/kv_transfer/disaggregation/types.py
@@ -97,6 +97,12 @@ def key(self) -> ConnectorCompletionKey:
return self.channel, self.operation_id
+# Region roles the DCP relayout dispatches on: an MLA latent region is always
+# interleave-sharded, a DSA index region may instead be replicated whole.
+MLA_KV_ROLE = "mla.kv"
+INDEX_CACHE_ROLE = "dsa.index_cache"
+
+
@dataclass
class KVTransferRegion:
"""One RDMA-registerable tensor region."""
@@ -113,6 +119,11 @@ class KVTransferRegion:
# and list positions are process-local implementation details; a named role
# makes equal-sized planes distinguishable across code versions.
semantic_role: str | None = None
+ # Plane split of a preshuffled DSA index page, which is not token-addressable
+ # (MFMA-16x16 tiled fp8 keys, then a plane of fp32 scales): a sub-page
+ # relayout moves each plane on its own. None for token-contiguous pages.
+ key_plane_bytes: int | None = None
+ scale_plane_bytes: int | None = None
def unit_addr(self, index: int) -> int:
if self.reverse_indexed:
diff --git a/atom/kv_transfer/offload/config.py b/atom/kv_transfer/offload/config.py
index 6bf4f573de..cb4b6504ac 100644
--- a/atom/kv_transfer/offload/config.py
+++ b/atom/kv_transfer/offload/config.py
@@ -23,6 +23,8 @@
import torch
+from atom.distributed.dcp_utils import dcp_replicated_index_cache_enabled
+
# Version 3 adds the effective index-cache dtype to the PAGE identity. FP4 and
# FP8 DSV4 indexers have different region counts and byte layouts, so they must
# never reuse one another's objects even when the HF model config is identical.
@@ -49,10 +51,12 @@
"qk_rope_head_dim",
"compress_ratios",
"indexer_dtype",
+ "indexer_types",
)
_HF_INTEGER_GEOMETRY_FIELDS = frozenset(_HF_PAGE_FIELDS) - {
"compress_ratios",
"indexer_dtype",
+ "indexer_types",
}
logger = logging.getLogger("atom")
@@ -200,6 +204,7 @@ def build_page_namespace(
),
minimum=1,
),
+ "replicated_index_cache": bool(dcp_replicated_index_cache_enabled(config)),
"hf_geometry": _stable_hf_geometry(hf),
"speculative_config": _stable_config_value(
getattr(config, "speculative_config", None)
diff --git a/atom/model_ops/attention_mla.py b/atom/model_ops/attention_mla.py
index 226852ac9e..b584b8d6eb 100644
--- a/atom/model_ops/attention_mla.py
+++ b/atom/model_ops/attention_mla.py
@@ -746,13 +746,8 @@ def __init__(
self.dcp_persistent_supported = dcp_persistent_supported()
self.dcp_prefill_merge_bf16_ok = dcp_prefill_merge_bf16_ok()
- # Scope sparse persistent DCP to native non-speculative attention.
- # Decode is q_len=1; sparse prefill is represented as per-token virtual
- # q_len=1 rows. Plugin DCP reconfigures its group after construction.
self.sparse_dcp_metadata_rebuild = (
- self.is_sparse_mla
- and self.dcp_world_size > 1
- and getattr(atom_config, "speculative_config", None) is None
+ self.is_sparse_mla and self.dcp_world_size > 1
)
# Compacted per-layer sparse offsets for DCP decode; rebound by the
@@ -1965,6 +1960,11 @@ def _rebuild_sparse_dcp_persistent_metadata(
"Sparse DCP persistent decode metadata rebuild currently "
"supports non-speculative q_len=1 only."
)
+ elif work_prefix == "sparse_mtp_":
+ assert attn_metadata.max_seqlen_q > 1, (
+ "sparse_mtp_ work buffers describe the per-token verify layout; "
+ "a q_len=1 step must use the unprefixed ones."
+ )
assert q.shape[1] == self.dcp_kernel_num_heads
get_mla_metadata_v1(
paged_cu_seqlens_q,
@@ -2101,8 +2101,8 @@ def _forward_decode(
paged_kv_indptr = attn_metadata.sparse_kv_indptr[: B + 1]
paged_kv_indices = self.sparse_kv_indices_buffer
paged_kv_last_page_lens = attn_metadata.sparse_kv_last_page_lens[:B]
- if self.dcp_world_size > 1:
- paged_kv_indptr = self.dcp_sparse_kv_indptr_buffer[: B + 1]
+ if self.dcp_world_size > 1:
+ paged_kv_indptr = self.dcp_sparse_kv_indptr_buffer[: B + 1]
dp_size = get_dp_group().world_size
use_persistent_mode = should_use_persistent_mode(
@@ -2112,10 +2112,10 @@ def _forward_decode(
dcp_world_size=self.dcp_world_size,
dcp_persistent_supported=self.dcp_persistent_supported,
)
- # Sparse DCP persistent decode is enabled only for ordinary q_len=1
- # native serving. Its full IndexShare layers rebuild the work plan
- # below from the layer-local compact indptr; plugin/speculative paths
- # retain the established non-persistent fallback.
+ # Sparse DCP persistent decode rebuilds the work plan below from
+ # the layer-local compact indptr on every full IndexShare layer. An
+ # MTP verify step is per-token q_len=1 rows and rebuilds into the
+ # sparse_mtp_ work buffers, so it stays on this path too.
if self.is_sparse_mla and self.dcp_world_size > 1:
use_persistent_mode = (
use_persistent_mode and self.sparse_dcp_metadata_rebuild
@@ -2134,6 +2134,7 @@ def _forward_decode(
paged_cu_seqlens_q,
paged_kv_indptr,
paged_kv_last_page_lens,
+ work_prefix="sparse_mtp_" if is_sparse_mtp else "",
)
if not use_persistent_mode:
diff --git a/atom/model_ops/attentions/aiter_mla.py b/atom/model_ops/attentions/aiter_mla.py
index 40c7572c1b..ea35940434 100644
--- a/atom/model_ops/attentions/aiter_mla.py
+++ b/atom/model_ops/attentions/aiter_mla.py
@@ -17,6 +17,7 @@
from atom.distributed.dcp_utils import (
dcp_persistent_supported,
+ dcp_replicated_index_cache_enabled,
get_dcp_rank,
get_dcp_world_size,
)
@@ -168,6 +169,84 @@ def cdiv(a, b):
return (a + b - 1) // b
+def _replicated_index_cache_transfer_supported(config) -> bool:
+ """Whether the target replicated-index transfer topology is configured."""
+
+ transfer_config = getattr(config, "kv_transfer_config", None)
+ if not isinstance(transfer_config, dict) or not transfer_config:
+ return False
+
+ from atom.kv_transfer.disaggregation.factory import KVConnectorFactory
+
+ def canonical(sub_config, path):
+ if not isinstance(sub_config, dict):
+ return None
+ try:
+ return KVConnectorFactory.canonical_name(
+ sub_config.get("kv_connector"), path=path
+ )
+ except (TypeError, ValueError):
+ return None
+
+ connector = canonical(transfer_config, "kv_transfer_config")
+ if connector == "lmcache_offload":
+ return True
+ if connector == "mooncake":
+ return transfer_config.get("kv_role", "kv_producer") in {
+ "kv_producer",
+ "kv_consumer",
+ }
+ if connector != "multi":
+ return False
+
+ sub_configs = transfer_config.get("connectors")
+ if not isinstance(sub_configs, list) or len(sub_configs) != 2:
+ return False
+ topology = {
+ (
+ canonical(sub_config, f"kv_transfer_config.connectors[{index}]"),
+ sub_config.get("kv_role") if isinstance(sub_config, dict) else None,
+ )
+ for index, sub_config in enumerate(sub_configs)
+ }
+ return topology == {
+ ("mooncake", "kv_producer"),
+ ("lmcache_offload", "offload"),
+ }
+
+
+def _replicated_index_cache_unsupported_reasons(
+ config,
+ hf_config,
+ *,
+ dcp_world_size: int,
+ mla_page_size: int,
+ pcp_world_size: int,
+) -> list[str]:
+ """Return startup blockers for the experimental replicated index layout."""
+ unsupported = []
+ if getattr(hf_config, "model_type", None) != "glm_moe_dsa":
+ unsupported.append("model_type must be glm_moe_dsa")
+ if dcp_world_size <= 1:
+ unsupported.append("decode context parallel size must be > 1")
+ if mla_page_size != 1:
+ unsupported.append("MLA page size must be 1")
+ if pcp_world_size > 1:
+ unsupported.append("PCP is not supported")
+ if getattr(config, "pipeline_parallel_size", 1) > 1:
+ unsupported.append("pipeline parallelism is not supported")
+ if getattr(config, "kv_transfer_config", None) and not (
+ _replicated_index_cache_transfer_supported(config)
+ ):
+ unsupported.append(
+ "KV transfer must be standalone LMCache, standalone Mooncake, or "
+ "multi[Mooncake producer + LMCache offload]"
+ )
+ if getattr(config, "enable_rapidserve", False):
+ unsupported.append("RapidServe disaggregation is not supported")
+ return unsupported
+
+
class AiterMLABackend(AttentionBackend):
@staticmethod
def get_name() -> str:
@@ -264,6 +343,47 @@ def __init__(self, model_runner):
self.dcp_world_size = get_dcp_world_size()
self.dcp_rank = get_dcp_rank()
+ self.replicate_index_cache = dcp_replicated_index_cache_enabled(config)
+ if envs.ATOM_DCP_REPLICATE_INDEX_CACHE:
+ unsupported = _replicated_index_cache_unsupported_reasons(
+ config,
+ hf_config,
+ dcp_world_size=self.dcp_world_size,
+ mla_page_size=self.block_size,
+ pcp_world_size=get_pcp_world_size(),
+ )
+ if unsupported:
+ raise ValueError(
+ "ATOM_DCP_REPLICATE_INDEX_CACHE=1 is unsupported: "
+ + "; ".join(unsupported)
+ )
+
+ if self.replicate_index_cache:
+ # The cache rows themselves come from _index_cache_layout(); this is
+ # only the startup guard that the schedule it reads is well formed,
+ # since a missing or short indexer_types silently degrades to
+ # one-row-per-layer and the replicated page width would then be
+ # applied to rows that hold no indexer.
+ indexer_types = getattr(hf_config, "indexer_types", None)
+ if not indexer_types or len(indexer_types) != hf_config.num_hidden_layers:
+ raise ValueError(
+ "Replicated GLM index cache requires one indexer_types entry "
+ "per target layer"
+ )
+ num_full_layers = sum(
+ 1 for indexer_type in indexer_types if indexer_type == "full"
+ )
+ if not num_full_layers:
+ raise ValueError(
+ "Replicated GLM index cache found no full IndexShare layers"
+ )
+ logger.info(
+ "Replicating GLM index cache for %d full IndexShare layers "
+ "across %d DCP ranks",
+ num_full_layers,
+ self.dcp_world_size,
+ )
+
self._publishes_dcp_local_lens = self.is_sparse and self.dcp_world_size > 1
self._tbo_full_running_bs = 0
@@ -281,7 +401,6 @@ def __init__(self, model_runner):
and self.dcp_world_size > 1
and dcp_persistent
and self.block_size == 1
- and config.speculative_config is None
)
if self.dcp_world_size > 1 and dcp_persistent:
self.persistent_num_heads = mla_dcp_kernel_num_heads(
@@ -317,6 +436,7 @@ def __init__(self, model_runner):
**_mla_seg_meta_kwargs(),
)
i32_kwargs = {"dtype": torch.int32, "device": self.device}
+ i64_kwargs = {"dtype": torch.int64, "device": self.device}
mla_metadata = {
# AITER MLA specific persistent buffers
@@ -362,6 +482,14 @@ def __init__(self, model_runner):
mla_metadata["kv_last_page_lens"].cpu.fill_(1)
mla_metadata["kv_last_page_lens"].copy_to_gpu()
if self.is_sparse:
+ if self.replicate_index_cache:
+ mla_metadata["index_slot_mapping"] = CpuGpuBuffer(
+ self.max_num_batched_tokens, **i64_kwargs
+ )
+ # -1 is the aiter cache kernels' skip sentinel; block 0 is a
+ # real allocatable block, so an unrefreshed row must not carry it.
+ mla_metadata["index_slot_mapping"].np.fill(-1)
+ mla_metadata["index_slot_mapping"].copy_to_gpu()
mla_metadata["cu_seqlen_ke"] = CpuGpuBuffer(
self.max_num_batched_tokens, **i32_kwargs
)
@@ -389,8 +517,8 @@ def __init__(self, model_runner):
device=self.device,
)
# DCP sparse decode compacts each rank's owned top-k slots to the
- # front (no -1 holes), so the per-request region length becomes data-
- # AND layer-dependent.
+ # front (no -1 holes), so the per-query-token region length becomes
+ # data- AND layer-dependent.
self._dcp_sparse_kv_indptr_gpu = torch.zeros(
self.max_num_batched_tokens + 1,
dtype=torch.int32,
@@ -445,6 +573,14 @@ def __init__(self, model_runner):
# Allocate a second set of persistent work buffers for sparse MTP
# per-token layout: max_bs*max_seqlen_qo virtual seqs, each q_len=1.
smt_max_bs = self.max_bs * max_seqlen_qo
+ # Same widening as sparse prefill: when the rebuild is live these
+ # descriptors are regenerated for the DCP-gathered query width, so
+ # they must be sized for it and not for a single rank's heads.
+ sparse_mtp_num_heads = (
+ self.persistent_num_heads
+ if self.sparse_dcp_metadata_rebuild
+ else self.padded_num_attention_heads
+ )
(
(smt_wmd_size, smt_wmd_type),
(smt_wi_size, smt_wi_type),
@@ -455,7 +591,7 @@ def __init__(self, model_runner):
) = get_mla_metadata_info_v1(
smt_max_bs,
1, # max_seqlen_qo=1 for per-token
- self.padded_num_attention_heads,
+ sparse_mtp_num_heads,
self.dtype_q,
self.dtype_kv,
is_sparse=True,
@@ -605,6 +741,13 @@ def _allocate_ubatch_buffers(
ub_max_bs * max_seqlen_qo,
**i64_kwargs,
)
+ if self.replicate_index_cache:
+ var[f"{p}index_slot_mapping"] = CpuGpuBuffer(
+ ub_max_bs * max_seqlen_qo,
+ **i64_kwargs,
+ )
+ var[f"{p}index_slot_mapping"].np.fill(-1)
+ var[f"{p}index_slot_mapping"].copy_to_gpu()
var[f"{p}block_tables"] = CpuGpuBuffer(
ub_max_bs, self.block_table_cols, **i32_kwargs
)
@@ -624,6 +767,20 @@ def _allocate_ubatch_buffers(
ub_max_bs + 1,
**i32_kwargs,
)
+ # Owning request of each query token, numbered within the ubatch:
+ # the DCP top-k filter uses it to index this ubatch's block_tables.
+ # Refreshed per step in _prepare_ubatch_decode; seeded here so a
+ # CUDAGraph capture before the first real step sees a valid map.
+ var[f"{p}token_to_seq_idxs"] = CpuGpuBuffer(
+ ub_max_bs * max_seqlen_qo,
+ **i32_kwargs,
+ )
+ var[f"{p}token_to_seq_idxs"].cpu.copy_(
+ torch.arange(ub_max_bs, dtype=torch.int32).repeat_interleave(
+ max_seqlen_qo
+ )
+ )
+ var[f"{p}token_to_seq_idxs"].copy_to_gpu()
# MLA work buffers per ubatch (GPU only)
var[f"{p}work_meta_data"] = torch.empty(
@@ -951,6 +1108,9 @@ def prepare_mtp_decode(
sparse_decode=True,
)
result["sparse_kv_indptr"] = sparse_kv_indptr
+ index_slots = self.rebuild_draft_index_slots(bs, running_bs)
+ if index_slots is not None:
+ result["index_slot_mapping"] = index_slots
else:
# `bs`, not `running_bs`, and paired with `num_reject_tokens`: this count
# becomes `cu_num`, and the update kernel loads
@@ -982,9 +1142,14 @@ def sub_pool_specs(self) -> list[SubPoolSpec]:
index_dim = hf_config.index_head_dim + 4
aligned_index_dim = ((index_dim + 15) // 16) * 16
index_cache_layer_ids, _ = self._index_cache_layout()
+ replicate_index_cache = getattr(self, "replicate_index_cache", False)
+ index_page_factor = (
+ getattr(self, "dcp_world_size", 1) if replicate_index_cache else 1
+ )
block_bytes += (
len(index_cache_layer_ids)
* runner.block_size
+ * index_page_factor
* aligned_index_dim
* dtypes.fp8.itemsize
)
@@ -1020,6 +1185,10 @@ def allocate_kv_cache_tensors(
index_dim = hf_config.index_head_dim + 4
aligned = ((index_dim + 15) // 16) * 16
index_cache_layer_ids, _ = self._index_cache_layout()
+ replicate_index_cache = getattr(self, "replicate_index_cache", False)
+ index_page_factor = (
+ getattr(self, "dcp_world_size", 1) if replicate_index_cache else 1
+ )
out["aligned_index_dim"] = aligned
out["index_cache_layer_ids"] = index_cache_layer_ids
out["index_cache_layer_map"] = {
@@ -1031,7 +1200,7 @@ def allocate_kv_cache_tensors(
out["index_cache"] = torch.zeros(
len(index_cache_layer_ids),
runner.num_physical_kvcache_blocks,
- runner.physical_block_size,
+ runner.physical_block_size * index_page_factor,
aligned,
dtype=dtypes.fp8,
device="cuda",
@@ -1063,6 +1232,10 @@ def build_kv_cache_tensor(self, layer_id: int, module):
module.max_model_len = runner.config.max_model_len
index_cache = None
if runner.is_deepseek_v32 and module.indexer is not None:
+ replicate_index_cache = getattr(self, "replicate_index_cache", False)
+ index_page_factor = (
+ getattr(self, "dcp_world_size", 1) if replicate_index_cache else 1
+ )
# `layer_id` is a PP-local cache-row counter, while the compact map
# is keyed by global model layer IDs. On a non-first PP stage they
# differ (for example local 0 may be global 39), so use layer_num
@@ -1077,7 +1250,9 @@ def build_kv_cache_tensor(self, layer_id: int, module):
index_cache = runner.index_cache[index_cache_layer_id]
# Use aligned dimension to avoid memory copy in torch inductor
module.indexer.k_cache.kv_cache[0] = index_cache.view(
- runner.num_physical_kvcache_blocks * runner.physical_block_size,
+ runner.num_physical_kvcache_blocks
+ * runner.physical_block_size
+ * index_page_factor,
1,
runner.aligned_index_dim,
)
@@ -1093,6 +1268,8 @@ def build_kv_cache_tensor(self, layer_id: int, module):
def get_kv_transfer_tensors(self):
from atom.kv_transfer.disaggregation.types import (
+ INDEX_CACHE_ROLE,
+ MLA_KV_ROLE,
KVTransferRegion,
KVTransferTensors,
)
@@ -1111,18 +1288,52 @@ def get_kv_transfer_tensors(self):
base_addr=t.data_ptr(),
total_bytes=t.numel() * t.element_size(),
unit_bytes=bpb,
+ semantic_role=MLA_KV_ROLE,
)
)
if hasattr(runner, "index_cache"):
+ index_head_dim = runner.config.hf_config.index_head_dim
+ if getattr(self, "replicate_index_cache", False):
+ expected_page_width = runner.physical_block_size * self.dcp_world_size
+ index_shape = tuple(runner.index_cache.shape)
+ if len(index_shape) < 3 or index_shape[2] != expected_page_width:
+ raise RuntimeError(
+ "Replicated MLA index page width mismatch: "
+ f"expected {expected_page_width}, got shape={index_shape}"
+ )
for layer_id in range(runner.index_cache.shape[0]):
t = runner.index_cache[layer_id]
bpb = t.stride(0) * t.element_size() * self.block_ratio
+ # One token occupies aligned_index_dim elements of an index row.
+ index_row_bytes = runner.aligned_index_dim * t.element_size()
+ if index_row_bytes <= 0 or bpb % index_row_bytes:
+ raise RuntimeError(
+ f"MLA index page of {bpb} bytes is not a whole number of "
+ f"{index_row_bytes}-byte token rows "
+ f"(aligned_index_dim={runner.aligned_index_dim}). The plane "
+ "sizes below would not match the real layout, and the "
+ "Mooncake relayout addresses the page through them."
+ )
+ tokens_per_page = bpb // index_row_bytes
+ # One fp32 scale per token holds only while the indexer
+ # quantizes a key row as a single block (Indexer.
+ # quant_block_size in deepseek_v2.py).
+ if index_head_dim != 128:
+ raise RuntimeError(
+ f"An index row of {index_head_dim} key bytes does not "
+ "match the indexer's 128-byte quantization block, so a "
+ "token's scales are not the single fp32 the key/scale "
+ "plane split assumes."
+ )
block_regions.append(
KVTransferRegion(
base_addr=t.data_ptr(),
total_bytes=t.numel() * t.element_size(),
unit_bytes=bpb,
+ semantic_role=INDEX_CACHE_ROLE,
+ key_plane_bytes=tokens_per_page * index_head_dim,
+ scale_plane_bytes=tokens_per_page * 4,
)
)
@@ -1251,6 +1462,28 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int):
bs = batch.total_seqs_num_prefill
sum_scheduled_tokens = batch.total_tokens_num_prefill
var = self.model_runner.forward_vars
+ if self.replicate_index_cache:
+ if batch.is_dummy_run:
+ var["index_slot_mapping"].np[:sum_scheduled_tokens] = -1
+ else:
+ index_slots = [
+ self._replicated_index_slot(block_table, pos)
+ for block_table, cached_len, seq_len in zip(
+ batch.block_tables,
+ batch.num_cached_tokens,
+ batch.context_lens,
+ )
+ for pos in range(cached_len, seq_len)
+ ]
+ if len(index_slots) != sum_scheduled_tokens:
+ raise RuntimeError(
+ "Replicated index slot count does not match prefill token "
+ f"count: {len(index_slots)} != {sum_scheduled_tokens}"
+ )
+ var["index_slot_mapping"].np[:sum_scheduled_tokens] = index_slots
+ attn_metadata.index_slot_mapping = var["index_slot_mapping"].copy_to_gpu(
+ sum_scheduled_tokens
+ )
if self.is_sparse and attn_metadata.max_seqlen_k > self.index_topk:
if attn_metadata.block_tables is None:
self.prepare_block_tables(batch)
@@ -1313,7 +1546,7 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int):
attn_metadata.sparse_kv_indptr = var["sparse_kv_indptr"].copy_to_gpu(
sum_scheduled_tokens + 1
)
- if self.dcp_world_size > 1:
+ if self.dcp_world_size > 1 and not self.replicate_index_cache:
self._build_dcp_indexer_prefill_meta(attn_metadata, bs, counts, var)
get_mla_metadata_v1(
attn_metadata.sparse_cu_seqlens_q,
@@ -1763,6 +1996,68 @@ def _dcp_round_robin_slot(self, block_table, pos: int) -> int:
+ dcp_local_index(pos, W, S) % block_size
)
+ def _replicated_index_slot(self, block_table, pos: int) -> int:
+ """Physical slot for a global token in an expanded replicated page.
+
+ Scheduler block IDs and lifetimes stay unchanged. Each scheduler block's
+ index-cache page is widened from ``block_size`` local tokens to
+ ``block_size * dcp_world_size`` global tokens.
+ """
+ expanded_page_size = self.model_runner.block_size * self.dcp_world_size
+ return (
+ block_table[pos // expanded_page_size] * expanded_page_size
+ + pos % expanded_page_size
+ )
+
+ def rebuild_draft_index_slots(self, scheduled_bs: int, running_bs: int):
+ """Replicated index slots for an MTP draft step: one row per sequence.
+
+ The target step publishes one entry per verify token; a draft step runs
+ one row per sequence at the advanced ``context_lens``, so the mapping has
+ to be rebuilt rather than sliced. No owner filter -- the index cache is
+ replicated, so every rank writes every token. Padded rows get the skip
+ sentinel. Returns the rebuilt view, or None when replication is off.
+ """
+ if not self.replicate_index_cache:
+ return None
+ var = self.model_runner.forward_vars
+ expanded_page_size = self.model_runner.block_size * self.dcp_world_size
+ # prepare_decode zeroes a padded row's context_len, so clamp before the
+ # gather: the sentinel below lands after it, too late to keep pos=-1 out
+ # of the block table.
+ pos = (var["context_lens"].gpu[:running_bs].to(torch.int64) - 1).clamp_(min=0)
+ page = (
+ var["block_tables"]
+ .gpu[:running_bs]
+ .gather(1, (pos // expanded_page_size).unsqueeze(1))
+ .squeeze(1)
+ )
+ slots = page.to(torch.int64) * expanded_page_size + pos % expanded_page_size
+ if running_bs > scheduled_bs:
+ slots[scheduled_bs:] = -1
+ out = var["index_slot_mapping"].gpu[:running_bs]
+ out.copy_(slots)
+ return out
+
+ def rebuild_draft_token_to_seq_idxs(self, running_bs: int):
+ """Owning request per query token for an MTP draft step.
+
+ The DCP top-k filter reads this to pick a row of ``block_tables``. The
+ target step publishes one entry per verify token
+ (``arange(bs).repeat_interleave(max_seqlen_q)``); a draft step runs one
+ row per sequence, which makes the map the identity. Padded rows carry no
+ top-k, but they still index the table, so cover ``running_bs``.
+
+ Returns the rebuilt view, or None when the model is not sparse (no
+ buffer, and no filter to read it).
+ """
+ if not self.is_sparse:
+ return None
+ self._token_to_seq_idxs_gpu[:running_bs] = torch.arange(
+ running_bs, dtype=torch.int32, device=self.device
+ )
+ return self._token_to_seq_idxs_gpu[:running_bs]
+
def prepare_decode(
self,
batch: ScheduledBatch,
@@ -1828,6 +2123,14 @@ def prepare_decode(
var["positions"].np[:sum_scheduled_tokens] = positions
var["context_lens"].np[:scheduled_bs] = context_lens
var["context_lens"].np[scheduled_bs:running_bs] = 0
+ if self.replicate_index_cache:
+ var["index_slot_mapping"].np[:running_tokens] = -1
+ if not batch.is_dummy_run:
+ var["index_slot_mapping"].np[:sum_scheduled_tokens] = [
+ self._replicated_index_slot(block_table, pos)
+ for block_table, seq_len in zip(block_tables, context_lens)
+ for pos in range(int(seq_len) - max_seqlen_q, int(seq_len))
+ ]
if self.dcp_world_size > 1:
from atom.model_ops.dcp_ops import get_dcp_local_seq_lens
@@ -1982,6 +2285,11 @@ def prepare_decode(
is_sparse_mtp = self.is_sparse and max_seqlen_q > 1
# metadata copies on main stream
positions = var["positions"].copy_to_gpu(sum_scheduled_tokens)
+ index_slot_mapping = (
+ var["index_slot_mapping"].copy_to_gpu(running_tokens)
+ if self.replicate_index_cache
+ else None
+ )
ctx.update({el: var[el].copy_to_gpu(num) for el, num in vars_for_metadata})
if is_sparse_mtp:
@@ -2021,6 +2329,8 @@ def prepare_decode(
**ctx,
)
attn_metadata.dtype_q = self.dtype_q
+ if self.replicate_index_cache:
+ attn_metadata.index_slot_mapping = index_slot_mapping
# Round-robin CP global kv_indptr (only under DCP; None otherwise so the
# non-DCP / qlen=1 paths keep the plain kernel). Consumed by
@@ -2057,6 +2367,15 @@ def prepare_decode(
attn_metadata.sparse_kv_last_page_lens = var[
"sparse_kv_last_page_lens"
].gpu[:running_bs]
+ if self.dcp_world_size > 1:
+ # One token per request makes the filter's token -> request map
+ # the identity; it must still exist, since MTP shares the filter.
+ self._token_to_seq_idxs_gpu[:running_bs] = torch.arange(
+ running_bs, dtype=torch.int32, device=self.device
+ )
+ attn_metadata.token_to_seq_idxs = self._token_to_seq_idxs_gpu[
+ :running_bs
+ ]
# running_bs, not scheduled_bs: the padded rows have to be split into the
# ubatches too, or accuracy drifts.
@@ -2115,6 +2434,15 @@ def _prepare_ubatch_decode(
tok_start : tok_start + ub_real_tokens
]
var[f"{p}slot_mapping"].np[ub_real_tokens:padded_tok_count] = -1
+ if self.replicate_index_cache:
+ var[f"{p}index_slot_mapping"].np[:ub_real_tokens] = var[
+ "index_slot_mapping"
+ ].np[tok_start : tok_start + ub_real_tokens]
+ var[f"{p}index_slot_mapping"].np[ub_real_tokens:padded_tok_count] = -1
+ if self.is_sparse:
+ var[f"{p}token_to_seq_idxs"].np[:padded_tok_count] = np.repeat(
+ np.arange(running_bs, dtype=np.int32), max_seqlen_q
+ )
var[f"{p}block_tables"].np[:ub_real_reqs] = var["block_tables"].np[
req_start : req_start + ub_real_reqs
@@ -2186,6 +2514,9 @@ def _prepare_ubatch_decode(
vars_used.append((f"{p}g_kv_indptr", running_bs + 1))
if self.is_sparse:
vars_used.append((f"{p}sparse_kv_indptr", running_bs + 1))
+ vars_used.append((f"{p}token_to_seq_idxs", padded_tok_count))
+ if self.replicate_index_cache:
+ vars_used.append((f"{p}index_slot_mapping", padded_tok_count))
for el, num in vars_used:
var[el].copy_to_gpu(num)
@@ -2328,6 +2659,10 @@ def build_for_cudagraph_capture(self, bs: int) -> AttentionMetaData:
**ctx_mla_ps,
)
attn_matadata.dtype_q = self.dtype_q
+ if self.replicate_index_cache:
+ attn_matadata.index_slot_mapping = var["index_slot_mapping"].gpu[
+ :sum_tokens
+ ]
# Attach the round-robin CP global kv_indptr for the captured graph so
# replay (which overwrites the buffer with real values) matches. Only
# consumed by _forward_decode when dcp>1 and max_q_len>1.
@@ -2357,6 +2692,13 @@ def build_for_cudagraph_capture(self, bs: int) -> AttentionMetaData:
attn_matadata.sparse_kv_last_page_lens = var[
"sparse_kv_last_page_lens"
].gpu[:bs]
+ if self.dcp_world_size > 1:
+ # Same identity map prepare_decode builds; the captured graph
+ # needs the buffer to exist.
+ self._token_to_seq_idxs_gpu[:bs] = torch.arange(
+ bs, dtype=torch.int32, device=self.device
+ )
+ attn_matadata.token_to_seq_idxs = self._token_to_seq_idxs_gpu[:bs]
positions = var["positions"].copy_to_gpu(sum_tokens)
context = Context(
positions=positions,
@@ -2422,6 +2764,14 @@ def build_ubatch_metadata(
reduce_partial_map=var[f"{p}reduce_partial_map"],
)
attn.dtype_q = self.dtype_q
+ if self.is_sparse:
+ attn.token_to_seq_idxs = var[f"{p}token_to_seq_idxs"].gpu[
+ : running_bs * max_q_len
+ ]
+ if self.replicate_index_cache:
+ attn.index_slot_mapping = var[f"{p}index_slot_mapping"].gpu[
+ : running_bs * max_q_len
+ ]
# Per-ubatch round-robin CP global kv_indptr (None when non-DCP). Consumed
# by _forward_decode when dcp>1 and max_q_len>1 (MTP).
attn.g_kv_indptr = (
@@ -2477,6 +2827,12 @@ def build_ubatch_prefill_metadata(
):
ub_attn.token_to_seq_idxs = attn_metadata.token_to_seq_idxs[ts] - req_start
+ if getattr(attn_metadata, "index_slot_mapping", None) is not None:
+ # Absolute cache addresses, so a token slice needs no rebase. Its
+ # presence is also what selects the replicated index layout in the
+ # indexer -- dropping it would silently fall back to the sharded one.
+ ub_attn.index_slot_mapping = attn_metadata.index_slot_mapping[ts]
+
total_tokens = (
attn_metadata.slot_mapping.shape[0]
if attn_metadata.slot_mapping is not None
diff --git a/atom/model_ops/dcp_ops.py b/atom/model_ops/dcp_ops.py
index 4344fc35dd..1cb1bde2f2 100644
--- a/atom/model_ops/dcp_ops.py
+++ b/atom/model_ops/dcp_ops.py
@@ -12,8 +12,8 @@
They are mathematically equivalent but not bitwise identical; see the A2A
section below for why one collective can replace two.
-Also here: the sparse indexer's DCP support -- the decode candidate exchange
-(fused into one aiter op) and the sparse-prefill owned-slot filter.
+Also here: the sparse indexer's DCP support -- the fused decode candidate
+exchange plus owned-slot filters for replicated-cache decode and sparse prefill.
"""
import numpy as np
@@ -936,15 +936,305 @@ def dcp_decode_candidate_exchange_fused(
# ---------------------------------------------------------------------------
-# DCP sparse index filter + round-robin localize (sparse PREFILL).
+# DCP sparse index filter + round-robin localize.
# Two-pass compacting filter: count this rank's owned top-k, then pack the
# owned slots to the front of each region (no -1 holes -- see cp_lse_ag_out_rs).
#
-# Decode has no twin here any more: aiter's flydsl_dcp_topk_merge emits this
-# rank's owned slots directly, so nothing is left to filter afterwards.
+# Sharded-cache decode uses flydsl_dcp_topk_merge and needs no filter. The
+# decode kernel below remains for replicated index caches, where every rank
+# scores the full context but attention still consumes a sharded main KV cache.
# ---------------------------------------------------------------------------
+@triton.jit
+def _count_owned_dcp_kernel(
+ token_indices_ptr, # int32 [num_tokens, NUM_TOPK_TOKENS] -- GLOBAL top-k positions
+ out_counts, # int32 [num_tokens] -- owned top-k count per query token
+ out_metadata_counts, # int32 [num_tokens] -- max(out_counts, 1)
+ DCP_RANK: tl.constexpr,
+ DCP_WORLD: tl.constexpr,
+ INTERLEAVE: tl.constexpr, # cp_kv_cache_interleave_size S (1 = round-robin)
+ NUM_TOPK_TOKENS: tl.constexpr,
+ BLOCK_N: tl.constexpr,
+ ti_stride0,
+ ti_stride1,
+):
+ """Pass 1 of the compacting DCP filter: how many of the global top-k
+ positions does this rank own, per QUERY TOKEN? Its exclusive cumsum gives
+ the compacted output offsets used by ``_compact_filter_dcp_kernel``.
+
+ The row unit is a query token, not a request: MTP verify forwards
+ max_seqlen_q draft positions per request and each one carries its own
+ top-k. At qlen==1 the two coincide. Owner of global position g is rank
+ (g//S)%W (S=INTERLEAVE; S=1 -> g%W).
+ """
+ token_id = tl.program_id(0)
+
+ count = 0
+ for tile_start in range(0, NUM_TOPK_TOKENS, BLOCK_N):
+ indice_id = tile_start + tl.arange(0, BLOCK_N)
+ col_valid = indice_id < NUM_TOPK_TOKENS
+ ti_ptr = token_indices_ptr + token_id * ti_stride0 + indice_id * ti_stride1
+ tok = tl.load(ti_ptr, mask=col_valid, other=-1)
+ owned = col_valid & (tok >= 0) & (((tok // INTERLEAVE) % DCP_WORLD) == DCP_RANK)
+ count += tl.sum(owned.to(tl.int32))
+
+ tl.store(out_counts + token_id, count)
+ tl.store(out_metadata_counts + token_id, tl.maximum(count, 1))
+
+
+@triton.jit
+def _compact_filter_dcp_kernel(
+ token_to_seq_idxs, # int32 [num_tokens] -- owning request of each query token
+ out_kv_indptr, # int32 [num_tokens + 1] -- COMPACTED offsets (cumsum of pass 1)
+ block_table, # int32 [num_req, max_num_blocks_per_req] -- logical(global) blocks
+ token_indices_ptr, # int32 [num_tokens, NUM_TOPK_TOKENS] -- GLOBAL top-k positions
+ out_kv_indices, # int32 [>= out_kv_indptr[-1]]
+ DCP_RANK: tl.constexpr,
+ DCP_WORLD: tl.constexpr,
+ INTERLEAVE: tl.constexpr, # cp_kv_cache_interleave_size S (1 = round-robin)
+ PAGE_SIZE: tl.constexpr, # runner (physical) block size
+ NUM_TOPK_TOKENS: tl.constexpr,
+ BLOCK_N: tl.constexpr,
+ ti_stride0,
+ ti_stride1,
+ bt_stride0: tl.int64,
+ bt_stride1: tl.constexpr,
+):
+ # DCP interleave-S: a GLOBAL position g is owned by rank ``(g // S) % W``; on
+ # the owner rank its physical slot follows the virtual-block layout used by
+ # _dcp_round_robin_slot / ATOM PR #847 (S=1 -> the original round-robin):
+ # vbs = PAGE_SIZE * W
+ # vb = g % vbs
+ # slot = block_table[req, g // vbs] * PAGE_SIZE
+ # + (vb // (W*S)) * S + (vb % S)
+ # token_indices holds GLOBAL positions (the indexer scored the full sequence
+ # via all-gathered logits). This rank keeps ONLY the positions it owns and
+ # writes them COMPACTED to the front of its region -- no -1 holes. Holes are
+ # exactly what breaks aiter's lse path (immediate fault on the persistent
+ # kernel, silently unwritten lse on the split-KV one).
+ #
+ # Compaction is order-preserving (tl.cumsum within a tile plus a running
+ # offset across tiles) rather than atomic-allocated like vLLM, so the KV
+ # order -- and hence the floating-point accumulation order -- is
+ # deterministic run to run, which the dcp=1 vs dcp=N comparison relies on.
+ #
+ # NOTE: the slot is computed from block_table directly (like vLLM) rather
+ # than gathered from a precomputed kv_indices -- the DCP round-robin
+ # per-token slot array does not exist on the sparse path (dense reads go
+ # through block_tables in-kernel).
+ token_id = tl.program_id(0)
+
+ out_kv_start = tl.load(out_kv_indptr + token_id)
+ # The block table is per request, so a draft position looks up its owner.
+ req_id = tl.load(token_to_seq_idxs + token_id)
+
+ vbs = PAGE_SIZE * DCP_WORLD
+ written = 0
+ for tile_start in range(0, NUM_TOPK_TOKENS, BLOCK_N):
+ indice_id = tile_start + tl.arange(0, BLOCK_N)
+ # Full top-k width; `tok >= 0` is the only valid-id guard. Must stay in
+ # lock-step with `_count_owned_dcp_kernel` (see the note there on why the
+ # old `indice_id < g_kv_len` mask is wrong) -- if the two disagree, the
+ # counted offsets and the written entries diverge.
+ col_valid = indice_id < NUM_TOPK_TOKENS
+
+ ti_ptr = token_indices_ptr + token_id * ti_stride0 + indice_id * ti_stride1
+ tok = tl.load(ti_ptr, mask=col_valid, other=-1) # GLOBAL position
+
+ idx_valid = (
+ col_valid & (tok >= 0) & (((tok // INTERLEAVE) % DCP_WORLD) == DCP_RANK)
+ )
+
+ block_id = tok // vbs
+ vb = tok % vbs
+ inblock_offset = (vb // (DCP_WORLD * INTERLEAVE)) * INTERLEAVE + (
+ vb % INTERLEAVE
+ )
+ physical_block = tl.load(
+ block_table + req_id * bt_stride0 + block_id * bt_stride1,
+ mask=idx_valid,
+ other=0,
+ )
+ slot = physical_block * PAGE_SIZE + inblock_offset
+
+ # Exclusive prefix sum of the owned mask -> destination inside this tile.
+ owned_i32 = idx_valid.to(tl.int32)
+ dst = written + tl.cumsum(owned_i32, axis=0) - owned_i32
+ tl.store(out_kv_indices + out_kv_start + dst, slot, mask=idx_valid)
+ written += tl.sum(owned_i32)
+
+ # Keep persistent fast-mode metadata valid for a rank that owns no selected
+ # KV. The A2A pack kernel neutralizes this dummy row's LSE without adding a
+ # separate pointwise launch.
+ tl.store(out_kv_indices + out_kv_start, 0, mask=written == 0)
+
+
+def _check_dcp_filter_rows(
+ num_tokens: int,
+ *,
+ token_to_seq_idxs: torch.Tensor,
+ topk_indices: torch.Tensor,
+ out_kv_indptr: torch.Tensor,
+ owned_counts: torch.Tensor,
+ block_table: torch.Tensor,
+) -> None:
+ """Check every tensor the ``(num_tokens,)`` grid indexes is long enough.
+
+ The kernels take raw pointers, and the slices built here
+ (``owned_counts[:num_tokens]``, ``out_kv_indptr[1 : num_tokens + 1]``)
+ truncate silently when the buffer is short -- the launch then reads and
+ writes past the end with no Python-side error. These buffers are sized
+ ``max_num_batched_tokens`` while ``num_tokens`` comes from the caller's
+ token count, so a padding or ubatch change is exactly what would break the
+ relation.
+
+ Raised rather than asserted because `python -O` strips asserts, which would
+ put the overrun back and leave it silent.
+ """
+ if num_tokens < 0:
+ raise ValueError(f"num_tokens must be non-negative, got {num_tokens}")
+ for name, tensor, needed in (
+ ("token_to_seq_idxs", token_to_seq_idxs, num_tokens),
+ ("topk_indices", topk_indices, num_tokens),
+ ("out_kv_indptr", out_kv_indptr, num_tokens + 1),
+ ("owned_counts", owned_counts, num_tokens),
+ ):
+ have = tensor.shape[0]
+ if have < needed:
+ raise ValueError(
+ f"{name} holds {have} rows but the filter runs {num_tokens} "
+ f"query tokens and needs {needed}"
+ )
+ for name, tensor in (
+ ("token_to_seq_idxs", token_to_seq_idxs),
+ ("out_kv_indptr", out_kv_indptr),
+ ("owned_counts", owned_counts),
+ ):
+ if tensor.dtype != torch.int32:
+ raise TypeError(f"{name} must be int32, got {tensor.dtype}")
+ if block_table.dim() != 2:
+ raise ValueError(
+ "block_table must be [num_requests, max_blocks_per_request]; "
+ f"got shape {tuple(block_table.shape)}"
+ )
+
+
+def triton_filter_and_convert_dcp_index(
+ token_to_seq_idxs: torch.Tensor, # int32 [num_tokens] owning request per token
+ num_tokens: int,
+ block_table: torch.Tensor, # int32 [num_req, max_num_blocks_per_req] logical
+ token_indices: torch.Tensor, # int32 [num_tokens, NUM_TOPK_TOKENS] GLOBAL pos
+ dcp_rank: int,
+ dcp_world_size: int,
+ block_size: int, # runner (physical) block size == PAGE_SIZE
+ out_kv_indptr: torch.Tensor, # int32 [num_tokens + 1] COMPACTED, written here
+ owned_counts: torch.Tensor, # int32 [>= num_tokens] scratch for pass 1
+ NUM_TOPK_TOKENS: int = 2048,
+ BLOCK_N: int = 128,
+ out: torch.Tensor | None = None,
+ cp_kv_cache_interleave_size: int = 1,
+):
+ """DCP (interleave-S) filter + localize of global top-k positions,
+ **compacting** each rank's owned slots to the front of its region.
+
+ ``token_indices[token_id, indice_id]`` is a GLOBAL token position selected by
+ the indexer (scored over the full sequence via all-gathered logits). This
+ rank keeps a position ``g`` only if ``g % W == dcp_rank`` and maps it to its
+ physical slot via the round-robin (virtual-block) layout, computed directly
+ from ``block_table`` (like vLLM):
+ vbs = block_size * W
+ slot = block_table[req, g // vbs] * block_size + (g % vbs) // W
+
+ Non-owned positions are **dropped**, not marked: the kept slots are packed
+ contiguously (original top-k order preserved) and ``out_kv_indptr`` is
+ rewritten to the resulting per-query-token lengths. This replaces the earlier
+ "fixed length + -1 sentinel" layout, whose holes broke aiter's lse output.
+ Because the kept count depends on the per-layer top-k selection,
+ ``out_kv_indptr`` is layer-dependent. Sparse+DCP persistent mode therefore
+ rebuilds its work metadata after each full IndexShare layer and reuses that
+ plan in the following shared layers.
+
+ The 8 ranks' kept sets are disjoint and their union is exactly the global
+ top-k, which is what makes the downstream ``cp_lse_ag_out_rs`` merge valid.
+ """
+ assert token_indices.dtype == torch.int32
+ assert token_indices.shape[1] == NUM_TOPK_TOKENS
+ assert NUM_TOPK_TOKENS % BLOCK_N == 0, (
+ f"NUM_TOPK_TOKENS ({NUM_TOPK_TOKENS}) must be divisible by"
+ f"BLOCK_N ({BLOCK_N})"
+ )
+ assert 0 <= dcp_rank < dcp_world_size
+ assert out is not None, "sparse_kv_indices_buffer (out) is required"
+ _check_dcp_filter_rows(
+ num_tokens,
+ token_to_seq_idxs=token_to_seq_idxs,
+ topk_indices=token_indices,
+ out_kv_indptr=out_kv_indptr,
+ owned_counts=owned_counts,
+ block_table=block_table,
+ )
+
+ token_to_seq_idxs_c = token_to_seq_idxs.contiguous()
+ block_table_c = block_table.contiguous()
+ token_indices_c = token_indices.contiguous()
+
+ ti_stride0, ti_stride1 = token_indices_c.stride()
+ bt_stride0, bt_stride1 = block_table_c.stride()
+ grid = (num_tokens,)
+
+ # Pass 1: per-query-token count of owned top-k positions.
+ counts = owned_counts[:num_tokens]
+ metadata_counts = out_kv_indptr[1 : num_tokens + 1]
+ _count_owned_dcp_kernel[grid](
+ token_indices_c,
+ counts,
+ metadata_counts,
+ dcp_rank,
+ dcp_world_size,
+ cp_kv_cache_interleave_size,
+ NUM_TOPK_TOKENS,
+ BLOCK_N,
+ ti_stride0,
+ ti_stride1,
+ )
+
+ # Exclusive cumsum -> compacted offsets. Written in place so the caller's
+ # tensor (and anything already holding a view of it) sees the update.
+ # dtype=int32 keeps the accumulation in int32 (torch would promote integral
+ # cumsum to int64 by default, which the kernels' int32 pointers reject).
+ # zero_() rather than `out_kv_indptr[0] = 0`: assigning a Python scalar goes
+ # through a host->device copy, which HIP rejects while a graph is capturing
+ # (hipErrorStreamCaptureUnsupported). Everything here must stay device-side.
+ out_kv_indptr[:1].zero_()
+ torch.cumsum(
+ metadata_counts,
+ dim=0,
+ dtype=torch.int32,
+ out=metadata_counts,
+ )
+
+ # Pass 2: write the owned slots packed to the front of each region.
+ _compact_filter_dcp_kernel[grid](
+ token_to_seq_idxs_c,
+ out_kv_indptr,
+ block_table_c,
+ token_indices_c,
+ out,
+ dcp_rank,
+ dcp_world_size,
+ cp_kv_cache_interleave_size,
+ block_size,
+ NUM_TOPK_TOKENS,
+ BLOCK_N,
+ ti_stride0,
+ ti_stride1,
+ bt_stride0,
+ bt_stride1,
+ )
+ return out
+
+
@triton.jit
def _count_owned_dcp_prefill_kernel(
dsa_kv_indptr, # int32 [num_tokens + 1] -- GLOBAL per-token candidate counts
@@ -1094,10 +1384,9 @@ def triton_filter_and_convert_dcp_index_prefill(
):
"""Filter a sparse-PREFILL top-k down to this rank's owned KV slots.
- Prefill has one row per query token and its ``topk_indices`` are flat KV
- indices rather than within-sequence positions. (The decode path no longer
- needs this at all -- the fused merge emits owned slots directly.) Two passes,
- in-place int32 cumsum, order-preserving compaction, and the same
+ Both key on query tokens; prefill's ``topk_indices`` are flat KV indices
+ rather than within-sequence positions. Everything else -- two passes, in-place int32
+ cumsum, order-preserving compaction -- is identical, and the same
layer-scoped buffers are reused (they are sized ``max_num_batched_tokens``,
which bounds the prefill token count too).
"""
@@ -1111,6 +1400,14 @@ def triton_filter_and_convert_dcp_index_prefill(
assert out is not None, "sparse_kv_indices_buffer (out) is required"
num_tokens = dsa_kv_indptr.shape[0] - 1
+ _check_dcp_filter_rows(
+ num_tokens,
+ token_to_seq_idxs=token_to_seq_idxs,
+ topk_indices=topk_indices,
+ out_kv_indptr=out_kv_indptr,
+ owned_counts=owned_counts,
+ block_table=block_table,
+ )
dsa_kv_indptr_c = dsa_kv_indptr.contiguous()
token_to_seq_idxs_c = token_to_seq_idxs.contiguous()
diff --git a/atom/models/deepseek_mtp.py b/atom/models/deepseek_mtp.py
index 70d1c2fd0c..3fa1a7a4b2 100644
--- a/atom/models/deepseek_mtp.py
+++ b/atom/models/deepseek_mtp.py
@@ -281,41 +281,6 @@ def compute_draft_ids(
mtp_layer = self.layers[str(self.mtp_start_layer_idx + current_step_idx)]
return mtp_layer.shared_head.head.compute_argmax_token(hidden_states, out=out)
- def set_skip_topk(self, skip: bool) -> None:
- """Toggle ``skip_topk`` on MTP sparse-attention layers.
-
- Used by ``EagleProposer`` for ``index_share_for_mtp_iteration``: draft
- step 0 sets ``skip=False`` (compute indexer top-k), steps 1+ set
- ``skip=True`` (reuse step 0's ``sparse_kv_indices_buffer``).
- Matches vLLM ``DeepSeekMultiTokenPredictor.set_skip_topk``.
- """
- for layer in self.layers.values():
- mtp_block = getattr(layer, "mtp_block", None)
- if mtp_block is None:
- continue
- self_attn = getattr(mtp_block, "self_attn", None)
- if self_attn is None or not hasattr(self_attn, "skip_topk"):
- continue
- if getattr(self_attn, "indexer", None) is not None:
- self_attn.skip_topk = skip
-
- def compact_topk_indices(self, slot_ids: torch.Tensor) -> None:
- """Gather sparse top-k rows at ``slot_ids`` to the front of each buffer."""
- num_slots = slot_ids.numel()
- for layer in self.layers.values():
- mtp_block = getattr(layer, "mtp_block", None)
- if mtp_block is None:
- continue
- self_attn = getattr(mtp_block, "self_attn", None)
- if self_attn is None:
- continue
- mla_attn = getattr(self_attn, "mla_attn", None)
- if mla_attn is None:
- continue
- sparse_buf = getattr(mla_attn, "sparse_kv_indices_buffer", None)
- if sparse_buf is not None and sparse_buf.numel() > 0:
- sparse_buf[:num_slots] = sparse_buf[slot_ids]
-
@support_torch_compile
class DeepSeekMTP(nn.Module):
@@ -421,12 +386,6 @@ def compute_draft_ids(
"""
return self.model.compute_draft_ids(hidden_states, spec_step_idx, out=out)
- def set_skip_topk(self, skip: bool) -> None:
- self.model.set_skip_topk(skip)
-
- def compact_topk_indices(self, slot_ids: torch.Tensor) -> None:
- self.model.compact_topk_indices(slot_ids)
-
def get_expert_mapping(self) -> list[tuple[str, str, int, str]]:
# Params for weights, fp8 weight scales, fp8 activation scales
# (param_name, weight_name, expert_id, shard_id)
diff --git a/atom/models/deepseek_v2.py b/atom/models/deepseek_v2.py
index 74dfdda06d..a262a38a28 100644
--- a/atom/models/deepseek_v2.py
+++ b/atom/models/deepseek_v2.py
@@ -24,7 +24,6 @@
"""Inference-only DeepseekV2/DeepseekV3 model."""
import logging
-from typing import Optional, Tuple, Union
import torch
from aiter import (
@@ -90,6 +89,7 @@
from atom.model_ops.base_attention import Attention
from atom.model_ops.dcp_ops import (
dcp_decode_candidate_exchange_fused,
+ triton_filter_and_convert_dcp_index,
triton_filter_and_convert_dcp_index_prefill,
)
from atom.model_ops.embed_head import (
@@ -260,8 +260,8 @@ def increment_version(tensor):
def _enable_non_triton_global_mxfp4_input_norm_quant(
config: PretrainedConfig,
- quant_config: Optional[QuantizationConfig],
- quant_dtype: Optional[torch.dtype],
+ quant_config: QuantizationConfig | None,
+ quant_dtype: torch.dtype | None,
is_mtp_block: bool,
) -> bool:
if (
@@ -323,7 +323,7 @@ def _is_neox_rope_style(
def _can_fuse_indexer_wk_weights_proj(
config: PretrainedConfig,
- quant_config: Optional[QuantizationConfig],
+ quant_config: QuantizationConfig | None,
indexer_prefixes: list[str],
) -> bool:
if not ENABLE_DS_INDEXER_QK_ROPE_CACHE_FUSION:
@@ -416,14 +416,14 @@ def _fuse_rmsnorm_fp4_quant_fake(
x1: torch.Tensor,
x1_weight: torch.Tensor,
x1_epsilon: float,
- x2: Optional[torch.Tensor] = None,
- x2_weight: Optional[torch.Tensor] = None,
- x2_epsilon: Optional[float] = None,
- res1: Optional[torch.Tensor] = None,
+ x2: torch.Tensor | None = None,
+ x2_weight: torch.Tensor | None = None,
+ x2_epsilon: float | None = None,
+ res1: torch.Tensor | None = None,
shuffle: bool = True,
scale_shuffle_padding: bool = True,
output_unquantized_inp1: bool = False,
-) -> Tuple[
+) -> tuple[
torch.Tensor,
torch.Tensor,
torch.Tensor,
@@ -458,7 +458,7 @@ def _fuse_rmsnorm_fp4_quant_fake(
return out1_quantized, out1_bs, out1_unquantized, out2, out_res1
-def _mxfp4_activation_quant_layout(num_tokens: int) -> Tuple[bool, bool]:
+def _mxfp4_activation_quant_layout(num_tokens: int) -> tuple[bool, bool]:
if use_fp4_non_shuffle_triton_gemm():
return False, False
if use_triton_gemm():
@@ -471,16 +471,16 @@ def _fused_rms_fp8_quant_fake(
x1: torch.Tensor,
x1_weight: torch.Tensor,
x1_epsilon: float,
- x2: Optional[torch.Tensor] = None,
- x2_weight: Optional[torch.Tensor] = None,
- x2_epsilon: Optional[float] = None,
- res1: Optional[torch.Tensor] = None,
+ x2: torch.Tensor | None = None,
+ x2_weight: torch.Tensor | None = None,
+ x2_epsilon: float | None = None,
+ res1: torch.Tensor | None = None,
dtype_quant: torch.dtype = dtypes.fp8,
group_size: int = 128,
- quant_type: Optional[int] = None,
+ quant_type: int | None = None,
output_unquantized_inp1: bool = False,
transpose_scale: bool = False,
-) -> Tuple[
+) -> tuple[
torch.Tensor,
torch.Tensor,
torch.Tensor,
@@ -516,14 +516,14 @@ def _fuse_rmsnorm_fp4_quant(
x1: torch.Tensor,
x1_weight: torch.Tensor,
x1_epsilon: float,
- x2: Optional[torch.Tensor] = None,
- x2_weight: Optional[torch.Tensor] = None,
- x2_epsilon: Optional[float] = None,
- res1: Optional[torch.Tensor] = None,
+ x2: torch.Tensor | None = None,
+ x2_weight: torch.Tensor | None = None,
+ x2_epsilon: float | None = None,
+ res1: torch.Tensor | None = None,
shuffle: bool = True,
scale_shuffle_padding: bool = True,
output_unquantized_inp1: bool = False,
-) -> Tuple[
+) -> tuple[
torch.Tensor,
torch.Tensor,
torch.Tensor,
@@ -554,16 +554,16 @@ def _fused_rms_fp8_quant(
x1: torch.Tensor,
x1_weight: torch.Tensor,
x1_epsilon: float,
- x2: Optional[torch.Tensor] = None,
- x2_weight: Optional[torch.Tensor] = None,
- x2_epsilon: Optional[float] = None,
- res1: Optional[torch.Tensor] = None,
+ x2: torch.Tensor | None = None,
+ x2_weight: torch.Tensor | None = None,
+ x2_epsilon: float | None = None,
+ res1: torch.Tensor | None = None,
dtype_quant: torch.dtype = dtypes.fp8,
group_size: int = 128,
- quant_type: Optional[int] = None,
+ quant_type: int | None = None,
output_unquantized_inp1: bool = False,
transpose_scale: bool = False,
-) -> Tuple[
+) -> tuple[
torch.Tensor,
torch.Tensor,
torch.Tensor,
@@ -617,15 +617,15 @@ def _fuse_rmsnorm_quant(
x1: torch.Tensor,
x1_weight: torch.Tensor,
x1_epsilon: float,
- x2: Optional[torch.Tensor] = None,
- x2_weight: Optional[torch.Tensor] = None,
- x2_epsilon: Optional[float] = None,
- res1: Optional[torch.Tensor] = None,
+ x2: torch.Tensor | None = None,
+ x2_weight: torch.Tensor | None = None,
+ x2_epsilon: float | None = None,
+ res1: torch.Tensor | None = None,
dtype_quant: torch.dtype = dtypes.fp8,
shuffle: bool = True,
scale_shuffle_padding: bool = False,
group_size: int = 128,
- quant_type: Optional[int] = None,
+ quant_type: int | None = None,
output_unquantized_inp1: bool = False,
transpose_scale: bool = False,
):
@@ -679,11 +679,11 @@ def _fuse_qkv_a_proj_reduce_rmsnorm_quant_fp4_fake(
q_lora_rank: int,
kv_lora_rank: int,
qk_rope_head_dim: int,
- hidden_states_quant_scale: Optional[torch.Tensor] = None,
- shuffle: Optional[bool] = True,
- scale_shuffle_padding: Optional[bool] = True,
- output_unquantized_inp1: Optional[bool] = False,
-) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
+ hidden_states_quant_scale: torch.Tensor | None = None,
+ shuffle: bool | None = True,
+ scale_shuffle_padding: bool | None = True,
+ output_unquantized_inp1: bool | None = False,
+) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
M = hidden_states_quant.shape[0]
device = hidden_states_quant.device
q_c = torch.empty((M, q_lora_rank // 2), dtype=torch.uint8, device=device)
@@ -715,10 +715,10 @@ def _fuse_qkv_a_proj_reduce_rmsnorm_quant_fp8_fake(
q_lora_rank: int,
kv_lora_rank: int,
qk_rope_head_dim: int,
- hidden_states_quant_scale: Optional[torch.Tensor] = None,
- output_unquantized_inp1: Optional[bool] = False,
+ hidden_states_quant_scale: torch.Tensor | None = None,
+ output_unquantized_inp1: bool | None = False,
transpose_scale: bool = True,
-) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
+) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
M = hidden_states_quant.shape[0]
FP8_QUANT_BLOCK_SIZE = 128
device = hidden_states_quant.device
@@ -748,11 +748,11 @@ def _fuse_qkv_a_proj_reduce_rmsnorm_quant_fp4(
q_lora_rank: int,
kv_lora_rank: int,
qk_rope_head_dim: int,
- hidden_states_quant_scale: Optional[torch.Tensor] = None,
- shuffle: Optional[bool] = True,
- scale_shuffle_padding: Optional[bool] = True,
- output_unquantized_inp1: Optional[bool] = False,
-) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
+ hidden_states_quant_scale: torch.Tensor | None = None,
+ shuffle: bool | None = True,
+ scale_shuffle_padding: bool | None = True,
+ output_unquantized_inp1: bool | None = False,
+) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
M = hidden_states_quant.shape[0]
if hidden_states_quant_scale is None:
@@ -872,10 +872,10 @@ def _fuse_qkv_a_proj_reduce_rmsnorm_quant_fp8(
q_lora_rank: int,
kv_lora_rank: int,
qk_rope_head_dim: int,
- hidden_states_quant_scale: Optional[torch.Tensor] = None,
- output_unquantized_inp1: Optional[bool] = False,
+ hidden_states_quant_scale: torch.Tensor | None = None,
+ output_unquantized_inp1: bool | None = False,
transpose_scale: bool = True,
-) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
+) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
M = hidden_states_quant.shape[0]
# NOTE: this fused path always calls aiter's *preshuffle* blockscale GEMMs,
@@ -985,12 +985,12 @@ def _fuse_qkv_a_proj_reduce_rmsnorm_quant(
kv_lora_rank: int,
qk_rope_head_dim: int,
dtype_quant=dtypes.fp8,
- hidden_states_quant_scale: Optional[torch.Tensor] = None,
- shuffle: Optional[bool] = False,
- scale_shuffle_padding: Optional[bool] = False,
- group_size: Optional[int] = 128,
- output_unquantized_inp1: Optional[bool] = False,
- transpose_scale: Optional[bool] = False,
+ hidden_states_quant_scale: torch.Tensor | None = None,
+ shuffle: bool | None = False,
+ scale_shuffle_padding: bool | None = False,
+ group_size: int | None = 128,
+ output_unquantized_inp1: bool | None = False,
+ transpose_scale: bool | None = False,
):
if dtype_quant == dtypes.fp4x2:
q_c, q_c_scale, kv_c_normed, k_pe = _fuse_qkv_a_proj_reduce_rmsnorm_quant_fp4(
@@ -1040,7 +1040,7 @@ def __init__(
hidden_size: int,
intermediate_size: int,
hidden_act: str,
- quant_config: Optional[QuantizationConfig] = None,
+ quant_config: QuantizationConfig | None = None,
reduce_results: bool = True,
prefix: str = "",
) -> None:
@@ -1077,10 +1077,10 @@ class DeepseekV2MoE(nn.Module):
def __init__(
self,
config: PretrainedConfig,
- quant_config: Optional[QuantizationConfig] = None,
+ quant_config: QuantizationConfig | None = None,
reduce_results: bool = True,
prefix: str = "",
- alt_stream: Optional[torch.cuda.Stream] = None,
+ alt_stream: torch.cuda.Stream | None = None,
):
super().__init__()
self.tp_size = get_tensor_model_parallel_world_size()
@@ -1176,7 +1176,7 @@ def routed_expert_forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
def combine_outputs(
self,
final_hidden_states: torch.Tensor,
- shared_output: Optional[torch.Tensor],
+ shared_output: torch.Tensor | None,
hidden_states: torch.Tensor,
) -> torch.Tensor:
if shared_output is not None:
@@ -1397,6 +1397,11 @@ def _dcp_gather_indexer_k_prefill(
return k_fp8, k_scale
+def _dcp_index_comm_required(dcp_world_size: int, replicated_index_cache: bool) -> bool:
+ """Whether index top-k still needs a cross-rank gather/merge."""
+ return dcp_world_size > 1 and not replicated_index_cache
+
+
def sparse_attn_indexer(
hidden_states: torch.Tensor,
k_cache_prefix: str,
@@ -1405,7 +1410,7 @@ def sparse_attn_indexer(
k: torch.Tensor,
weights: torch.Tensor,
quant_block_size: int,
- scale_fmt: Optional[str],
+ scale_fmt: str | None,
topk_tokens: int,
head_dim: int,
max_model_len: int,
@@ -1434,7 +1439,14 @@ def sparse_attn_indexer(
forward_context = get_forward_context()
attn_metadata = forward_context.attn_metadata
context = forward_context.context
- slot_mapping = attn_metadata.slot_mapping
+ replicated_index_cache = (
+ getattr(attn_metadata, "index_slot_mapping", None) is not None
+ )
+ slot_mapping = (
+ attn_metadata.index_slot_mapping
+ if replicated_index_cache
+ else attn_metadata.slot_mapping
+ )
# Skip for dummy runs to avoid corrupting KV cache
if forward_context.context.is_dummy_run:
# dummy runner
@@ -1447,7 +1459,10 @@ def sparse_attn_indexer(
)
runner_block_size = get_current_atom_config().kv_cache_block_size
cp_kv_cache_interleave_size = get_current_atom_config().dcp_config.interleave_size
- kv_cache = kv_cache.view(-1, runner_block_size, kv_cache.shape[-1])
+ index_block_size = runner_block_size * (
+ get_dcp_world_size() if replicated_index_cache else 1
+ )
+ kv_cache = kv_cache.view(-1, index_block_size, kv_cache.shape[-1])
# PCP prefill: `k` (and `positions`) arrive as the full PADDED key set
# [S_pad] produced by an all-gather of the round-robin shards. The KV-cache
# write (driven by slot_mapping) and the gathered-KV sizing (total_kv =
@@ -1525,7 +1540,7 @@ def sparse_attn_indexer(
dtype=torch.long,
device=prefill_metadata.block_tables.device,
)
- if get_dcp_world_size() > 1:
+ if _dcp_index_comm_required(get_dcp_world_size(), replicated_index_cache):
k_fp8, k_scale = _dcp_gather_indexer_k_prefill(
kv_cache, prefill_metadata, head_dim, k.device
)
@@ -1661,14 +1676,12 @@ def sparse_attn_indexer(
batch_size, next_n, _heads, _ = padded_q_fp8_decode_tokens.shape
num_rows = batch_size * next_n
dcp_world_size = get_dcp_world_size()
+ dcp_rank = get_dcp_rank() if dcp_world_size > 1 else 0
assert topk_tokens == 2048, "top_k_per_row assumes size 2048"
- if dcp_world_size > 1:
- # The fused exchange scores this rank's own shard and writes the KV
- # slots it owns -- ownership filter, slot localize and compaction all
- # inside the op -- straight into sparse_kv_indices_buffer and
- # dcp_sparse_kv_indptr_buffer. So there is no global logits plane
- # left to rank and no topk_indices left to convert: everything the
- # non-DCP path does below has already happened, and we return here.
+ if _dcp_index_comm_required(dcp_world_size, replicated_index_cache):
+ # A sharded index cache scores only this rank's shard. The fused
+ # exchange selects the global winners and emits this rank's owned
+ # physical slots directly, so no filtering remains afterward.
dcp_decode_candidate_exchange_fused(
attn_metadata,
padded_q_fp8_decode_tokens,
@@ -1686,8 +1699,9 @@ def sparse_attn_indexer(
owned_counts=dcp_owned_counts_buffer,
)
return weights
- # Non-DCP: this rank holds the whole plane, so its top-k is already the
- # global one.
+
+ # Either DCP is disabled or the index cache is replicated. In both
+ # cases this rank can score the complete context without communication.
logits = torch.empty(
[num_rows, max_model_len], dtype=torch.float32, device="cuda"
)
@@ -1699,7 +1713,7 @@ def sparse_attn_indexer(
decode_metadata.context_lens,
attn_metadata.block_tables,
max_model_len,
- KVBlockSize=runner_block_size,
+ KVBlockSize=index_block_size,
Preshuffle=True,
)
topk_indices_decode = topk_indices[:num_decode_tokens, :topk_tokens]
@@ -1713,7 +1727,40 @@ def sparse_attn_indexer(
logits.stride(1),
stable=stable_topk,
)
- if attn_metadata.max_seqlen_q > 1:
+ if dcp_world_size > 1:
+ # The replicated index cache produced global positions. Attention's
+ # main KV cache is still sharded, so retain only this rank's owned
+ # positions and map them to its local physical slots.
+ # topk_indices now hold GLOBAL positions. Keep only this rank's owned
+ # tokens ((p//S)%W == r), de-interleave to the local index, map to the
+ # local main-KV slot, and COMPACT them to the front -- non-owned
+ # positions are dropped, not marked with -1, because holes break
+ # aiter's lse output. The compacted per-request lengths are written
+ # into dcp_sparse_kv_indptr_buffer for this layer's attention. The
+ # row unit is the query token, so each MTP draft position gets its
+ # own compacted region.
+ # Cover the padded row count, not just the scheduled one: the
+ # attention kernel reads dcp_sparse_kv_indptr_buffer[: B + 1] with B
+ # padded, and any row left over from the previous step makes that
+ # window non-monotonic. Padded rows hold no top-k, so -1 them first
+ # -- the filter's only validity guard is `tok >= 0`.
+ num_index_tokens = topk_indices.shape[0]
+ topk_indices[num_decode_tokens:num_index_tokens].fill_(-1)
+ triton_filter_and_convert_dcp_index(
+ attn_metadata.token_to_seq_idxs,
+ num_index_tokens,
+ attn_metadata.block_tables,
+ topk_indices,
+ dcp_rank,
+ dcp_world_size,
+ runner_block_size,
+ out_kv_indptr=dcp_sparse_kv_indptr_buffer,
+ owned_counts=dcp_owned_counts_buffer,
+ NUM_TOPK_TOKENS=topk_tokens,
+ out=sparse_kv_indices_buffer,
+ cp_kv_cache_interleave_size=cp_kv_cache_interleave_size,
+ )
+ elif attn_metadata.max_seqlen_q > 1:
triton_gather_kv_indices_sparse(
attn_metadata.sparse_kv_indptr,
attn_metadata.token_to_seq_idxs,
@@ -1744,7 +1791,7 @@ def sparse_attn_indexer_fake(
k: torch.Tensor,
weights: torch.Tensor,
quant_block_size: int,
- scale_fmt: Optional[str],
+ scale_fmt: str | None,
topk_tokens: int,
head_dim: int,
max_model_len: int,
@@ -1841,8 +1888,8 @@ def __init__(
n_head: int,
prefix: str = "",
):
- self._wk_pending_weight: Optional[torch.Tensor] = None
- self._wk_pending_scale: Optional[torch.Tensor] = None
+ self._wk_pending_weight: torch.Tensor | None = None
+ self._wk_pending_scale: torch.Tensor | None = None
self._wk_loaded = False
super().__init__(
hidden_size,
@@ -1882,7 +1929,7 @@ def weight_loader(
self,
param: nn.Parameter,
loaded_weight: torch.Tensor,
- loaded_shard_id: Optional[int] = None,
+ loaded_shard_id: int | None = None,
):
if param is self.weight_scale:
if loaded_shard_id == 0:
@@ -1924,7 +1971,7 @@ def process_weights_after_loading(self):
def _indexer_with_output_fake(
hidden_states: torch.Tensor,
qr: torch.Tensor,
- qr_scale: Optional[torch.Tensor],
+ qr_scale: torch.Tensor | None,
positions: torch.Tensor,
layer_name: str,
sparse_kv_indices_buffer: torch.Tensor,
@@ -1939,7 +1986,7 @@ def _indexer_with_output_fake(
def indexer_with_output(
hidden_states: torch.Tensor,
qr: torch.Tensor,
- qr_scale: Optional[torch.Tensor],
+ qr_scale: torch.Tensor | None,
positions: torch.Tensor,
layer_name: str,
sparse_kv_indices_buffer: torch.Tensor,
@@ -2008,7 +2055,7 @@ def __init__(
config: PretrainedConfig,
hidden_size: int,
q_lora_rank: int,
- quant_config: Optional[QuantizationConfig],
+ quant_config: QuantizationConfig | None,
cache_config: str,
use_wk_weights_proj_fusion: bool = True,
prefix: str = "",
@@ -2109,7 +2156,7 @@ def forward(
self,
hidden_states: torch.Tensor,
qr: torch.Tensor,
- qr_scale: Optional[torch.Tensor],
+ qr_scale: torch.Tensor | None,
positions,
rotary_emb=None,
) -> torch.Tensor:
@@ -2139,7 +2186,7 @@ def forward_impl(
self,
hidden_states: torch.Tensor,
qr: torch.Tensor,
- qr_scale: Optional[torch.Tensor],
+ qr_scale: torch.Tensor | None,
positions,
rotary_emb=None,
) -> torch.Tensor:
@@ -2255,14 +2302,14 @@ def __init__(
qk_nope_head_dim: int,
qk_rope_head_dim: int,
v_head_dim: int,
- q_lora_rank: Optional[int],
+ q_lora_rank: int | None,
kv_lora_rank: int,
max_position_embeddings: int = 8192,
cache_config: str = "bf16",
- quant_config: Optional[QuantizationConfig] = None,
+ quant_config: QuantizationConfig | None = None,
prefix: str = "",
layer_num: int = 0,
- use_indexer_wk_weights_proj_fusion: Optional[bool] = None,
+ use_indexer_wk_weights_proj_fusion: bool | None = None,
) -> None:
super().__init__()
self.hidden_size = hidden_size
@@ -2675,11 +2722,11 @@ def __init__(
config: PretrainedConfig,
prefix: str,
cache_config: str = "bf16",
- quant_config: Optional[QuantizationConfig] = None,
+ quant_config: QuantizationConfig | None = None,
layer_num: int = 0,
is_mtp_block: bool = False,
- alt_stream: Optional[torch.cuda.Stream] = None,
- use_indexer_wk_weights_proj_fusion: Optional[bool] = None,
+ alt_stream: torch.cuda.Stream | None = None,
+ use_indexer_wk_weights_proj_fusion: bool | None = None,
) -> None:
super().__init__()
self.hidden_size = config.hidden_size
@@ -2852,7 +2899,7 @@ def forward(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
- residual: Optional[torch.Tensor],
+ residual: torch.Tensor | None,
) -> torch.Tensor:
# Self Attention
if self.fuse_input_norm_quant:
@@ -3000,7 +3047,7 @@ def __init__(
atom_config: Config,
prefix: str = "",
layer_type: type[nn.Module] = DeepseekV2DecoderLayer,
- use_indexer_wk_weights_proj_fusion: Optional[bool] = None,
+ use_indexer_wk_weights_proj_fusion: bool | None = None,
):
super().__init__()
@@ -3032,7 +3079,7 @@ def __init__(
else:
self.embed_tokens = PPMissingLayer()
- self.alt_stream: Optional[torch.cuda.Stream] = None
+ self.alt_stream: torch.cuda.Stream | None = None
if getattr(config, "n_shared_experts", None) is not None:
self.alt_stream = torch.cuda.Stream()
@@ -3074,11 +3121,9 @@ def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
- intermediate_tensors: Optional[IntermediateTensors],
- inputs_embeds: Optional[torch.Tensor] = None,
- ) -> Union[
- torch.Tensor, IntermediateTensors, Tuple[torch.Tensor, list[torch.Tensor]]
- ]:
+ intermediate_tensors: IntermediateTensors | None,
+ inputs_embeds: torch.Tensor | None = None,
+ ) -> torch.Tensor | IntermediateTensors | tuple[torch.Tensor, list[torch.Tensor]]:
if get_pp_group().is_first_rank:
if inputs_embeds is not None:
hidden_states = inputs_embeds
@@ -3199,9 +3244,9 @@ def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
- intermediate_tensors: Optional[IntermediateTensors] = None,
- inputs_embeds: Optional[torch.Tensor] = None,
- ) -> Union[torch.Tensor, IntermediateTensors]:
+ intermediate_tensors: IntermediateTensors | None = None,
+ inputs_embeds: torch.Tensor | None = None,
+ ) -> torch.Tensor | IntermediateTensors:
# ---- Prefill Context Parallel (PCP) query split ------------------
# During prefill with pcp_size > 1 the token sequence is round-robin
# split so each PCP rank runs the whole model (embed / norm / q-proj /
@@ -3245,7 +3290,7 @@ def forward(
def compute_logits(
self,
hidden_states: torch.Tensor,
- ) -> Optional[torch.Tensor]:
+ ) -> torch.Tensor | None:
logits = self.lm_head(hidden_states)
return logits
diff --git a/atom/spec_decode/eagle_proposer.py b/atom/spec_decode/eagle_proposer.py
index 9279716017..e97c2bb35a 100644
--- a/atom/spec_decode/eagle_proposer.py
+++ b/atom/spec_decode/eagle_proposer.py
@@ -45,28 +45,6 @@ class EagleProposer(Drafter):
flavor is its sibling ``DSparkProposer``; both share the ``Drafter`` base.
"""
- def __init__(self, atom_config, device: torch.device, runner):
- super().__init__(atom_config, device, runner)
- # GLM-5.2 draft index sharing: step 0 runs the MTP indexer, steps 1+
- # reuse sparse_kv_indices_buffer via skip_topk + compact_topk_indices.
- # Gated on method=mtp, DSA index_topk and the config flag, so other
- # draft backends are unchanged. (DSpark is DSparkProposer, not this
- # class, so it cannot reach here.)
- draft_hf = self.speculative_config.draft_model_hf_config
- mtp_inner = getattr(self.model, "model", None)
- self._share_mtp_indices = (
- self.speculative_config.method == "mtp"
- and getattr(draft_hf, "index_share_for_mtp_iteration", False)
- and hasattr(draft_hf, "index_topk")
- and mtp_inner is not None
- and hasattr(mtp_inner, "set_skip_topk")
- )
- if self._share_mtp_indices:
- logger.info(
- "MTP draft index_share_for_mtp_iteration enabled: "
- "step 0 computes indexer top-k, steps 1+ reuse the buffer."
- )
-
def _resolve_mtp_k(self) -> int:
return self.speculative_config.num_speculative_tokens or 0
@@ -345,15 +323,7 @@ def _step_warmup_inputs(self, running_bs, **staged):
steps 1+ run one row per sequence, and every kernel they compile is
chosen from that shape. Replaying the same rewrite serving uses is the
point -- a warmup that built its own would warm a shape nobody asks for.
-
- The same goes for state the model carries: `skip_topk` is read as a
- Python branch inside the draft's attention, and `propose` turns it on
- only after step 0. Warming without it records the branch that recomputes
- the index, which a replay then repeats every step -- the sharing this
- flavor logs as enabled would be dead, silently.
"""
- if self._share_mtp_indices:
- self.model.model.set_skip_topk(True)
fc = get_forward_context()
# Where each synthetic sequence actually is. Warming at position 0 would
# compile a masked, near-empty window -- a shape steady-state decode
@@ -396,6 +366,13 @@ def _enter_decode_metadata(
cu_seqlens_q = var["cu_seqlens_q"].gpu[: running_bs + 1]
attn_metadata.cu_seqlens_q = cu_seqlens_q
attn_metadata.slot_mapping = slot_mapping
+ if getattr(builder, "replicate_index_cache", False):
+ # The replicated index cache has its own mapping, and the indexer
+ # picks it over slot_mapping -- it has to drop to one row per
+ # sequence here too. `prepare_mtp_decode` refreshes it per step.
+ attn_metadata.index_slot_mapping = builder.rebuild_draft_index_slots(
+ scheduled_bs, running_bs
+ )
if has_flat_kv:
kv_indptr = var["kv_indptr"].gpu[: running_bs + 1]
kv_indices = var["kv_indices"].gpu
@@ -416,6 +393,12 @@ def _enter_decode_metadata(
attn_metadata.sparse_kv_last_page_lens = var[
"sparse_kv_last_page_lens"
].gpu[:running_bs]
+ # The DCP top-k filter's row unit is the query token, and it reads
+ # this to pick the owning request's block table. The target left one
+ # entry per verify token; at one row per sequence it is the identity.
+ t2s = builder.rebuild_draft_token_to_seq_idxs(running_bs)
+ if t2s is not None:
+ attn_metadata.token_to_seq_idxs = t2s
# block_tables, context_lens, and sparse_kv_indptr are
# needed by both MHA and MLA+sparse attention
attn_metadata.block_tables = var["block_tables"].gpu[:running_bs]
@@ -578,10 +561,6 @@ def propose(
positions,
hidden_states,
)
- # index_share_for_mtp_iteration: step 0 runs the MTP indexer;
- # steps 1+ skip it and read the compacted sparse_kv buffer.
- if self._share_mtp_indices and i == 0:
- self.model.model.set_skip_topk(False)
if i and self.step is not None:
# Steps 1+ are the declared pass: one row per sequence, at
# a batch the startup sweep already warmed. The head rides
@@ -603,9 +582,6 @@ def propose(
ret_hidden_states = pcp_allgather_rerange(
ret_hidden_states, pcp_ws
)[:n_global_draft]
- if self._share_mtp_indices and i == 0:
- self.model.model.set_skip_topk(True)
- self.model.model.compact_topk_indices(last_token_indices)
# Step 0 gathers one row per sequence out of the token stream;
# steps 1+ already are one row per sequence -- sliced back off
diff --git a/atom/utils/envs.py b/atom/utils/envs.py
index 6331c87027..52e2a33bd8 100644
--- a/atom/utils/envs.py
+++ b/atom/utils/envs.py
@@ -70,6 +70,12 @@
"ATOM_USE_TRITON_MLA_SHUFFLE_KV": lambda: (
os.getenv("ATOM_USE_TRITON_MLA_SHUFFLE_KV", "0") == "1"
),
+ # Experimental native GLM-5.2 DCP mode: keep MLA KV sharded but replicate
+ # full IndexShare-layer index caches on every DCP rank. This removes the
+ # indexer candidate all-gather/global-merge path. Disabled by default.
+ "ATOM_DCP_REPLICATE_INDEX_CACHE": lambda: (
+ os.getenv("ATOM_DCP_REPLICATE_INDEX_CACHE", "0") == "1"
+ ),
"ATOM_USE_TRITON_MOE": lambda: os.getenv("ATOM_USE_TRITON_MOE", "0") == "1",
"ATOM_USE_TRITON_MOE_DECODE": lambda: os.getenv("ATOM_USE_TRITON_MOE_DECODE", "0")
== "1",
diff --git a/atom/utils/forward_context.py b/atom/utils/forward_context.py
index fb688f867d..98421b9b88 100644
--- a/atom/utils/forward_context.py
+++ b/atom/utils/forward_context.py
@@ -587,6 +587,10 @@ class AttentionMetaData:
# prefill/MTP-verify and per seq in decode. Separate from kv_last_page_lens
# (the dense per-seq buffer) so the two never clobber each other.
sparse_kv_last_page_lens: torch.Tensor | None = None
+ # Absolute addresses into the replicated DSA index cache, one per query
+ # token. Its presence is the switch: not None makes the indexer use it in
+ # place of slot_mapping, None leaves the dcp-sharded index layout.
+ index_slot_mapping: torch.Tensor | None = None
# Precomputed per-rank context lengths for qlen=1 sparse DSA + DCP decode.
# The padded running-width buffer is shared by eager, graph and TBO paths.
dcp_local_context_lens: torch.Tensor | None = None
diff --git a/docs/context_parallel_guide.md b/docs/context_parallel_guide.md
index e1f4f31ca1..d887a09e6e 100644
--- a/docs/context_parallel_guide.md
+++ b/docs/context_parallel_guide.md
@@ -263,9 +263,12 @@ existing TP ranks, so `world = tp` and `tp` must be divisible by `dcp`.
> the ATOM server path is validated; the **vllm-atom plugin path is not yet
> verified**.
>
-> **Not yet supported: DSA + DCP + MTP.** Speculative decode (MTP, `q > 1`) is
-> only available for dense MLA (gfx950); combining it with sparse attention under
-> DCP is rejected at runtime — see the DCP Constraints & Compatibility table below.
+> **DSA + DCP + MTP needs the replicated index cache.** Speculative decode
+> (MTP, `q > 1`) on sparse MLA under DCP requires
+> `ATOM_DCP_REPLICATE_INDEX_CACHE=1`. With the default *sharded* index cache the
+> global top-k comes from a cross-rank candidate all-gather that only handles
+> `qlen=1`, and an MTP verify step is rejected by an assert. Dense MLA is
+> unaffected. See the DCP Constraints & Compatibility table below.
## When to use DCP
@@ -291,6 +294,7 @@ existing TP ranks, so `world = tp` and `tp` must be divisible by `dcp`.
| `--kv-cache-dtype fp8` | `auto` | Supported with DCP (per-tensor scale). `auto`/`bf16` also fine |
| `--enable_prefix_caching` | off | Supported with DCP |
| `--enable_chunked_prefill` / `--no-enable_chunked_prefill` | on | Chunked prefill is supported with DCP; on by default |
+| `ATOM_DCP_REPLICATE_INDEX_CACHE=1` | off | Sparse MLA (DSA) only: replicate the index cache on every DCP rank instead of sharding it, which drops the indexer's candidate all-gather and is what makes DSA + DCP + MTP work. See [Sparse MLA (DSA) + MTP](#sparse-mla-dsa--mtp) |
| Goal (8 GPUs) | Command |
|------|---------|
@@ -533,11 +537,12 @@ compacts its rank-local top-k. The rebuild consumes that layer's
indptr, and work plan. This makes the persistent descriptors and the actual
sparse regions agree without rebuilding metadata on shared layers.
-The implementation is scoped to native, non-speculative serving on gfx950 with
-page size 1: decode is q_len=1, while sparse prefill is represented as per-token
-virtual q_len=1 rows. Unsupported paths (including gfx942 and plugin or
-speculative sparse DCP paths without the per-layer rebuild) remain
-non-persistent and round a gathered 64 up to **128**.
+The implementation covers native serving on gfx950 with page size 1, speculative
+decode included: decode is q_len=1, sparse prefill is represented as per-token
+virtual q_len=1 rows, and an MTP verify step rebuilds into its own
+`sparse_mtp_`-prefixed work buffers for the per-token layout. Paths without the
+per-layer rebuild (gfx942, or a page size above 1) remain non-persistent and
+round a gathered 64 up to **128**.
**GLM-5.2 is the model that benefits**: 64 query heads at `-tp 8` is 8 per
rank, so `-dcp 8` gathers exactly 64 and now dispatches the native persistent
@@ -576,9 +581,12 @@ mask, so the intra-block mask has to be applied on global positions. This is
handled by a dedicated **round-robin CP (`cprr`) MLA kernel**, selected
automatically when DCP is on, `q > 1`, and the decode is causal.
-> **Dense MLA only.** This applies to dense MLA (V3 / R1). **DSA / sparse MLA
-> (V3.2-Exp) does not support MTP under DCP yet** — sparse decode with `q > 1` is
-> rejected by an assert. Serve DSA + DCP without `--method mtp`.
+> **Sparse MLA takes a different route.** The `cprr` kernel above is the dense
+> MLA (V3 / R1) path. **DSA / sparse MLA under DCP verifies MTP through the
+> per-query-token top-k filter instead**, and only with
+> `ATOM_DCP_REPLICATE_INDEX_CACHE=1` — see
+> [Sparse MLA (DSA) + MTP](#sparse-mla-dsa--mtp). Left on the default sharded
+> index cache, sparse decode with `q > 1` is still rejected by an assert.
**Support matrix:**
@@ -608,6 +616,29 @@ Plugin path (`vllm serve`): add `--speculative-config '{"method":"mtp","num_spec
Accuracy: gsm8k (DeepSeek-R1, tp8, 5-shot) matches the non-speculative DCP
baseline (≈0.95) across bf16/fp8 and `num_speculative_tokens` 1/2/3.
+### Sparse MLA (DSA) + MTP
+
+Sparse decode does not use the `cprr` kernel. Its intra-block causality is
+already carried by the per-layer top-k, so the only thing MTP changes is the
+*row unit*: a verify step forwards `max_seqlen_q` draft positions per request,
+each with its own top-k, while the block table stays per request.
+`triton_filter_and_convert_dcp_index` therefore keys on the **query token**
+(`token_to_seq_idxs` maps a row back to its owning request) rather than on
+`qo_indptr`. At `qlen == 1` the two are identical, so the non-speculative path
+is unchanged.
+
+This requires **`ATOM_DCP_REPLICATE_INDEX_CACHE=1`**. The default sharded index
+cache reaches its global top-k through `dcp_decode_candidate_exchange`, which
+asserts `max_seqlen_q == 1`; replicating the index cache removes that exchange
+entirely, so every rank scores the whole sequence and each draft position gets
+its own compacted region. The cost is `dcp_size` × index-cache memory, and an
+index page widens from `block_size` to `block_size * dcp_size` tokens. See
+[environment variables](environment_variables.md#decode-context-parallelism-dcp)
+for the full startup gate.
+
+In a P/D deployment this is a **decode-side** flag. A CPP prefill node runs
+`pp > 1` / `dcp = 1`, and the layout rejects both.
+
### DSpark (Kimi-K3)
[DSpark](../recipes/DSpark.md) drafts a whole block in one parallel backbone
@@ -655,7 +686,8 @@ contributing under DCP rather than collapsing back to single-token decode.
| gathered head width 64 + fp8 KV | Native sparse / DSA prefill and q_len=1 decode on gfx950 rebuild persistent metadata per full IndexShare layer and run gqa64 directly; unsupported non-persistent paths still pad to 128 — see [Sparse DCP persistent attention and gqa=64](#sparse-dcp-persistent-attention-and-gqa64) |
| speculative decode (MTP), dense MLA | Supported on **gfx950 only** (bf16/fp8, `num_speculative_tokens` 1–3); raises at startup on gfx942 |
| speculative decode (DSpark), Kimi-K3 | Supported on gfx950; validated at `tp8 -dcp 8` with `num_speculative_tokens 2` |
-| speculative decode (MTP), sparse / DSA | **Not supported** — sparse decode with `q > 1` under DCP is rejected by an assert |
+| speculative decode (MTP), sparse / DSA | Supported **only with `ATOM_DCP_REPLICATE_INDEX_CACHE=1`**; on the default sharded index cache sparse decode with `q > 1` is rejected by an assert |
+| `ATOM_DCP_REPLICATE_INDEX_CACHE` | Experimental, decode-side, `glm_moe_dsa` + `-dcp > 1` + page size 1 only; rejects PP, PCP and RapidServe at startup. Costs `dcp_size` × index-cache memory |
| vllm-atom plugin | Validated for dense MLA only; the sparse / DSA and Kimi-K3 plugin paths are not yet verified |
| DCP + PCP | Independent dimensions (different phases); combined use not validated here |
diff --git a/docs/environment_variables.md b/docs/environment_variables.md
index be971c8a3c..a68a60b76c 100644
--- a/docs/environment_variables.md
+++ b/docs/environment_variables.md
@@ -155,6 +155,17 @@ that ever changes, so a fusion left inert by an unrecognised layout says so.
|----------|------|---------|-------------|
| **ATOM_DSPARK_FUSED_CTX_KV** | bool | 1 (true) | Write the context rows with one Triton kernel (RMSNorm + RoPE + concat + paged store) instead of four launches plus a throwaway `empty_like` for the RoPE's query side. Falls back per call when the cache layout or the RoPE is not the plain one the kernel understands (seg / shuffled-KV layouts keep their own write kernels), and until the RoPE's cos/sin cache has reached the device. Measured on Kimi-K3 (MI355X, TP8, fp8 KV): one 4.65 µs kernel replaces a 14 µs three-kernel chain, saving ~39 µs per drafting step at B=1 and ~36 µs at B=64. Set to `0` to force the per-op chain; that chain is the fallback above rather than debug code, so it stays reachable either way (it runs the first write of every layer). |
+## Decode context parallelism (DCP)
+
+Shape of the sparse-MLA (DSA) index cache under DCP. The index cache is normally
+sharded the same way the MLA latent KV is, which makes the indexer's global
+top-k a cross-rank candidate all-gather. See the
+[context parallel guide](context_parallel_guide.md).
+
+| Variable | Type | Default | Description |
+|----------|------|---------|-------------|
+| **ATOM_DCP_REPLICATE_INDEX_CACHE** | bool | 0 (false) | Experimental. Keep the MLA latent KV sharded but replicate every full IndexShare layer's index cache on all DCP ranks, so each rank scores the whole sequence and the candidate all-gather/global-merge disappears. Costs `dcp_size` × index-cache memory, and widens an index page from `block_size` to `block_size * dcp_size` tokens. Enabling it raises at startup unless the whole supported topology holds: `glm_moe_dsa`, `-dcp > 1`, `ATOM_MLA_PAGE_SIZE=1`, no PCP, no PP, no RapidServe, and a KV-transfer config that is standalone LMCache, standalone Mooncake, or multi[Mooncake producer + LMCache offload]. In a P/D deployment this is a **decode-side** flag: a CPP prefill node runs pp>1/dcp=1 and the layout rejects both. |
+
## V4 attention backend (Migration)
Selects between the legacy per-seq Python dispatch path in `atom/models/deepseek_v4.py`
diff --git a/tests/test_dcp_sparse_filter.py b/tests/test_dcp_sparse_filter.py
index d916117665..9a4947eb4f 100644
--- a/tests/test_dcp_sparse_filter.py
+++ b/tests/test_dcp_sparse_filter.py
@@ -1,10 +1,10 @@
# SPDX-License-Identifier: MIT
-"""DCP sparse PREFILL index filter: a top-k -> this rank's compacted slot list.
+"""DCP sparse index filter: a global top-k -> this rank's compacted slot list.
-Covers ``triton_filter_and_convert_dcp_index_prefill`` in
-``atom/model_ops/dcp_ops.py``. Decode had a twin until the fused merge
-(``aiter.flydsl_dcp_topk_merge``) took over, emitting this rank's owned slots
-directly with nothing left to filter; its tests live in test_dcp_topk.py.
+Both halves of the same kernel family in ``atom/model_ops/dcp_ops.py``:
+
+ * decode -- replicated-index-cache fallback only
+ * prefill -- ``triton_filter_and_convert_dcp_index_prefill``
Why the checks go past "it does not crash": ``cp_lse_ag_out_rs`` rebuilds a
global softmax out of per-rank partial attentions, and that is only valid when
@@ -31,7 +31,10 @@
try:
from atom.model_ops.attentions.aiter_mla import AiterMLAMetadataBuilder
- from atom.model_ops.dcp_ops import triton_filter_and_convert_dcp_index_prefill
+ from atom.model_ops.dcp_ops import (
+ triton_filter_and_convert_dcp_index,
+ triton_filter_and_convert_dcp_index_prefill,
+ )
except ImportError as _e: # triton/aiter absent on a CPU-only runner
pytest.skip(f"requires full atom import env: {_e}", allow_module_level=True)
@@ -63,6 +66,216 @@ def writer_slot(block_table_row, pos, rank, world, page, interleave=1):
)
+# ─────────────────────────────────────────────────────────────── decode side ──
+
+DEC_W = 4 # dcp world size
+DEC_K = 256 # NUM_TOPK_TOKENS (must be a multiple of BLOCK_N=128)
+DEC_PAGE = 16 # runner physical block size
+
+
+def _build_decode_case(g_ctxs, max_blocks, seed, max_seqlen_q=1):
+ """Random global top-k selections + block table for the given contexts.
+
+ ``max_seqlen_q`` is the MTP verify width: a request contributes that many
+ query rows, each with its own top-k and its own context (a draft chain
+ extends the context by one per step). Rows then outnumber requests, and
+ ``token_to_seq_idxs`` is the only thing that resolves a row to a block
+ table -- at width 1 it degenerates to the identity the decode path had.
+ """
+ gen = torch.Generator().manual_seed(seed)
+ bs = len(g_ctxs)
+
+ token_to_seq_idxs = torch.arange(bs, dtype=torch.int32).repeat_interleave(
+ max_seqlen_q
+ )
+
+ # Physical blocks are deliberately shuffled so a wrong slot formula cannot
+ # accidentally match a "logical == physical" identity mapping.
+ block_table = (
+ torch.randperm(bs * max_blocks, generator=gen)[: bs * max_blocks]
+ .reshape(bs, max_blocks)
+ .to(torch.int32)
+ )
+
+ ctxs = [g + j for g in g_ctxs for j in range(max_seqlen_q)]
+ token_indices = torch.full((len(ctxs), DEC_K), -1, dtype=torch.int32)
+ for t, ctx in enumerate(ctxs):
+ n = min(ctx, DEC_K)
+ # distinct global positions in [0, ctx), in the indexer's (arbitrary) order
+ picks = torch.randperm(ctx, generator=gen)[:n]
+ token_indices[t, :n] = picks.to(torch.int32)
+ return token_to_seq_idxs, block_table, token_indices, ctxs
+
+
+def _decode_reference(
+ ctxs, token_to_seq_idxs, block_table, token_indices, rank, interleave=1
+):
+ """Expected compacted slots per query row, taken from the write side."""
+ out = []
+ for t, ctx in enumerate(ctxs):
+ b = int(token_to_seq_idxs[t])
+ n = min(ctx, DEC_K)
+ slots = []
+ for c in range(n):
+ tok = int(token_indices[t, c])
+ if tok < 0:
+ continue
+ # ownership AND placement both come from the writer
+ slot = writer_slot(block_table[b], tok, rank, DEC_W, DEC_PAGE, interleave)
+ if slot >= 0:
+ slots.append(slot)
+ out.append(slots)
+ return out
+
+
+@pytest.mark.parametrize(
+ "name, g_ctxs, max_seqlen_q, seed",
+ [
+ ("short ctx (< topk)", [13], 1, 1),
+ ("multi-request mixed", [13, 100, 7, 300], 1, 2),
+ ("ctx > topk (clipped)", [1000, 4096], 1, 3),
+ ("page boundary", [DEC_PAGE * DEC_W, DEC_PAGE * DEC_W + 1], 1, 4),
+ # ctx=2 with W=4 leaves ranks 2 and 3 owning nothing for that request.
+ ("zero-owned ranks", [2, 1], 1, 5),
+ # MTP verify: rows outnumber requests, so the block table can only be
+ # reached through token_to_seq_idxs. Shuffled physical blocks make
+ # resolving a row against the wrong request land on the wrong KV.
+ ("mtp k=3", [37, 128, 5, 260], 4, 6),
+ ("mtp k=1, page boundary", [DEC_PAGE * DEC_W, 3], 2, 7),
+ ],
+)
+def test_decode_filter(name, g_ctxs, max_seqlen_q, seed):
+ rows = len(g_ctxs) * max_seqlen_q
+ span = max(g_ctxs) + max_seqlen_q - 1 # the last draft row's context
+ max_blocks = max(1, (span + DEC_PAGE * DEC_W - 1) // (DEC_PAGE * DEC_W)) + 1
+ token_to_seq_idxs, block_table, token_indices, ctxs = _build_decode_case(
+ g_ctxs, max_blocks, seed, max_seqlen_q
+ )
+
+ t2s_g = token_to_seq_idxs.to(DEV)
+ bt_g = block_table.to(DEV)
+ ti_g = token_indices.to(DEV)
+
+ per_rank_lens = []
+ for rank in range(DEC_W):
+ out_buf = torch.full((rows * DEC_K,), -999, dtype=torch.int32, device=DEV)
+ out_indptr = torch.zeros(rows + 1, dtype=torch.int32, device=DEV)
+ counts = torch.zeros(rows, dtype=torch.int32, device=DEV)
+
+ triton_filter_and_convert_dcp_index(
+ t2s_g,
+ rows,
+ bt_g,
+ ti_g,
+ rank,
+ DEC_W,
+ DEC_PAGE,
+ out_kv_indptr=out_indptr,
+ owned_counts=counts,
+ NUM_TOPK_TOKENS=DEC_K,
+ out=out_buf,
+ )
+ torch.cuda.synchronize()
+
+ exp = _decode_reference(
+ ctxs, token_to_seq_idxs, block_table, token_indices, rank
+ )
+ indptr = out_indptr.cpu().tolist()
+ true_counts = counts.cpu().tolist()
+
+ for t in range(rows):
+ got_len = indptr[t + 1] - indptr[t]
+ assert true_counts[t] == len(exp[t])
+ assert got_len == max(len(exp[t]), 1), (
+ f"[{name}] rank{rank} row{t}: region length {got_len} "
+ f"!= metadata length {max(len(exp[t]), 1)}"
+ )
+ for t in range(rows):
+ got = out_buf[indptr[t] : indptr[t + 1]].cpu().tolist()
+ expected = exp[t] if exp[t] else [0]
+ assert got == expected, f"[{name}] rank{rank} row{t}: {got} != {expected}"
+
+ written = out_buf[: indptr[rows]]
+ assert (
+ int((written < 0).sum()) == 0
+ ), f"[{name}] rank{rank}: -1 hole inside the compacted region"
+
+ per_rank_lens.append(true_counts)
+
+ # Partition: every valid top-k token is claimed by exactly one rank.
+ # Checked on COUNTS, not on slot values -- slots are per-rank local
+ # addresses (each rank holds its own 1/W KV shard), so equal slot numbers
+ # across ranks are expected and carry no information. Counts summing to n
+ # rules out both dropped and double-claimed tokens, which is what
+ # cp_lse_ag_out_rs needs.
+ for t, ctx in enumerate(ctxs):
+ n = min(ctx, DEC_K)
+ total = sum(per_rank_lens[rank][t] for rank in range(DEC_W))
+ assert total == n, f"[{name}] row{t}: kept {total} of {n} top-k tokens"
+
+
+@pytest.mark.parametrize("interleave", [2, 4]) # both divide DEC_PAGE=16
+@pytest.mark.parametrize(
+ "g_ctxs, seed",
+ [([13, 100, 7, 300], 12), ([1000, 4096], 13), ([DEC_PAGE * DEC_W + 1], 14)],
+)
+def test_decode_filter_block_interleave(interleave, g_ctxs, seed):
+ """Same partition + writer-agreement checks as test_decode_filter, but with
+ block-level interleave S>1 (cp_kv_cache_interleave_size). The reference
+ slots come from the same _dcp_round_robin_slot writer, now with S, so the
+ filter kernel's owner/offset math is pinned to the write side at S>1."""
+ bs = len(g_ctxs)
+ max_blocks = max(1, (max(g_ctxs) + DEC_PAGE * DEC_W - 1) // (DEC_PAGE * DEC_W)) + 1
+ token_to_seq_idxs, block_table, token_indices, ctxs = _build_decode_case(
+ g_ctxs, max_blocks, seed
+ )
+ t2s_g = token_to_seq_idxs.to(DEV)
+ bt_g = block_table.to(DEV)
+ ti_g = token_indices.to(DEV)
+
+ per_rank_lens = []
+ for rank in range(DEC_W):
+ out_buf = torch.full((bs * DEC_K,), -999, dtype=torch.int32, device=DEV)
+ out_indptr = torch.zeros(bs + 1, dtype=torch.int32, device=DEV)
+ counts = torch.zeros(bs, dtype=torch.int32, device=DEV)
+
+ triton_filter_and_convert_dcp_index(
+ t2s_g,
+ bs,
+ bt_g,
+ ti_g,
+ rank,
+ DEC_W,
+ DEC_PAGE,
+ out_kv_indptr=out_indptr,
+ owned_counts=counts,
+ NUM_TOPK_TOKENS=DEC_K,
+ out=out_buf,
+ cp_kv_cache_interleave_size=interleave,
+ )
+ torch.cuda.synchronize()
+
+ exp = _decode_reference(
+ ctxs, token_to_seq_idxs, block_table, token_indices, rank, interleave
+ )
+ indptr = out_indptr.cpu().tolist()
+ true_counts = counts.cpu().tolist()
+ for b in range(bs):
+ got = out_buf[indptr[b] : indptr[b + 1]].cpu().tolist()
+ assert true_counts[b] == len(exp[b])
+ expected = exp[b] if exp[b] else [0]
+ assert (
+ got == expected
+ ), f"S={interleave} rank{rank} req{b}: {got} != {expected}"
+ assert int((out_buf[: indptr[bs]] < 0).sum()) == 0, "-1 hole in region"
+ per_rank_lens.append(true_counts)
+
+ for b, g in enumerate(g_ctxs):
+ n = min(g, DEC_K)
+ total = sum(per_rank_lens[rank][b] for rank in range(DEC_W))
+ assert total == n, f"S={interleave} req{b}: kept {total} of {n}"
+
+
# ────────────────────────────────────────────────────────────── prefill side ──
PRE_W = 8 # overridden per-test below
@@ -251,3 +464,34 @@ def _writer(b, p):
# the zero-owned path must have been exercised -- guard against the empty-row
# assertions passing vacuously.
assert n_empty > 0, "expected some rank to own nothing for the early tokens"
+
+
+def test_filter_row_guard_raises_rather_than_asserts():
+ """A short buffer must surface as an exception `python -O` cannot strip.
+
+ The kernels take raw pointers, so this guard is the only thing between an
+ undersized scratch buffer and an out-of-bounds launch.
+ """
+ from atom.model_ops.dcp_ops import _check_dcp_filter_rows
+
+ rows = torch.zeros(8, dtype=torch.int32, device=DEV)
+ block_table = torch.zeros(2, 4, dtype=torch.int32, device=DEV)
+ kwargs = {
+ "token_to_seq_idxs": rows,
+ "topk_indices": rows,
+ "out_kv_indptr": rows,
+ "owned_counts": rows,
+ "block_table": block_table,
+ }
+
+ _check_dcp_filter_rows(7, **kwargs)
+
+ # out_kv_indptr needs num_tokens + 1, so 8 tokens is one row short.
+ with pytest.raises(ValueError, match="out_kv_indptr holds 8 rows"):
+ _check_dcp_filter_rows(8, **kwargs)
+
+ with pytest.raises(TypeError, match="owned_counts must be int32"):
+ _check_dcp_filter_rows(7, **{**kwargs, "owned_counts": rows.to(torch.int64)})
+
+ with pytest.raises(ValueError, match="block_table must be"):
+ _check_dcp_filter_rows(7, **{**kwargs, "block_table": block_table[0]})
diff --git a/tests/test_dcp_topk.py b/tests/test_dcp_topk.py
index bd8eebb482..aa89f82192 100644
--- a/tests/test_dcp_topk.py
+++ b/tests/test_dcp_topk.py
@@ -274,8 +274,6 @@ def test_reruns_stay_valid_even_when_the_set_shifts():
# owned twice. ``cp_lse_ag_out_rs`` combines the per-rank partial attentions by
# summing them, so a token counted on two ranks is silently double-weighted --
# no crash, just a wrong answer.
-
-
@pytest.mark.parametrize(
"name, rows, ctx, world",
[
diff --git a/tests/test_lmcache_offload_config.py b/tests/test_lmcache_offload_config.py
index 6cd92f0db8..2d0d856906 100644
--- a/tests/test_lmcache_offload_config.py
+++ b/tests/test_lmcache_offload_config.py
@@ -118,6 +118,24 @@ def test_page_namespace_changes_for_meaningful_geometry(mutate):
assert offcfg.build_page_namespace(config, cfg, 4) != original
+def test_page_namespace_separates_replicated_index_layout(monkeypatch):
+ config = _config()
+ config.speculative_config = None
+ config.hf_config.model_type = "glm_moe_dsa"
+ config.hf_config.indexer_types = ["full", "shared", "shared"]
+ cfg = _lmcache_config()
+
+ monkeypatch.delenv("ATOM_DCP_REPLICATE_INDEX_CACHE", raising=False)
+ sharded = offcfg.build_page_namespace(config, cfg, 4)
+ monkeypatch.setenv("ATOM_DCP_REPLICATE_INDEX_CACHE", "1")
+ replicated = offcfg.build_page_namespace(config, cfg, 4)
+ assert replicated != sharded
+
+ config.hf_config.indexer_types = ["full", "full", "shared"]
+ changed_schedule = offcfg.build_page_namespace(config, cfg, 4)
+ assert changed_schedule != replicated
+
+
def test_page_namespace_changes_when_code_layout_version_changes():
current = offcfg.build_page_namespace(_config(), _lmcache_config(), 4)
future = offcfg.build_page_namespace(
diff --git a/tests/test_mla_index_cache.py b/tests/test_mla_index_cache.py
index 01693cde22..7b6b99b3ab 100644
--- a/tests/test_mla_index_cache.py
+++ b/tests/test_mla_index_cache.py
@@ -129,6 +129,10 @@ def _builder(
),
block_size=16,
is_deepseek_v32=True,
+ # allocate_kv_cache_tensors() sets this on the real ModelRunner
+ # (index_head_dim + one fp32 scale, rounded up to 16 bytes); the
+ # transfer-region builder divides page bytes by it.
+ aligned_index_dim=((index_head_dim + 4 + 15) // 16) * 16,
_get_total_num_layers=lambda: total_local_layers,
)
builder = object.__new__(AiterMLAMetadataBuilder)
@@ -277,28 +281,29 @@ def view(self, *shape):
class _FakeTransferTensor:
- def __init__(self, address):
+ def __init__(self, address, page_stride=1):
self._address = address
+ self._page_stride = page_stride
def stride(self, dim):
assert dim == 0
- return 1
+ return self._page_stride
def element_size(self):
return 1
def numel(self):
- return 8
+ return 8 * self._page_stride
def data_ptr(self):
return self._address
class _FakeTransferStack:
- def __init__(self, num_layers, address_base):
+ def __init__(self, num_layers, address_base, page_stride=1):
self.shape = (num_layers,)
self._layers = [
- _FakeTransferTensor(address_base + layer_id)
+ _FakeTransferTensor(address_base + layer_id, page_stride)
for layer_id in range(num_layers)
]
@@ -376,13 +381,20 @@ def test_transfer_regions_use_explicit_compact_consumer_map(monkeypatch):
_mock_pp(monkeypatch, rank=1, world_size=2)
builder.block_ratio = 1
runner.kv_cache = _FakeTransferStack(4, 100)
- runner.index_cache = _FakeTransferStack(3, 200)
+ # One index page is block_size tokens of aligned_index_dim bytes.
+ runner.index_cache = _FakeTransferStack(3, 200, page_stride=16 * 144)
runner.index_cache_layer_ids = (3, 5, 6)
runner.config.num_kvcache_blocks = 8
transfer_tensors = builder.get_kv_transfer_tensors()
assert len(transfer_tensors.block_regions) == 7
+ # The three index regions follow the four MLA KV regions. Each page holds
+ # 16 keys of index_head_dim bytes plus 16 fp32 scales; the rest of
+ # aligned_index_dim is padding and is not part of either plane.
+ index_regions = transfer_tensors.block_regions[4:]
+ assert [r.key_plane_bytes for r in index_regions] == [16 * 128] * 3
+ assert [r.scale_plane_bytes for r in index_regions] == [16 * 4] * 3
assert transfer_tensors.block_region_consumer_indices == [
3,
4,
@@ -394,6 +406,36 @@ def test_transfer_regions_use_explicit_compact_consumer_map(monkeypatch):
]
+def test_draft_index_slots_skip_padded_rows_without_indexing_the_block_table():
+ """A padded row's context_len is 0, so ``pos = ctx - 1`` is -1.
+
+ ``_enter_decode_metadata`` rebuilds the mapping before the step's
+ ``context_lens += 1``, so the pad rows really are at 0 there, and gathering
+ the block table at -1 faults before the sentinel below can mask them.
+ """
+ scheduled_bs, running_bs = 2, 4
+ var = {
+ "context_lens": SimpleNamespace(
+ gpu=torch.tensor([70, 5, 0, 0], dtype=torch.int32)
+ ),
+ "block_tables": SimpleNamespace(
+ gpu=torch.tensor([[3, 9], [7, 0], [0, 0], [0, 0]], dtype=torch.int32)
+ ),
+ "index_slot_mapping": SimpleNamespace(
+ gpu=torch.zeros(running_bs, dtype=torch.int64)
+ ),
+ }
+ builder = object.__new__(AiterMLAMetadataBuilder)
+ builder.replicate_index_cache = True
+ builder.dcp_world_size = 4
+ builder.model_runner = SimpleNamespace(block_size=16, forward_vars=var)
+
+ slots = builder.rebuild_draft_index_slots(scheduled_bs, running_bs)
+
+ # Page width is block_size * dcp_world_size = 64.
+ assert slots.tolist() == [9 * 64 + 5, 7 * 64 + 4, -1, -1]
+
+
class _FakeMetadataBuffer:
def __init__(self, size):
self.np = np.zeros(size, dtype=np.int32)
diff --git a/tests/test_mtp_index_share.py b/tests/test_mtp_index_share.py
deleted file mode 100644
index fc4722a278..0000000000
--- a/tests/test_mtp_index_share.py
+++ /dev/null
@@ -1,131 +0,0 @@
-"""Unit tests for MTP draft index_share_for_mtp_iteration helpers.
-
-Avoid importing ``atom.models.deepseek_mtp`` here: that module pulls in
-``atom.config`` (torch / aiter / transformers) and is brittle in lightweight
-test collection. The helpers below mirror
-``DeepSeekMultiTokenPredictor.set_skip_topk`` /
-``compact_topk_indices`` — keep them in sync when editing production code.
-"""
-
-from __future__ import annotations
-
-from types import SimpleNamespace
-from typing import Any
-
-import pytest
-import torch
-
-
-class _FakeMlaAttn:
- def __init__(self, buf: torch.Tensor):
- self.sparse_kv_indices_buffer = buf
-
-
-class _FakeSelfAttn:
- def __init__(self, *, has_indexer: bool, buf: torch.Tensor):
- self.skip_topk = False
- self.indexer = object() if has_indexer else None
- self.mla_attn = _FakeMlaAttn(buf)
-
-
-class _FakeMtpBlock:
- def __init__(self, self_attn: _FakeSelfAttn):
- self.self_attn = self_attn
-
-
-class _FakeLayer:
- def __init__(self, self_attn: _FakeSelfAttn):
- self.mtp_block = _FakeMtpBlock(self_attn)
-
-
-def _set_skip_topk(layers: dict[str, Any], skip: bool) -> None:
- """Mirror of ``DeepSeekMultiTokenPredictor.set_skip_topk``."""
- for layer in layers.values():
- mtp_block = getattr(layer, "mtp_block", None)
- if mtp_block is None:
- continue
- self_attn = getattr(mtp_block, "self_attn", None)
- if self_attn is None or not hasattr(self_attn, "skip_topk"):
- continue
- if getattr(self_attn, "indexer", None) is not None:
- self_attn.skip_topk = skip
-
-
-def _compact_topk_indices(layers: dict[str, Any], slot_ids: torch.Tensor) -> None:
- """Mirror of ``DeepSeekMultiTokenPredictor.compact_topk_indices``."""
- num_slots = slot_ids.numel()
- for layer in layers.values():
- mtp_block = getattr(layer, "mtp_block", None)
- if mtp_block is None:
- continue
- self_attn = getattr(mtp_block, "self_attn", None)
- if self_attn is None:
- continue
- mla_attn = getattr(self_attn, "mla_attn", None)
- if mla_attn is None:
- continue
- sparse_buf = getattr(mla_attn, "sparse_kv_indices_buffer", None)
- if sparse_buf is not None and sparse_buf.numel() > 0:
- sparse_buf[:num_slots] = sparse_buf[slot_ids]
-
-
-def test_set_skip_topk_only_layers_with_indexer():
- buf0 = torch.zeros(4, 8, dtype=torch.int32)
- buf1 = torch.zeros(4, 8, dtype=torch.int32)
- layers = {
- "80": _FakeLayer(_FakeSelfAttn(has_indexer=True, buf=buf0)),
- "81": _FakeLayer(_FakeSelfAttn(has_indexer=False, buf=buf1)),
- }
-
- _set_skip_topk(layers, True)
-
- assert layers["80"].mtp_block.self_attn.skip_topk is True
- assert layers["81"].mtp_block.self_attn.skip_topk is False
-
-
-def test_compact_topk_indices_gathers_rows_to_front():
- buf = torch.arange(20, dtype=torch.int32).reshape(10, 2)
- layers = {"80": _FakeLayer(_FakeSelfAttn(has_indexer=True, buf=buf))}
-
- slot_ids = torch.tensor([3, 7], dtype=torch.int64)
- _compact_topk_indices(layers, slot_ids)
-
- expected_row0 = torch.arange(6, 8, dtype=torch.int32)
- expected_row1 = torch.arange(14, 16, dtype=torch.int32)
- assert torch.equal(buf[0], expected_row0)
- assert torch.equal(buf[1], expected_row1)
-
-
-def test_compact_topk_indices_skips_empty_buffer():
- empty = torch.empty(0, dtype=torch.int32)
- layers = {"80": _FakeLayer(_FakeSelfAttn(has_indexer=True, buf=empty))}
- _compact_topk_indices(layers, torch.tensor([0], dtype=torch.int64))
-
-
-@pytest.mark.parametrize(
- "method,index_share,index_topk,has_api,expected",
- [
- ("mtp", True, 2048, True, True),
- ("mtp", False, 2048, True, False),
- ("eagle3", True, 2048, True, False),
- ("mtp", True, None, True, False),
- ("mtp", True, 2048, False, False),
- ],
-)
-def test_share_mtp_indices_gate(method, index_share, index_topk, has_api, expected):
- draft_hf = SimpleNamespace(
- index_share_for_mtp_iteration=index_share,
- )
- if index_topk is not None:
- draft_hf.index_topk = index_topk
- mtp_inner = (
- SimpleNamespace(set_skip_topk=lambda _: None) if has_api else SimpleNamespace()
- )
-
- share = (
- method == "mtp"
- and getattr(draft_hf, "index_share_for_mtp_iteration", False)
- and hasattr(draft_hf, "index_topk")
- and hasattr(mtp_inner, "set_skip_topk")
- )
- assert share is expected
diff --git a/tests/test_pd_pp.py b/tests/test_pd_pp.py
index e95d1cd4a3..e63ddfae71 100644
--- a/tests/test_pd_pp.py
+++ b/tests/test_pd_pp.py
@@ -7,9 +7,11 @@
import threading
import types
from collections import deque
+from itertools import pairwise
from types import SimpleNamespace
from unittest.mock import MagicMock
+import numpy as np
import pytest
# Ensure aiter.dist.parallel_state exposes symbols mooncake_connector needs.
@@ -170,10 +172,12 @@ def test_producer_advertises_remote_pp_size():
assert seq.kv_transfer_params_output["remote_block_ids"] == [1, 2, 3]
-def _mooncake_consumer_scheduler(mc, hash_block_size=64):
+def _mooncake_consumer_scheduler(mc, block_size=64, dcp_size=1):
sched = object.__new__(mc.MooncakeConnectorScheduler)
sched.is_producer = False
- sched.hash_block_size = hash_block_size
+ sched.block_size = block_size
+ sched.dcp_size = dcp_size
+ sched.hash_block_size = block_size * dcp_size
sched.request_id_to_transfer_id = {}
sched.transfer_id_to_request_id = {}
sched._reqs_need_recv = {}
@@ -195,7 +199,7 @@ def _remote_prefill_seq(remote_hash_block_size):
)
-def test_matching_hash_block_size_enables_incremental_transfer():
+def test_matching_block_size_enables_incremental_transfer():
mc = pytest.importorskip(
"atom.kv_transfer.disaggregation.mooncake.mooncake_connector"
)
@@ -207,8 +211,23 @@ def test_matching_hash_block_size_enables_incremental_transfer():
assert seq.kv_transfer_params["num_computed_blocks"] == 2
+def test_dcp_consumer_stays_incremental_against_a_non_dcp_producer():
+ # CPP prefill (dcp=1, 16-token blocks) -> DCP decode (dcp=4). The consumer
+ # addresses 64-token virtual blocks, so the two sides' hash_block_size
+ # differ by exactly dcp_size; the block slicing applies that factor itself.
+ mc = pytest.importorskip(
+ "atom.kv_transfer.disaggregation.mooncake.mooncake_connector"
+ )
+ sched = _mooncake_consumer_scheduler(mc, block_size=16, dcp_size=4)
+ seq = _remote_prefill_seq(remote_hash_block_size=16)
+
+ sched.update_state_after_alloc(seq)
+
+ assert seq.kv_transfer_params["num_computed_blocks"] == 2
+
+
@pytest.mark.parametrize("remote_hash_block_size", [32, None])
-def test_mismatched_or_missing_hash_block_size_forces_full_transfer(
+def test_mismatched_or_missing_block_size_forces_full_transfer(
remote_hash_block_size, caplog
):
mc = pytest.importorskip(
@@ -665,3 +684,239 @@ def test_pp_downstream_skips_forward_for_request_less_batch():
]
assert forwards == []
stage.pp_transport.send_tokens.assert_not_called()
+
+
+# ---------------------------------------------------------------------------
+# DCP relayout planners
+# ---------------------------------------------------------------------------
+
+
+def _expand(plan):
+ """Runs -> the (src, dst) unit pairs they actually move."""
+ src, dst, length = plan
+ return [
+ (int(s) + i, int(d) + i)
+ for s, d, n in zip(src, dst, length)
+ for i in range(int(n))
+ ]
+
+
+def _sharded_reference(src_ids, dst_ids, block_size, dcp_size, rank, interleave):
+ """Per-token expectation, taken from the write-side ownership rule.
+
+ Local slot ``s`` on rank ``r`` holds global position
+ ``((s // S) * W + r) * S + (s % S)``, which the producer stored in its own
+ block table at ``g // block_size``. Deriving it a token at a time is what
+ makes this independent of the planner's run arithmetic.
+ """
+ pairs = []
+ for slot in range(len(dst_ids) * block_size):
+ g = ((slot // interleave) * dcp_size + rank) * interleave + slot % interleave
+ if g // block_size >= len(src_ids):
+ continue
+ pairs.append(
+ (
+ src_ids[g // block_size] * block_size + g % block_size,
+ dst_ids[slot // block_size] * block_size + slot % block_size,
+ )
+ )
+ return pairs
+
+
+@pytest.mark.parametrize(
+ "name, src_ids, dst_ids, block_size, dcp_size, rank, interleave",
+ [
+ ("round robin, rank 0", [5, 6, 7, 8], [2, 3], 2, 4, 0, 1),
+ ("round robin, last rank", [5, 6, 7, 8], [2, 3], 2, 4, 3, 1),
+ ("interleave 2", [11, 12, 13, 14, 15, 16, 17, 18], [3, 9], 8, 4, 2, 2),
+ ("whole-block interleave", [0, 1, 2, 3], [7, 4], 4, 2, 1, 4),
+ # The block manager sizes every rank's table from rank 0's share, so a
+ # higher rank can own a trailing virtual block with no source tokens.
+ ("trailing block has no source", [5, 6, 7], [2, 3], 2, 4, 2, 1),
+ ("shuffled ids", [9, 2, 30, 4], [8, 1], 2, 4, 1, 1),
+ ],
+)
+def test_plan_sharded_matches_the_write_side_ownership(
+ name, src_ids, dst_ids, block_size, dcp_size, rank, interleave
+):
+ mc = pytest.importorskip(
+ "atom.kv_transfer.disaggregation.mooncake.mooncake_connector"
+ )
+ plan = mc.plan_sharded(src_ids, dst_ids, block_size, dcp_size, rank, interleave)
+ assert _expand(plan) == _sharded_reference(
+ src_ids, dst_ids, block_size, dcp_size, rank, interleave
+ )
+
+
+def test_plan_sharded_ranks_partition_the_producer_tokens():
+ mc = pytest.importorskip(
+ "atom.kv_transfer.disaggregation.mooncake.mooncake_connector"
+ )
+ block_size, dcp_size, interleave = 8, 4, 2
+ dst_ids = [3, 9]
+ src_ids = [11, 12, 13, 14, 15, 16, 17, 18] # exactly the global width
+
+ moved = []
+ for rank in range(dcp_size):
+ plan = mc.plan_sharded(src_ids, dst_ids, block_size, dcp_size, rank, interleave)
+ moved += [s for s, _ in _expand(plan)]
+
+ every_token = [b * block_size + t for b in src_ids for t in range(block_size)]
+ assert sorted(moved) == sorted(every_token)
+
+
+@pytest.mark.parametrize("interleave", [0, 3, 5, 12])
+def test_plan_sharded_rejects_an_interleave_that_does_not_tile_a_block(interleave):
+ mc = pytest.importorskip(
+ "atom.kv_transfer.disaggregation.mooncake.mooncake_connector"
+ )
+ with pytest.raises(ValueError, match="must divide block_size"):
+ mc.plan_sharded([0, 1], [0], 8, 4, 0, interleave)
+
+
+def test_coalesce_merges_only_runs_contiguous_on_both_sides():
+ mc = pytest.importorskip(
+ "atom.kv_transfer.disaggregation.mooncake.mooncake_connector"
+ )
+ src, dst, length = mc._coalesce(
+ np.array([0, 4, 8, 100]),
+ np.array([0, 4, 20, 24]), # the 8 -> 20 hop breaks the destination run
+ np.array([4, 4, 4, 4]),
+ )
+ assert src.tolist() == [0, 8, 100]
+ assert dst.tolist() == [0, 20, 24]
+ assert length.tolist() == [8, 4, 4]
+
+
+def test_coalesce_handles_an_empty_plan():
+ mc = pytest.importorskip(
+ "atom.kv_transfer.disaggregation.mooncake.mooncake_connector"
+ )
+ src, dst, length = mc._coalesce(
+ np.empty(0, dtype=np.int64),
+ np.empty(0, dtype=np.int64),
+ np.empty(0, dtype=np.int64),
+ )
+ assert src.size == dst.size == length.size == 0
+
+
+def test_plan_replicated_index_splits_each_page_into_a_key_and_a_scale_run():
+ mc = pytest.importorskip(
+ "atom.kv_transfer.disaggregation.mooncake.mooncake_connector"
+ )
+ key, scale, dcp_size = 512, 64, 4
+ page = key + scale
+ src_ids = [10, 11, 12, 13, 14] # 5 pages for 8 destination sub-pages
+ dst_ids = [2, 7]
+
+ src, dst, length = mc.plan_replicated_index(
+ src_ids, dst_ids, dcp_size, page, key, scale
+ )
+
+ kept_dst = [2, 2, 2, 2, 7] # the last 3 sub-pages have no source page
+ kept_sub = [0, 1, 2, 3, 0]
+ assert src.tolist() == [b * page for b in src_ids] + [
+ b * page + key for b in src_ids
+ ]
+ assert dst.tolist() == [
+ d * page * dcp_size + s * key for d, s in zip(kept_dst, kept_sub)
+ ] + [
+ d * page * dcp_size + dcp_size * key + s * scale
+ for d, s in zip(kept_dst, kept_sub)
+ ]
+ assert length.tolist() == [key] * 5 + [scale] * 5
+
+ # Distinct source pages must not overwrite each other on the destination.
+ runs = sorted(zip(dst.tolist(), length.tolist()))
+ assert all(a + n <= b for (a, n), (b, _) in pairwise(runs))
+
+
+def _dcp_block_transfer_producer(mc, roles):
+ """A pp1 producer with one region per role, ready for _execute_block_transfer."""
+ conn = object.__new__(mc.MooncakeConnector)
+ conn.pp_size = 1
+ conn.pp_rank = 0
+ conn._num_local_layers = len(roles)
+ conn._start_layer = 0
+ conn._block_region_consumer_indices = None
+ conn._block_region_roles = list(roles)
+ conn._block_region_planes = [None] * len(roles)
+ conn.kv_caches_base_addr = [0x1000 * (i + 1) for i in range(len(roles))]
+ conn._per_block_bytes_list = [576 * 16] * len(roles)
+ conn.block_size = 16
+ conn.dcp_size = 1
+ return conn
+
+
+def _dcp_block_transfer_request(conn, roles):
+ return {
+ "consumer_base_addrs": [0x9000 * (i + 1) for i in range(len(roles))],
+ "consumer_num_layers": len(roles),
+ "consumer_region_roles": list(roles),
+ "consumer_dcp_size": 4,
+ "consumer_dcp_rank": 0,
+ "consumer_dcp_interleave": 16,
+ "consumer_replicates_index_cache": False,
+ }
+
+
+def test_block_transfer_refuses_a_roleless_region_under_dcp():
+ mc = pytest.importorskip(
+ "atom.kv_transfer.disaggregation.mooncake.mooncake_connector"
+ )
+ # An EAGLE3 draft region is appended to the MLA builder's list with no
+ # semantic_role; the token-interleave relayout does not describe its layout.
+ roles = [mc.MLA_KV_ROLE, None]
+ conn = _dcp_block_transfer_producer(mc, roles)
+
+ with pytest.raises(RuntimeError, match="semantic_role"):
+ conn._execute_block_transfer(
+ _dcp_block_transfer_request(conn, roles),
+ "host:1",
+ list(range(8)),
+ [0, 1],
+ "req-0",
+ )
+
+
+def test_block_transfer_descriptors_match_scalar_arithmetic():
+ mc = pytest.importorskip(
+ "atom.kv_transfer.disaggregation.mooncake.mooncake_connector"
+ )
+ roles = [mc.MLA_KV_ROLE, mc.MLA_KV_ROLE]
+ conn = _dcp_block_transfer_producer(mc, roles)
+ request = _dcp_block_transfer_request(conn, roles)
+ src_block_ids, dst_block_ids = list(range(8)), [0, 1]
+ captured = {}
+
+ def _capture(target, src_addrs, dst_addrs, sizes, req_id, kind):
+ captured.update(src=src_addrs, dst=dst_addrs, sizes=sizes)
+ return True
+
+ conn._rdma_write_with_retry = _capture
+ assert conn._execute_block_transfer(
+ request, "host:1", src_block_ids, dst_block_ids, "req-0"
+ )
+
+ plan = mc.plan_sharded(
+ src_block_ids,
+ dst_block_ids,
+ conn.block_size,
+ request["consumer_dcp_size"],
+ request["consumer_dcp_rank"],
+ request["consumer_dcp_interleave"],
+ )
+ unit = conn._per_block_bytes_list[0] // conn.block_size
+ exp_src, exp_dst, exp_sizes = [], [], []
+ for region_idx in range(len(roles)):
+ for src_off, dst_off, run_len in zip(*plan):
+ exp_src.append(conn.kv_caches_base_addr[region_idx] + int(src_off) * unit)
+ exp_dst.append(
+ request["consumer_base_addrs"][region_idx] + int(dst_off) * unit
+ )
+ exp_sizes.append(int(run_len) * unit)
+
+ assert captured["src"] == exp_src
+ assert captured["dst"] == exp_dst
+ assert captured["sizes"] == exp_sizes
+ assert all(type(a) is int for a in captured["src"])