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"])