diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 733128b3c..dcaf4c95d 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -25,8 +25,8 @@ class StateType(str, enum.Enum): # DeepSeek-V4 unified_kv SWA ring: addressed per-row by ring slot # (req_pool_idx * ring_stride + pos % ring_stride), needs its own component. SWA_RING = "swa_ring" - # DeepSeek-V4 online C128 request-scoped state. - C128_STATE = "c128_state" + # DeepSeek-V4 request-scoped compression state; preserve the legacy wire value. + DSV4_REQUEST_STATE = "c128_state" # A block-scaled KV dtype keeps its per-block scales in buffers parallel to # K/V, one component per sub-pool so each carries the index payload of the # KV it describes (whole sequence for full attention, window for SWA). diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 7f486522e..5ebcaf811 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -972,10 +972,8 @@ class CommonKVManager(BaseKVManager): mla_ratios = getattr(self.kv_args, "mla_compression_ratios", None) if mla_ratios: - # Compressed-MLA (e.g. DeepSeek V4): the flat list is organized - # by buffer type (compression-ratio bucket) rather than by - # layer, so we locate the sub-range for this PP stage inside each - # section of the dst flat list. + # Compressed-MLA pointers are grouped by buffer type; + # each group needs its own PP-stage slice. sliced_src_kv_ptrs, sliced_dst_kv_ptrs = self._mla_slice_ptrs_for_pp( src_kv_ptrs, dst_kv_ptrs, mla_ratios, state_type ) @@ -999,37 +997,11 @@ class CommonKVManager(BaseKVManager): mla_ratios: List[int], state_type: Optional[StateType] = None, ) -> Tuple[List[int], List[int]]: - """Produce aligned (src, dst) pointer lists for compressed-MLA - pools (e.g. DeepSeek V4) under PP. - - The pool produces two possible flat-list layouts (selected via dst - length): - - - kv_data layout, length = 2 * c4_L + c128_L: - [c4_layer_{0..c4_L-1}, - c4_indexer_layer_{0..c4_L-1}, - c128_layer_{0..c128_L-1}] - Each section is indexed by compressed-layer id within that - compression bucket. - - - SWA state_data layout, length = swa_L + 2 * c4_L: - [swa_layer_{0..swa_L-1}, - c4_compress_state_{0..c4_L-1}, - c4_indexer_compress_state_{0..c4_L-1}] - ``swa_L`` is the SWA pool's actual buffer count - (``num_effective_layers``), which can be smaller than - ``len(mla_ratios)`` when the HF config's ``compress_ratios`` - list contains entries for layers not materialized into the SWA - pool (e.g. an MTP/nextn slot at the tail). - - - C128_STATE layout, length = c128_L: - [c128_compress_state_{0..c128_L-1}] - - src is already PP-filtered on the prefill side. dst is the - decode-side full-model list (when decode is PP=1). We slice dst to - match src's PP stage. If src itself is also full-model, it is - returned unchanged. - """ + # Match the pool's flat buffer order, with layers grouped by ratio: + # KV: [C4 KV, C4 indexer KV, C128 KV]. + # SWA state: [SWA KV, C4 compressor state, C4 indexer state]. + # DSV4_REQUEST_STATE: [C128 compressor state]; SWA_RING: [SWA rings]. + # Prefill src is stage-local; decode dst may cover the full model. start_layer = self.kv_args.prefill_start_layer end_layer = self.kv_args.prefill_end_layer assert end_layer is not None, ( @@ -1045,7 +1017,7 @@ class CommonKVManager(BaseKVManager): c128_off_s = sum(1 for r in mla_ratios[:start_layer] if r == 128) c128_off_e = sum(1 for r in mla_ratios[:end_layer] if r == 128) - if state_type == StateType.C128_STATE: + if state_type == StateType.DSV4_REQUEST_STATE: return src_kv_ptrs, list(dst_kv_ptrs[c128_off_s:c128_off_e]) if state_type == StateType.SWA_RING: @@ -1054,7 +1026,8 @@ class CommonKVManager(BaseKVManager): return src_kv_ptrs, list(dst_kv_ptrs[swa_s:swa_e]) if ( - state_type not in (StateType.SWA, StateType.SWA_RING, StateType.C128_STATE) + state_type + not in (StateType.SWA, StateType.SWA_RING, StateType.DSV4_REQUEST_STATE) and len(dst_kv_ptrs) == kv_layout_len ): sliced_dst = ( @@ -1064,11 +1037,8 @@ class CommonKVManager(BaseKVManager): ) return src_kv_ptrs, sliced_dst - # SWA state-data layout. ``swa_L`` is derived from the actual dst - # length so we tolerate cases where the SWA pool has fewer buffers - # than ``len(mla_ratios)`` (e.g. nextn padding). C128 state ships as - # a separate StateType.C128_STATE component and must not be counted - # here. + # SWA may omit nextn entries present in mla_ratios; use its buffer count. + # C128 state ships separately as DSV4_REQUEST_STATE. swa_L = len(dst_kv_ptrs) - 2 * c4_full if swa_L < 0 or swa_L > len(mla_ratios): raise ValueError( diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 9ac2121e3..b16c5cec2 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -1449,7 +1449,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): ) state_types = self.kv_manager.kv_args.state_types - if StateType.C128_STATE in state_types: + if StateType.DSV4_REQUEST_STATE in state_types: clear_c128_state = getattr( self.token_to_kv_pool, "clear_c128_req_state", None ) @@ -1462,7 +1462,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): StateType.DSA_TAIL: _dsa_tail_payload, StateType.MINIMAX_INDEX_K: _full_kv_pages_payload, StateType.SWA_RING: _swa_ring_payload, - StateType.C128_STATE: _c128_state_payload, + StateType.DSV4_REQUEST_STATE: _c128_state_payload, StateType.BLOCK_SCALE: _full_kv_pages_payload, StateType.BLOCK_SCALE_SWA: _swa_payload, } diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 0f9e9ce38..a37600e44 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -1343,14 +1343,14 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): StateType.SWA, StateType.DSA, StateType.SWA_RING, - StateType.C128_STATE, + StateType.DSV4_REQUEST_STATE, StateType.BLOCK_SCALE, StateType.BLOCK_SCALE_SWA, ) def _requires_exact_state_index_match(self, st: StateType) -> bool: """State types whose page lists are positional and must not be truncated.""" - return st in (StateType.SWA_RING, StateType.C128_STATE) + return st in (StateType.SWA_RING, StateType.DSV4_REQUEST_STATE) def maybe_send_extra( self, @@ -1495,7 +1495,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): src_indices = list(indices) dst_indices_local = list(dst_indices) if ( - st == StateType.C128_STATE + st == StateType.DSV4_REQUEST_STATE and len(src_indices) == 0 and len(dst_indices_local) == 0 ): diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index a7fc753ad..6bf84a189 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -2458,13 +2458,17 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): state_type=st, force_flat=True, ) - elif st in (StateType.SWA, StateType.SWA_RING, StateType.C128_STATE): + elif st in ( + StateType.SWA, + StateType.SWA_RING, + StateType.DSV4_REQUEST_STATE, + ): if not self.is_mla_backend and self.attn_tp_size != decode_tp_size: raise RuntimeError( f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {st.upper()} hybrid models yet." ) if ( - st == StateType.C128_STATE + st == StateType.DSV4_REQUEST_STATE and len(src_indices) == 0 and len(dst_indices) == 0 ): diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index a797d375a..1afebb284 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -1384,7 +1384,7 @@ class SchedulerDisaggregationPrefillMixin: StateType.DSA_TAIL: _dsa_tail_payload, StateType.MINIMAX_INDEX_K: _full_kv_pages_payload, StateType.SWA_RING: _swa_ring_payload, - StateType.C128_STATE: _c128_state_payload, + StateType.DSV4_REQUEST_STATE: _c128_state_payload, StateType.BLOCK_SCALE: _full_kv_pages_payload, StateType.BLOCK_SCALE_SWA: _swa_payload, } diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index ca0919bfa..4ed40446d 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -89,13 +89,8 @@ def get_dsv4_c4_state_indices( *, ring_size: int, ) -> np.ndarray: - """Return physical rows for the live C4 compressor history. - - Prefill and decode may use different C4 ring sizes (8 without speculative - decoding and 16 with EAGLE/MTP). State transfer must therefore pair rows - by logical token position instead of copying a whole request-local bank. - The C4 overlap compressor keeps ``seq_len % 4 + 4`` live rows. - """ + # Prefill and decode can have different ring sizes (8 or 16 with EAGLE/MTP); + # pair the overlap compressor's live rows by logical token position. if ring_size < 8 or ring_size % 4 != 0: raise ValueError( f"C4 ring_size must be a multiple of 4 and at least 8, got {ring_size}" @@ -1276,10 +1271,6 @@ def setup_state_kv_args( total_kv_layers: int = None, req_to_token_pool=None, ) -> None: - """Populate ``kv_args`` state-buffer fields from the given pool. - Shared by prefill and decode bootstrap paths so the state_type dispatch - lives in one place. - """ from sglang.srt.disaggregation.base.conn import StateType from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMLATokenToKVPool from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool @@ -1373,14 +1364,14 @@ def setup_state_kv_args( ring_lens, ring_item_lens, ) - if hasattr(token_to_kv_pool, "get_c128_state_buf_infos"): + if hasattr(token_to_kv_pool, "get_request_state_buf_infos"): c128_ptrs, c128_lens, c128_item_lens = ( - token_to_kv_pool.get_c128_state_buf_infos() + token_to_kv_pool.get_request_state_buf_infos() ) if c128_ptrs: append_state_component( kv_args, - StateType.C128_STATE, + StateType.DSV4_REQUEST_STATE, c128_ptrs, c128_lens, c128_item_lens, diff --git a/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_common_hooks.py b/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_common_hooks.py index 0e0033efb..c6045a06d 100644 --- a/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_common_hooks.py +++ b/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_common_hooks.py @@ -75,7 +75,7 @@ def dsv4_state_payloads( """Build NPU-specific DSV4 PD payloads. Returns payloads for components that are addressed differently from the - cross-hardware ``StateType.SWA`` / ``StateType.C128_STATE`` defaults: + cross-hardware ``StateType.SWA`` / ``StateType.DSV4_REQUEST_STATE`` defaults: * ``DSV4_C128`` — C128 KV pages from ``req_to_c128_sidecar``. * ``DSV4_C4_STATE`` (A5 only) — live C4 compress-state rows. Prefill diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index 799834048..3a3fc1f8e 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -179,8 +179,7 @@ class DSV4AttnMetadata: swa_topk_lengths: torch.Tensor c4_sparse_topk: int - # SWA KV-store write target (out_cache_loc translated to SWA space), computed - # once per iteration in make_core_attn_metadata and read by the store path. + # Shared by all layer stores; locations are in SWA space. swa_out_cache_loc: Optional[torch.Tensor] = None c4_out_loc: Optional[torch.Tensor] = None c4_topk_lengths_raw: Optional[torch.Tensor] = None @@ -205,7 +204,7 @@ class DSV4AttnMetadata: trtllm_prefill_c4_indices: Optional[torch.Tensor] = None trtllm_prefill_c128: Optional[tuple] = None - c1_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) + c0_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) c4_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) c128_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) @@ -215,7 +214,7 @@ class DSV4AttnMetadata: def get_flashmla_metadata(self, compress_ratio: Literal[0, 4, 128]): if compress_ratio == 0: - return self.c1_flashmla_metadata + return self.c0_flashmla_metadata elif compress_ratio == 4: return self.c4_flashmla_metadata elif compress_ratio == 128: @@ -258,7 +257,7 @@ class DSV4AttnMetadata: # Recomputed by the recorded init_forward_metadata_in_graph op # each forward; not copied across replays. "swa_out_cache_loc", - "c1_flashmla_metadata", + "c0_flashmla_metadata", "c4_flashmla_metadata", "c128_flashmla_metadata", # Eager-only lazy caches are assigned, not content-copied. @@ -296,7 +295,7 @@ class DSV4AttnMetadata: "swa_topk_lengths", "c128_page_indices", "c128_topk_lengths_clamp1", - "c1_flashmla_metadata", + "c0_flashmla_metadata", "c4_flashmla_metadata", "c128_flashmla_metadata", # Reset eager-only caches so a replay cannot reuse another shape. @@ -428,7 +427,7 @@ class DSV4AttnMetadata: self.c4_sparse_page_indices = _pad_last_dim(self.c4_sparse_page_indices) if is_prefill: self.c4_sparse_raw_indices = torch.empty_like(self.c4_sparse_page_indices) - self.c1_flashmla_metadata = _create_flashmla_metadata() + self.c0_flashmla_metadata = _create_flashmla_metadata() self.c4_flashmla_metadata = _create_flashmla_metadata() self.c128_flashmla_metadata = _create_flashmla_metadata() @@ -780,7 +779,7 @@ class DeepseekV4AttnBackend( return PagedIndexerMetadata( page_size=self.page_size, page_table=core_attn_metadata.page_table, - c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw, + compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw, use_topk_v2=self.dsa_topk_backend.should_use_topk_v2() and not _is_xpu, # The SM120 FP4 kernel schedules split_kv=128, while the generic # JIT metadata planner encodes split_kv=256. @@ -1176,9 +1175,7 @@ class DeepseekV4AttnBackend( return buf[:n] def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: - # Upgrade Raw->Full so the c4/c128 compress + core_attn + indexer - # materialization is recorded inside the cuda graph; a no-op (Full - # already) when PREP_IN_CUDA_GRAPH=0. + # Raw metadata must be materialized inside the graph to refresh on replay. if isinstance(self.forward_metadata, DSV4RawVerifyMetadata): self.forward_metadata = self.make_forward_metadata_from_raw_verify( raw_metadata=self.forward_metadata, @@ -1189,11 +1186,9 @@ class DeepseekV4AttnBackend( raw_metadata=self.forward_metadata, ) - # Compute the SWA KV-store write target once per forward and cache it on - # the metadata for every layer's store. This is recorded inside the cuda - # graph, so replay re-reads the live out_cache_loc buffer (spec-v2 and DP - # padding rebind out_cache_loc after out-graph metadata prep). flash_mla - # kernels require int32 indices. + # Spec-v2 and DP padding can rebind out_cache_loc after out-graph prep; + # capture the translation here so replay reads live locations. + # FlashMLA requires int32 indices. metadata = self.forward_metadata if ( isinstance(metadata, DSV4Metadata) @@ -1703,7 +1698,7 @@ class DeepseekV4AttnBackend( metadata.core_attn_metadata, DSV4AttnMetadata ): core = metadata.core_attn_metadata - core.c1_flashmla_metadata = _create_flashmla_metadata() + core.c0_flashmla_metadata = _create_flashmla_metadata() core.c4_flashmla_metadata = _create_flashmla_metadata() core.c128_flashmla_metadata = _create_flashmla_metadata() @@ -1714,15 +1709,7 @@ class DeepseekV4AttnBackend( self.forward_metadata = current_raw def get_swa_out_cache_loc(self, forward_batch: ForwardBatch) -> torch.Tensor: - """Resolve the SWA KV-store write target for the current forward. - - Prefer the value cached by the metadata init: in-graph for - decode/verify, the hoisted cuda_graph_swa_out_cache_loc buffer for - draft-extend. Translate at store time when nothing matching is cached - (paths that skip the init, or a batch re-padded after init). Idle - always falls back: its metadata may be stale, and - translating the zero-padded out_cache_loc writes to the dummy slot. - """ + # Idle metadata may be stale; zero-padded locations target the dummy slot. out_cache_loc = forward_batch.out_cache_loc core = getattr(self.forward_metadata, "core_attn_metadata", None) cached = core.swa_out_cache_loc if core is not None else None @@ -2366,7 +2353,7 @@ class DeepseekV4AttnBackend( core_attn_metadata.c4_sparse_topk_lengths = None core_attn_metadata.c4_sparse_page_indices = None core_attn_metadata.c4_sparse_raw_indices = None - core_attn_metadata.c1_flashmla_metadata = _create_flashmla_metadata() + core_attn_metadata.c0_flashmla_metadata = _create_flashmla_metadata() core_attn_metadata.c4_flashmla_metadata = None core_attn_metadata.c128_flashmla_metadata = None if self.trtllm_attn: diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py index 8be31dff2..f78bd1c73 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -162,8 +162,7 @@ class DSV4AttnMetadata: swa_topk_lengths: torch.Tensor c4_sparse_topk: int - # SWA KV-store write target (out_cache_loc translated to SWA space), computed - # once per iteration in make_core_attn_metadata and read by the store path. + # Shared by all layer stores; locations are in SWA space. swa_out_cache_loc: Optional[torch.Tensor] = None c4_out_loc: Optional[torch.Tensor] = None c4_topk_lengths_raw: Optional[torch.Tensor] = None @@ -181,7 +180,7 @@ class DSV4AttnMetadata: # unified-kv metadata unified: Optional[UnifiedKvMetadata] = None - c1_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) + c0_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) c4_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) c128_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) @@ -191,7 +190,7 @@ class DSV4AttnMetadata: def get_flashmla_metadata(self, compress_ratio: Literal[0, 4, 128]): if compress_ratio == 0: - return self.c1_flashmla_metadata + return self.c0_flashmla_metadata elif compress_ratio == 4: return self.c4_flashmla_metadata elif compress_ratio == 128: @@ -232,7 +231,7 @@ class DSV4AttnMetadata: # Recomputed by the recorded init_forward_metadata_in_graph op # each forward; not copied across replays. "swa_out_cache_loc", - "c1_flashmla_metadata", + "c0_flashmla_metadata", "c4_flashmla_metadata", "c128_flashmla_metadata", ], @@ -347,7 +346,7 @@ class DSV4AttnMetadata: self.c4_sparse_page_indices = _pad_last_dim(self.c4_sparse_page_indices) if is_prefill: self.c4_sparse_raw_indices = torch.empty_like(self.c4_sparse_page_indices) - self.c1_flashmla_metadata = _create_flashmla_metadata() + self.c0_flashmla_metadata = _create_flashmla_metadata() self.c4_flashmla_metadata = _create_flashmla_metadata() self.c128_flashmla_metadata = _create_flashmla_metadata() @@ -509,7 +508,7 @@ class DeepseekV4HipRadixBackend( return PagedIndexerMetadata( page_size=self.page_size, page_table=core_attn_metadata.page_table, - c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw, + compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw, use_topk_v2=self.dsa_topk_backend.should_use_topk_v2(), ) @@ -842,9 +841,7 @@ class DeepseekV4HipRadixBackend( ) def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: - # Upgrade Raw->Full so the c4/c128 compress + core_attn + indexer - # materialization is recorded inside the cuda graph; a no-op (Full - # already) when PREP_IN_CUDA_GRAPH=0. + # Raw metadata must be materialized inside the graph to refresh on replay. if isinstance(self.forward_metadata, DSV4RawVerifyMetadata): self.forward_metadata = self.make_forward_metadata_from_raw_verify( raw_metadata=self.forward_metadata, @@ -854,11 +851,9 @@ class DeepseekV4HipRadixBackend( raw_metadata=self.forward_metadata, ) - # Compute the SWA KV-store write target once per forward and cache it on - # the metadata for every layer's store. This is recorded inside the cuda - # graph, so replay re-reads the live out_cache_loc buffer (spec-v2 and DP - # padding rebind out_cache_loc after out-graph metadata prep). flash_mla - # kernels require int32 indices. + # Spec-v2 and DP padding can rebind out_cache_loc after out-graph prep; + # capture the translation here so replay reads live locations. + # FlashMLA requires int32 indices. metadata = self.forward_metadata if ( isinstance(metadata, DSV4Metadata) @@ -911,7 +906,7 @@ class DeepseekV4HipRadixBackend( indexer_metadata = metadata.indexer_metadata metadata.fp4_decode_workspace = prepare_fp4_decode_workspace( indexer_metadata.page_table, - indexer_metadata.c4_seq_lens, + indexer_metadata.compressed_seq_lens, ) def _fp4_workspaces_enabled(self, metadata) -> bool: @@ -957,7 +952,7 @@ class DeepseekV4HipRadixBackend( indexer_metadata = metadata.indexer_metadata metadata.fp4_prefill_workspace = prepare_fp4_prefill_workspace( indexer_metadata.page_table, - indexer_metadata.c4_seq_lens, + indexer_metadata.compressed_seq_lens, workspace=metadata.fp4_prefill_workspace, ) @@ -1215,7 +1210,7 @@ class DeepseekV4HipRadixBackend( metadata.core_attn_metadata, DSV4AttnMetadata ): core = metadata.core_attn_metadata - core.c1_flashmla_metadata = _create_flashmla_metadata() + core.c0_flashmla_metadata = _create_flashmla_metadata() core.c4_flashmla_metadata = _create_flashmla_metadata() core.c128_flashmla_metadata = _create_flashmla_metadata() @@ -1228,12 +1223,8 @@ class DeepseekV4HipRadixBackend( def _attach_unified_kv_decode_streams( self, core: DSV4AttnMetadata, state_slot: torch.Tensor ) -> None: - """build the ragged decode index streams once per forward. - - ``state_slot`` is the per-row req-slot map: decode passes - ``req_pool_indices`` (1 token per req), target-verify passes - ``req_pool_indices_repeated`` (the per-token num_draft*bs -> bs map) so - the same builder produces per-draft-token decode streams.""" + # state_slot maps each query token to its request slot; + # target-verify repeats request slots for the draft tokens. from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, ) @@ -1502,18 +1493,7 @@ class DeepseekV4HipRadixBackend( return o def get_swa_out_cache_loc(self, forward_batch: ForwardBatch) -> torch.Tensor: - """Resolve the SWA KV-store write target for the current forward. - - Fast path: the per-forward value cached by init_forward_metadata_in_graph - (recorded inside cuda graphs, so replay re-reads live buffers). Fallback: - translate at store time, matching the pre-cache behavior, for paths that - never run the in-graph init — eager idle (forward_idle skips attn init), - runners that only run the out-graph prep (e.g. - EAGLEDraftExtendCudaGraphRunner) — or whose batch was re-padded after - init (shape mismatch). Idle always falls back: its metadata is absent or - left over from a previous forward, and translating the zero-padded - out_cache_loc writes to the dummy slot. - """ + # Idle metadata may be stale; zero-padded locations target the dummy slot. out_cache_loc = forward_batch.out_cache_loc core = getattr(self.forward_metadata, "core_attn_metadata", None) cached = core.swa_out_cache_loc if core is not None else None @@ -1528,18 +1508,7 @@ class DeepseekV4HipRadixBackend( ) def get_unified_swa_loc(self, forward_batch: ForwardBatch) -> torch.Tensor: - """SWA ring write target for unified_kv, shared by all layers. - - Fast path: the per-forward value cached in _attach_unified_kv_decode_streams - (recorded inside cuda graphs, so replay re-reads live buffers). Fallback: - recompute at store time, matching the pre-cache per-layer behavior, for - paths that never ran the decode-stream init (eager prefill/extend, idle, - or a batch re-padded after init -> shape mismatch). - - Cached swa_loc is computed once from committed positions, so every draft-decode - step would reuse the same ring slot and break the chain. Recompute from the live - per-step positions; only the draft path is affected, the rest keeps the fast path. - """ + # Cached slots use committed positions; draft steps need live positions. positions = forward_batch.positions core = getattr(self.forward_metadata, "core_attn_metadata", None) unified = getattr(core, "unified", None) if core is not None else None @@ -1818,7 +1787,7 @@ class DeepseekV4HipRadixBackend( core_attn_metadata.c4_sparse_topk_lengths_raw = None core_attn_metadata.c4_sparse_page_indices = None core_attn_metadata.c4_sparse_raw_indices = None - core_attn_metadata.c1_flashmla_metadata = _create_flashmla_metadata() + core_attn_metadata.c0_flashmla_metadata = _create_flashmla_metadata() core_attn_metadata.c4_flashmla_metadata = None core_attn_metadata.c128_flashmla_metadata = None return core_attn_metadata diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index 1be191040..0e4085a91 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -632,7 +632,7 @@ class C4IndexerBackendMixin: # reading logits, so DeepGEMM can receive an empty range for them. if self.dsa_topk_backend.is_sgl_kernel(): ke = torch.where(ke - ks > c4_indexer.index_topk, ke, ks) - c4_page_size = indexer_metadata.c4_page_size + c4_page_size = indexer_metadata.compressed_page_size max_seqlen_k = (final_c4_len + c4_page_size - 1) // c4_page_size * c4_page_size plan = NonPagedIndexerPlan( page_table=request_page_table, @@ -793,7 +793,7 @@ class C4IndexerBackendMixin: return F.pad(tensor, pad, value=value) c4_seq_lens = match_num_queries( - indexer_metadata.c4_seq_lens, value=0 if use_aiter_fp4 else 1 + indexer_metadata.compressed_seq_lens, value=0 if use_aiter_fp4 else 1 ) _c4sl = c4_seq_lens page_table = match_num_queries(indexer_metadata.page_table, value=0) @@ -856,7 +856,7 @@ class C4IndexerBackendMixin: c4_seq_lens[rows], page_table[rows], c4_sparse_page_indices[rows], - indexer_metadata.c4_page_size, + indexer_metadata.compressed_page_size, row_raw_indices, ) elif self.dsa_topk_backend.is_flashinfer(): @@ -865,7 +865,7 @@ class C4IndexerBackendMixin: c4_seq_lens[rows], page_table[rows], c4_sparse_page_indices[rows], - indexer_metadata.c4_page_size, + indexer_metadata.compressed_page_size, row_raw_indices, ) elif self.dsa_topk_backend.should_use_topk_v2() and raw_indices is None: @@ -874,7 +874,7 @@ class C4IndexerBackendMixin: c4_seq_lens[rows], page_table[rows], c4_sparse_page_indices[rows], - indexer_metadata.c4_page_size, + indexer_metadata.compressed_page_size, # The cached plan routes rows by their index in the full # range, so a chunk needs one built over its own rows. ( @@ -889,7 +889,7 @@ class C4IndexerBackendMixin: c4_seq_lens[rows], page_table[rows], c4_sparse_page_indices[rows], - indexer_metadata.c4_page_size, + indexer_metadata.compressed_page_size, row_raw_indices, ) @@ -964,7 +964,7 @@ class C4IndexerBackendMixin: _c4sl[rows], page_table[rows], metadata, - indexer_metadata.max_c4_seq_len, + indexer_metadata.max_compressed_seq_len, False, ) run_topk_transform(rows, logits) @@ -990,7 +990,7 @@ class C4IndexerBackendMixin: core_metadata.c4_sparse_page_indices = ( hisparse_coordinator.swap_in_selected_pages( req_pool_indices=forward_batch.req_pool_indices, - compressed_seq_lens=indexer_metadata.c4_seq_lens, + compressed_seq_lens=indexer_metadata.compressed_seq_lens, top_k_result=raw_indices, layer_id=compress_layer_id, ) diff --git a/python/sglang/srt/layers/attention/dsv4/metadata.py b/python/sglang/srt/layers/attention/dsv4/metadata.py index bf7ce9670..67e07be27 100644 --- a/python/sglang/srt/layers/attention/dsv4/metadata.py +++ b/python/sglang/srt/layers/attention/dsv4/metadata.py @@ -45,8 +45,8 @@ positions: Some other notes: c4_ / c128_: means "compressed by 4" / "compressed by 128". - c4_page_size: page_size // 4 - c4_seq_lens: seq_lens // 4, but bounded by at least 1, due to flash_mla requirement. + compressed_page_size: page_size // 4 + compressed_seq_lens: seq_lens // 4, but bounded by at least 1, due to flash_mla requirement. c4_sparse: means "compressed by 4" but only attend to top-512 tokens. all related length will be clipped to 512. """ @@ -115,7 +115,7 @@ class NonPagedIndexerPlan: class PagedIndexerMetadata: page_size: int page_table: torch.Tensor - c4_seq_lens: torch.Tensor + compressed_seq_lens: torch.Tensor use_topk_v2: bool force_deep_gemm_metadata: bool = False use_prefill_cuda_graph: bool = False @@ -135,7 +135,7 @@ class PagedIndexerMetadata: use_jit_indexer = not self.force_deep_gemm_metadata and ( envs.SGLANG_OPT_USE_JIT_INDEXER_METADATA.get() - or self.c4_seq_lens.numel() > _LARGE_INDEXER_QUERY_THRESHOLD + or self.compressed_seq_lens.numel() > _LARGE_INDEXER_QUERY_THRESHOLD ) if use_jit_indexer: from sglang.kernels.ops.attention.dsv4 import ( @@ -144,25 +144,25 @@ class PagedIndexerMetadata: else: from deep_gemm import get_paged_mqa_logits_metadata - _c4 = self.c4_seq_lens.to(torch.int32) - if _c4.dim() == 1: - _c4 = _c4.unsqueeze(-1) - if _IS_SM120 and _c4.shape[0] > _SM120_INDEXER_M_CHUNK: - # Chunk metadata is identical for every layer in the forward - # pass; compute the per-chunk list once here instead of per - # layer in the indexer. + compressed_seq_lens = self.compressed_seq_lens.to(torch.int32) + if compressed_seq_lens.dim() == 1: + compressed_seq_lens = compressed_seq_lens.unsqueeze(-1) + if _IS_SM120 and compressed_seq_lens.shape[0] > _SM120_INDEXER_M_CHUNK: + # Chunk metadata is shared by all indexer layers in this forward. self.deep_gemm_metadata = [ get_paged_mqa_logits_metadata( - _c4[_s : _s + _SM120_INDEXER_M_CHUNK], - self.c4_page_size, + compressed_seq_lens[_s : _s + _SM120_INDEXER_M_CHUNK], + self.compressed_page_size, deep_gemm.get_num_sms(), ) - for _s in range(0, _c4.shape[0], _SM120_INDEXER_M_CHUNK) + for _s in range( + 0, compressed_seq_lens.shape[0], _SM120_INDEXER_M_CHUNK + ) ] else: self.deep_gemm_metadata = get_paged_mqa_logits_metadata( - _c4, - self.c4_page_size, + compressed_seq_lens, + self.compressed_page_size, deep_gemm.get_num_sms(), ) @@ -171,14 +171,14 @@ class PagedIndexerMetadata: if self.use_topk_v2: from sglang.kernels.ops.attention.dsv4 import plan_topk_v2 - self.topk_metadata = plan_topk_v2(self.c4_seq_lens) + self.topk_metadata = plan_topk_v2(self.compressed_seq_lens) else: self.topk_metadata = torch.empty((0,)) assert self.page_size == 256, "the system hardcodes page_size=256" @property - def c4_page_size(self) -> int: + def compressed_page_size(self) -> int: return self.page_size // 4 @property @@ -186,15 +186,15 @@ class PagedIndexerMetadata: return self.page_table.shape[1] * self.page_size @property - def max_c4_seq_len(self) -> int: - return self.page_table.shape[1] * self.c4_page_size + def max_compressed_seq_len(self) -> int: + return self.page_table.shape[1] * self.compressed_page_size def copy_(self, other: PagedIndexerMetadata): if is_hip(): - copy_fields = ["page_table", "c4_seq_lens"] + copy_fields = ["page_table", "compressed_seq_lens"] assign_fields = ["deep_gemm_metadata", "nonpaged_plan"] else: - copy_fields = ["page_table", "c4_seq_lens", "deep_gemm_metadata"] + copy_fields = ["page_table", "compressed_seq_lens", "deep_gemm_metadata"] assign_fields = ["nonpaged_plan"] copy_fields += ["topk_metadata"] copy_metadata( diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index ddc7c7eca..a6f9b5672 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -42,9 +42,8 @@ def get_compress_state_ring_size( compress_ratio: int, is_speculative: bool = False ) -> int: assert compress_ratio in [4, 128], f"Unsupported {compress_ratio = }" - # Online c128 keeps a single (max, sum, kv) state per index instead of a - # 128-slot ring buffer of raw tokens, so ring_size collapses to 1. Online - # is incompatible with speculative decode for now. + # Online C128 stores one (max, sum, kv) state per index; + # speculative decoding requires the experimental online C128 MTP path. if compress_ratio == 128 and ONLINE_C128: if is_speculative and not envs.SGLANG_EXPERIMENTAL_ONLINE_C128_MTP.get(): raise AssertionError("online c128 does not support MTP") @@ -56,9 +55,8 @@ def get_compress_state_ring_size( def get_compress_state_write_pad(compress_ratio: int, ring_size: int) -> int: - """Largest draft-token count this ring can serve; mirrors `mtp_pad` in `c_plan.cuh` - (the bound is derived there). Zero for a non-speculative ring, which is exactly one - window wide.""" + # Draft-token capacity must match mtp_pad in c_plan.cuh; + # a non-speculative ring has no write padding. window_size = compress_ratio * (2 if compress_ratio == 4 else 1) return ring_size - window_size + 2 if ring_size > window_size else 0 @@ -620,9 +618,8 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): ) self.max_num_reqs = max_num_reqs - # SWA ring needs one slot per addressable req_pool_idx. PD decode inflates - # req_to_token past max_num_reqs (pre-alloc), so the caller passes the real - # capacity; sizing as max_num_reqs+1 overflows ("length out of range"). + # PD preallocation can exceed max_num_reqs; + # the SWA ring must cover every addressable req_pool_idx. self.num_req_slots = ( num_req_slots if num_req_slots is not None else max_num_reqs + 1 ) @@ -648,9 +645,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): self.c4_state_pool_size = c4_state_pool_size c128_ring_size = self.get_ring_size(128) if ONLINE_C128: - # Request-scoped online C128 state is indexed by req_pool_idx. - # PD decode can allocate pre-transfer slots beyond - # max_num_reqs, so size to the actual req_to_token row count. + # Request-scoped C128 state must also cover PD preallocation slots. c128_state_pool_size = max(c128_state_pool_size, self.num_req_slots) else: # Offline C128 keeps a per-request raw state ring. @@ -860,8 +855,8 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): return data_ptrs, data_lens, item_lens def get_unified_swa_ring_buf_infos(self) -> Tuple[List[int], List[int], List[int]]: - """SWA-ring region [0, swa_pages) of every unified_kv layer, addressed - per-row by ring slot. Shipped as the StateType.SWA_RING PD component.""" + # StateType.SWA_RING transfers [0, swa_pages) of each unified_kv layer; + # its indices address individual ring rows. # TODO(billishyahao): validate PP layer-slicing for SWA_RING. data_ptrs: List[int] = [] data_lens: List[int] = [] @@ -878,12 +873,8 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): return data_ptrs, data_lens, item_lens def unified_region_buffers(self, ratio: int) -> Tuple[List[torch.Tensor], int]: - """ - In unified_kv, swa/c4/c128 share one buffer with one slot per row. But the - HiCache host pool transfers a whole page per indexed row, so we reshape the - compressed region into the layout it expects: skip the SWA segment, reshape to - one row per page, then cast to uint8. - """ + # HiCache expects byte rows containing whole pages; + # the unified pool stores individual token rows after its SWA region. assert self._unified_kv, "unified_region_buffers requires unified_kv layout" assert ratio in (4, 128), f"unsupported compression ratio: {ratio}" @@ -943,7 +934,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): return data_ptrs, data_lens, item_lens - def get_c128_state_buf_infos( + def get_request_state_buf_infos( self, ) -> Tuple[List[int], List[int], List[int]]: data_ptrs: List[int] = [] @@ -1076,7 +1067,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): ) def _init_compressed_layer_mapping(self): - c1_cnt = c4_cnt = c128_cnt = 0 + c0_cnt = c4_cnt = c128_cnt = 0 total_L = len(self.compression_ratios) self.layer_mapping: List[Optional[DeepSeekV4LayerItem]] = [None] * total_L @@ -1085,9 +1076,9 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): if ratio == 0: self.layer_mapping[idx] = DeepSeekV4LayerItem( compress_ratio=0, - compress_layer_id=c1_cnt, + compress_layer_id=c0_cnt, ) - c1_cnt += 1 + c0_cnt += 1 elif ratio == 4: self.layer_mapping[idx] = DeepSeekV4LayerItem( compress_ratio=4, @@ -1188,9 +1179,8 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): accept_lens: torch.Tensor, num_draft_tokens: int, ) -> None: - """Clear offline C128 ring slots written for rejected speculative tokens. - C4 needs no counterpart: its draft states are overwritten in position order - before any read; a C128 compression boundary can read a stale draft slot.""" + # C128 compression can read rejected draft slots at a boundary; + # C4 overwrites its draft slots before reading them. if ONLINE_C128 or num_draft_tokens <= 1 or req_pool_indices.numel() == 0: return diff --git a/test/registered/attention/unittests/dsv4/test_deepseek_v4.py b/test/registered/attention/unittests/dsv4/test_deepseek_v4.py index 20d1a74fc..acbef7463 100644 --- a/test/registered/attention/unittests/dsv4/test_deepseek_v4.py +++ b/test/registered/attention/unittests/dsv4/test_deepseek_v4.py @@ -365,7 +365,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase): metadata.c128_topk_lengths_clamp1 = torch.tensor( [base + 39, base + 40], dtype=torch.int32 ) - metadata.c1_flashmla_metadata = object() + metadata.c0_flashmla_metadata = object() metadata.c4_flashmla_metadata = object() metadata.c128_flashmla_metadata = object() return metadata @@ -377,10 +377,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase): ) from sglang.srt.server_args import ServerArgs - # cg-refactor folded the legacy enable_breakable_cuda_graph flag - # into cuda_graph_config. Verify the per-phase backend selectors - # default to None (i.e. nothing opted into BREAKABLE without an - # explicit CLI flag). + # Breakable graphs require explicit opt-in for each phase. sa = ServerArgs(model_path="dummy") self.assertNotEqual(sa.cuda_graph_backend_decode, "breakable") self.assertNotEqual(sa.cuda_graph_backend_prefill, "breakable") @@ -519,7 +516,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase): "swa_topk_lengths", "c128_page_indices", "c128_topk_lengths_clamp1", - "c1_flashmla_metadata", + "c0_flashmla_metadata", "c4_flashmla_metadata", "c128_flashmla_metadata", ] @@ -688,13 +685,8 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase): class TestDSV4SwaOutCacheLocResolution(CustomTestCase): - """`get_swa_out_cache_loc`: cached fast path vs store-time fallback. - - The KV-store consumers run in paths that never invoke - `init_forward_metadata_in_graph` (eager idle, runners that only run the - out-graph prep) or whose batch is re-padded after init (DP attention). - The resolver must use the per-forward cached value only when it is - provably current and fall back to translating `out_cache_loc` otherwise. + """SWA writes must translate live locations for idle or missing/mismatched caches. + A matching cache on an active forward must be reused. """ def _make_backend(self, mapping: torch.Tensor): diff --git a/test/registered/unit/disaggregation/test_disaggregation_wire.py b/test/registered/unit/disaggregation/test_disaggregation_wire.py index 3ba0f3cb3..c2e354f2e 100644 --- a/test/registered/unit/disaggregation/test_disaggregation_wire.py +++ b/test/registered/unit/disaggregation/test_disaggregation_wire.py @@ -654,7 +654,7 @@ def _make_dsv4_target(*, unified, mapping=None): pool.get_unified_swa_ring_buf_infos = lambda: ( _buf_infos(12) if unified else ([], [], []) ) - pool.get_c128_state_buf_infos = lambda: ([], [], []) + pool.get_request_state_buf_infos = lambda: ([], [], []) return pool diff --git a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py index 2f08b5c26..06fe11e66 100644 --- a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py +++ b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py @@ -46,7 +46,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase): metadata = PagedIndexerMetadata( page_size=256, page_table=torch.zeros((1, 1), dtype=torch.int32), - c4_seq_lens=torch.tensor([65], dtype=torch.int32), + compressed_seq_lens=torch.tensor([65], dtype=torch.int32), use_topk_v2=False, force_deep_gemm_metadata=True, ) @@ -68,7 +68,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase): metadata = PagedIndexerMetadata( page_size=256, page_table=torch.zeros((1, 1), dtype=torch.int32), - c4_seq_lens=torch.tensor([65], dtype=torch.int32), + compressed_seq_lens=torch.tensor([65], dtype=torch.int32), use_topk_v2=False, ) @@ -84,7 +84,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase): metadata = PagedIndexerMetadata( page_size=256, page_table=torch.zeros((1, 1), dtype=torch.int32), - c4_seq_lens=torch.tensor([65], dtype=torch.int32), + compressed_seq_lens=torch.tensor([65], dtype=torch.int32), use_topk_v2=False, ) @@ -253,7 +253,7 @@ class TestDSV4NonPagedIndexer(CustomTestCase): extend_start_loc=torch.tensor([0], dtype=torch.int32), extend_num_tokens=query_rows, ) - metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64) + metadata = SimpleNamespace(nonpaged_plan=None, compressed_page_size=64) page_table = torch.tensor([[3, 1]], dtype=torch.int32).repeat(query_rows, 1) c4_seq_lens = torch.tensor([62, 63, 64, 65], dtype=torch.int32) @@ -301,7 +301,7 @@ class TestDSV4NonPagedIndexer(CustomTestCase): extend_start_loc=torch.tensor([0], dtype=torch.int32), extend_num_tokens=query_rows, ) - metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64) + metadata = SimpleNamespace(nonpaged_plan=None, compressed_page_size=64) page_table = torch.zeros((query_rows, 1), dtype=torch.int32) c4_seq_lens = torch.tensor( [124_997, 124_998, 124_999, 125_000], dtype=torch.int32 @@ -339,7 +339,7 @@ class TestDSV4NonPagedIndexer(CustomTestCase): backend = SimpleNamespace(_can_use_nonpaged_indexer=can_use_nonpaged_indexer) backend.dsa_topk_backend = SimpleNamespace(is_sgl_kernel=lambda: True) c4_indexer = SimpleNamespace(use_fp4_indexer=False, index_topk=512) - metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64) + metadata = SimpleNamespace(nonpaged_plan=None, compressed_page_size=64) def build_plan(query_rows): batch = SimpleNamespace(