Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
18 commits
Select commit Hold shift + click to select a range
f5d16ca
dsv4 metadata: index_topk; present_ratios; build only present ratios
hnyls2002 Sep 17, 2026
581445b
dsv4 metadata: ratio-keyed sparse accessors
hnyls2002 Sep 17, 2026
9527e41
sparse prefill: explicit query_pos; per-ratio gather; kernel test
hnyls2002 Sep 17, 2026
3adb9ef
dsv4 pool: request_scoped state; PD request-state helper
hnyls2002 Sep 17, 2026
d3f6d7f
dsv4 pool: kv layout getters
hnyls2002 Sep 17, 2026
48feef6
merge main
hnyls2002 Sep 17, 2026
cfde99d
request_scoped from pool factory; present_ratios on pool; request-sta…
hnyls2002 Sep 17, 2026
996a1ea
bytes getters read pool dim; drop CompressedGather.page_size; DEFAULT…
hnyls2002 Sep 17, 2026
c3210f1
drop unused is_dsv4_c128_online_enabled
hnyls2002 Sep 17, 2026
0912107
state transfer indices on CompressStatePool; c128 in sparse accessors…
hnyls2002 Sep 17, 2026
a3edf13
present_ratios required on DSV4AttnMetadata
hnyls2002 Sep 17, 2026
683daf2
has_c4/has_c128 properties; single cp reindex loop
hnyls2002 Sep 17, 2026
2b1c0c0
c128 gather into CompressedGather; layer_inputs; shared workspace bases
hnyls2002 Sep 17, 2026
ea422d9
comments
hnyls2002 Sep 17, 2026
6764f1f
tests: present_ratios gating; accessor routing; cp reindex; pool requ…
hnyls2002 Sep 17, 2026
fde85f7
Merge branch 'main' into lsyin/dsv4-ratio-generalization
hnyls2002 Sep 17, 2026
90607ab
request-scoped state pools are per layer; share one ring layout
hnyls2002 Sep 17, 2026
2322571
Merge branch 'main' into lsyin/dsv4-ratio-generalization
hnyls2002 Sep 17, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@ def _combine_topk_swa_indices_kernel(
topk_indices_ptr,
topk_indices_stride,
query_start_loc_ptr,
query_pos_ptr,
seq_lens_ptr,
gather_lens_ptr,
compressed_base_ptr,
Expand All @@ -62,23 +63,19 @@ def _combine_topk_swa_indices_kernel(
base = tl.load(query_start_loc_ptr)
query_start = tl.load(query_start_loc_ptr + batch_idx) - base
query_end = tl.load(query_start_loc_ptr + batch_idx + 1) - base
query_len = query_end - query_start
seq_len = tl.load(seq_lens_ptr + batch_idx)
gather_len = tl.load(gather_lens_ptr + batch_idx)
compressed_base = tl.load(compressed_base_ptr + batch_idx)
swa_base = tl.load(swa_base_ptr + batch_idx)
start_pos = seq_len - query_len
# SWA portion of the gathered buffer starts from position
# (seq_len - gather_len), not 0. The +pos-gather_start formula maps a
# query's window back into the workspace's SWA region.
gather_start = seq_len - gather_len

for token_idx in range(query_start + worker_id, query_end, num_workers):
token_idx_in_query = token_idx - query_start
pos = start_pos + token_idx_in_query
# Both the C4 indexer and the C128 metadata builder emit
# min((pos+1)//compress_ratio, topk_tokens) valid entries. Caller
# passes top_k=0 for SWA-only layers to zero this out.
pos = tl.load(query_pos_ptr + token_idx)
# -1 entries inside the top-k span stay -1 (attention skips them).
# top_k=0 disables the compressed portion for SWA-only layers.
topk_len = tl.minimum((pos + 1) // COMPRESS_RATIO, top_k)
swa_len = tl.minimum(pos + 1, WINDOW_SIZE)

Expand All @@ -93,7 +90,7 @@ def _combine_topk_swa_indices_kernel(
)
tl.store(
combined_indices_ptr + combined_row + offset,
topk_vals + compressed_base,
tl.where(topk_vals >= 0, topk_vals + compressed_base, -1),
mask=mask,
)

Expand Down
23 changes: 8 additions & 15 deletions python/sglang/srt/disaggregation/decode.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,10 +61,8 @@
build_kv_layer_ids,
build_staging_slot_metadata,
get_dsa_tail_state_indices,
get_dsv4_c128_state_indices,
get_kv_class,
get_qsa_pending_state_indices,
is_dsv4_c128_online_enabled,
is_mla_backend,
is_unadmitted_reject,
poll_and_all_reduce,
Expand Down Expand Up @@ -1498,23 +1496,18 @@ def _swa_ring_payload():
ring_rows = state_slot * ring_stride + (positions % ring_stride)
return ring_rows.astype(np.int32)

def _c128_state_payload():
online = is_dsv4_c128_online_enabled()
ring_size = 1 if online else self.token_to_kv_pool.get_ring_size(128)
return get_dsv4_c128_state_indices(
int(decode_req.req.kv.req_pool_idx),
seq_len,
online=online,
ring_size=ring_size,
def _request_state_payload():
return self.token_to_kv_pool.request_state_transfer_indices(
int(decode_req.req.kv.req_pool_idx), seq_len
)

state_types = self.kv_manager.kv_args.state_types
if StateType.DSV4_REQUEST_STATE in state_types:
clear_c128_state = getattr(
self.token_to_kv_pool, "clear_c128_req_state", None
clear_request_state = getattr(
self.token_to_kv_pool, "clear_request_scoped_state", None
)
if clear_c128_state is not None:
clear_c128_state(int(decode_req.req.kv.req_pool_idx))
if clear_request_state is not None:
clear_request_state(int(decode_req.req.kv.req_pool_idx))
payloads = {
StateType.MAMBA: _mamba_payload,
StateType.QSA_PENDING: _qsa_pending_payload,
Expand All @@ -1524,7 +1517,7 @@ def _c128_state_payload():
StateType.DSA_TAIL: _dsa_tail_payload,
StateType.MINIMAX_INDEX_K: _full_kv_pages_payload,
StateType.SWA_RING: _swa_ring_payload,
StateType.DSV4_REQUEST_STATE: _c128_state_payload,
StateType.DSV4_REQUEST_STATE: _request_state_payload,
StateType.BLOCK_SCALE: _full_kv_pages_payload,
StateType.BLOCK_SCALE_SWA: _swa_payload,
}
Expand Down
30 changes: 5 additions & 25 deletions python/sglang/srt/disaggregation/prefill.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,11 +52,9 @@
build_kv_layer_ids,
build_staging_slot_metadata,
get_dsa_tail_state_indices,
get_dsv4_c128_state_indices,
get_kv_class,
get_qsa_pending_state_indices,
is_aborted,
is_dsv4_c128_online_enabled,
is_mla_backend,
is_unadmitted_reject,
poll_and_all_reduce_attn_cp_tp_group,
Expand Down Expand Up @@ -257,7 +255,6 @@ def _init_kv_manager(self) -> CommonKVManager:
hf_text_config=self.scheduler.model_config.hf_text_config,
)
)
kv_args.mla_compression_ratios = None
kv_data_ptrs, kv_data_lens, kv_item_lens = (
self.token_to_kv_pool.get_contiguous_buf_infos()
)
Expand Down Expand Up @@ -316,13 +313,6 @@ def _init_kv_manager(self) -> CommonKVManager:
req_to_token_pool=req_to_token_pool,
)

if isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool):
# V4's KVCache is organized by compression-ratio
# buckets rather than by layer.
kv_args.mla_compression_ratios = list(
self.token_to_kv_pool.compression_ratios
)

kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
kv_manager = kv_manager_class(
kv_args,
Expand Down Expand Up @@ -1415,20 +1405,10 @@ def _swa_ring_payload():
ring_rows = state_slot * ring_stride + (positions % ring_stride)
return ring_rows.astype(np.int32)

def _c128_state_payload():
online = is_dsv4_c128_online_enabled()
ring_size = (
1
if online
else self.token_to_kv_pool_allocator.get_kvcache().get_ring_size(
128
)
)
return get_dsv4_c128_state_indices(
int(req.kv.req_pool_idx),
c128_seq_len,
online=online,
ring_size=ring_size,
def _request_state_payload():
kvcache = self.token_to_kv_pool_allocator.get_kvcache()
return kvcache.request_state_transfer_indices(
int(req.kv.req_pool_idx), c128_seq_len
)

state_types = (
Expand All @@ -1443,7 +1423,7 @@ def _c128_state_payload():
StateType.DSA_TAIL: _dsa_tail_payload,
StateType.MINIMAX_INDEX_K: _full_kv_pages_payload,
StateType.SWA_RING: _swa_ring_payload,
StateType.DSV4_REQUEST_STATE: _c128_state_payload,
StateType.DSV4_REQUEST_STATE: _request_state_payload,
StateType.BLOCK_SCALE: _full_kv_pages_payload,
StateType.BLOCK_SCALE_SWA: _swa_payload,
}
Expand Down
53 changes: 7 additions & 46 deletions python/sglang/srt/disaggregation/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
from sglang.srt.runtime_context import (
get_disagg,
)
from sglang.srt.utils import is_hip, is_npu
from sglang.srt.utils import is_npu

if TYPE_CHECKING:
from sglang.srt.disaggregation.base.conn import KVArgs, StateType
Expand All @@ -46,7 +46,6 @@
# Constants & Enums
#########################
FAKE_BOOTSTRAP_HOST = "2.2.2.2"
_IS_HIP = is_hip()


def poll_and_all_reduce_pp(
Expand Down Expand Up @@ -78,50 +77,6 @@ def get_dsa_seed_metadata_dim(hf_config) -> int:
return get_dsa_mtp_topk_width(hf_config)


def is_dsv4_c128_online_enabled() -> bool:
"""Return whether DSV4 C128 uses request-scoped online state."""
return not _IS_HIP and envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get()


def get_dsv4_c4_state_indices(
req_pool_idx: int,
seq_len: int,
*,
ring_size: int,
) -> np.ndarray:
# 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}"
)

seq_len = max(0, int(seq_len))
state_len = seq_len % 4 + 4
positions = np.arange(max(0, seq_len - state_len), seq_len, dtype=np.int64)
rows = int(req_pool_idx) * int(ring_size) + positions % int(ring_size)
return rows.astype(np.int32)


def get_dsv4_c128_state_indices(
req_pool_idx: int,
seq_len: int,
*,
online: bool,
ring_size: int,
) -> np.ndarray:
"""Return the PD transfer row/page indices for DSV4 C128 state."""
if seq_len == 0 or seq_len % 128 == 0:
return np.empty((0,), dtype=np.int32)
if online:
return np.array([int(req_pool_idx)], dtype=np.int32)

assert ring_size % 128 == 0, f"C128 ring_size must be 128-aligned, got {ring_size}"
pages_per_req = ring_size // 128
page = int(req_pool_idx) * pages_per_req + ((seq_len - 1) % ring_size) // 128
return np.array([page], dtype=np.int32)


def get_qsa_pending_state_indices(req: Req) -> np.ndarray:
"""Return the request-pool row that owns a QSA pending-state ring."""
req_pool_idx = req.kv.req_pool_idx
Expand Down Expand Up @@ -1380,6 +1335,12 @@ def setup_state_kv_args(
kv_args.state_layer_ids = []
kv_args.is_hybrid_mla_backend = False
kv_args.state_conv_shard_groups = []
# V4's KVCache is organized by compression-ratio buckets rather than by layer.
kv_args.mla_compression_ratios = (
list(token_to_kv_pool.compression_ratios)
if isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
else None
)

def append_dsa_tail(pool) -> None:
if not pool.kpool_use_compress:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -544,7 +544,7 @@ def free(
row = req_to_token_pool.req_to_c128_sidecar[int(req_pool_idx)]
self.release_c128_pages(row[row > 0])
row.zero_()
self.get_kvcache().clear_c128_req_state(int(req_pool_idx))
self.get_kvcache().clear_request_scoped_state(int(req_pool_idx))

def available_size(self):
return min(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -88,8 +88,10 @@ def dsv4_state_payloads(
import numpy as np

from sglang.srt.disaggregation.ascend.conn import AscendStateType
from sglang.srt.disaggregation.utils import get_dsv4_c4_state_indices
from sglang.srt.hardware_backend.npu.utils import is_npu_arch35
from sglang.srt.mem_cache.deepseek_v4_compress_state import (
c4_state_transfer_indices,
)

seq_len = max(0, int(seq_len))
prefix_len = max(0, min(int(prefix_len), seq_len))
Expand All @@ -113,7 +115,7 @@ def c128_kv_pages():
if is_npu_arch35():

def c4_state_indices():
return get_dsv4_c4_state_indices(
return c4_state_transfer_indices(
req_pool_idx,
seq_len,
ring_size=req_to_token_pool.get_dsv4_c4_state_ring_size(),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -118,6 +118,7 @@ def __init__(
enable_memory_saver: bool,
ratio: int,
ring_size: int,
request_scoped: bool,
swa_page_size: int,
):
assert ratio in (
Expand All @@ -139,6 +140,7 @@ def __init__(
enable_memory_saver=enable_memory_saver,
ratio=ratio,
online=False,
request_scoped=request_scoped,
swa_page_size=swa_page_size,
state_cache_page_size=ring_size,
)
Expand Down Expand Up @@ -352,6 +354,7 @@ def _make_compress_state_pool(
device=self.device,
enable_memory_saver=enable_memory_saver,
ratio=ratio,
request_scoped=ratio == 128,
swa_page_size=self.swa_page_size,
)

Expand Down
Loading
Loading