[NPU] Support DeepSeek-V4 DSpark and refactor DSV4 cache management (#33676)

Co-authored-by: JiaruiChang5268 <jc5268@columbia.edu>
Co-authored-by: Kelon <kelonlu@163.com>
Co-authored-by: unknown <z8ruev42yk@gmail.com>
Co-authored-by: Talantan1102 <545811257@qq.com>
Co-authored-by: Talantan1102 <44429302+Talantan1102@users.noreply.github.com>
This commit is contained in:
Zed
2026-08-17 16:27:44 +08:00
committed by GitHub
co-authored by JiaruiChang5268 Kelon unknown Talantan1102 Talantan1102
parent e03c53fc13
commit b83d507cd7
37 changed files with 2010 additions and 2018 deletions
@@ -9,7 +9,11 @@ from sglang.kernels.ops.speculative.cache_locs import assign_extend_cache_locs_f
from sglang.kernels.ops.speculative.dspark.dispatch import inputs_on_cuda
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
from sglang.srt.utils import is_npu
from sglang.srt.utils import (
is_npu,
)
_is_npu = is_npu()
class RaggedVerifyWindow(msgspec.Struct, frozen=True):
@@ -798,7 +802,7 @@ def build_commit_inject_layout_triton(
class BuildOutTokens:
@classmethod
def execute(cls, *args, **kwargs) -> torch.Tensor:
if not is_npu() and inputs_on_cuda(*args, **kwargs):
if inputs_on_cuda(*args, **kwargs) and not _is_npu:
return cls.triton(*args, **kwargs)
return cls.torch(*args, **kwargs)
@@ -277,9 +277,9 @@ def _target_checkpoint_bundles_dspark_draft(server_args: ServerArgs) -> bool:
def _handle_dspark(server_args: ServerArgs) -> None:
_is_npu = server_args.device.startswith("npu")
if not server_args.device.startswith("cuda") and not _is_npu:
if not server_args.device.startswith(("cuda", "npu")):
raise ValueError(
"DSpark speculative decoding only supports CUDA and NPU devices."
"DSpark speculative decoding only supports CUDA or NPU device."
)
# dp_size==1 with dp_attention is a degenerate flag under DSV4 CP; skip DP-only checks.
@@ -21,15 +21,9 @@ logger = logging.getLogger(__name__)
class AscendStateType(str, enum.Enum):
"""DSV4-on-NPU per-pool PD components, kept out of the cross-hardware
StateType enum. Sent via the same page-indexed path as SWA."""
"""DSV4-on-NPU PD components without a cross-hardware equivalent."""
DSV4_SWA = "dsv4_swa"
DSV4_C4 = "dsv4_c4"
DSV4_C128 = "dsv4_c128"
DSV4_INDEXER = "dsv4_indexer"
DSV4_C4_STATE = "dsv4_c4_state"
DSV4_C128_STATE = "dsv4_c128_state"
_DSV4_KVCACHE_STATE_TYPES = tuple(AscendStateType)
@@ -71,6 +65,32 @@ class AscendKVManager(MooncakeKVManager):
def get_mla_kv_ptrs_with_pp(
self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int], state_type=None
) -> Tuple[List[int], List[int], int]:
mla_ratios = getattr(self.kv_args, "mla_compression_ratios", None)
if mla_ratios:
if len(src_kv_ptrs) == len(dst_kv_ptrs):
return src_kv_ptrs, dst_kv_ptrs, len(src_kv_ptrs)
start_layer = self.kv_args.prefill_start_layer
end_layer = self.kv_args.prefill_end_layer
c4_full = sum(ratio == 4 for ratio in mla_ratios)
c4_start = sum(ratio == 4 for ratio in mla_ratios[:start_layer])
c4_end = sum(ratio == 4 for ratio in mla_ratios[:end_layer])
c128_start = sum(ratio == 128 for ratio in mla_ratios[:start_layer])
c128_end = sum(ratio == 128 for ratio in mla_ratios[:end_layer])
if state_type == AscendStateType.DSV4_C128:
dst = dst_kv_ptrs[c128_start:c128_end]
return src_kv_ptrs, dst, len(src_kv_ptrs)
# NPU main KV layout: [C4 KV, index K, index scale].
if state_type is None and len(dst_kv_ptrs) == 3 * c4_full:
dst = []
for offset in (0, c4_full, 2 * c4_full):
dst.extend(dst_kv_ptrs[offset + c4_start : offset + c4_end])
return src_kv_ptrs, dst, len(src_kv_ptrs)
return super().get_mla_kv_ptrs_with_pp(src_kv_ptrs, dst_kv_ptrs, state_type)
# src_kv_ptrs: k_data, v_data, index_k_data(optional)
# dst_kv_ptrs: k_data, v_data, index_k_data(optional)
# state_type is accepted for parity with the common disaggregation path;
+2 -23
View File
@@ -1102,17 +1102,6 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
self.tree_cache.dec_lock_ref(decode_req.req.last_node)
break
if total_prefix_len != 0 and hasattr(
self.token_to_kv_pool_allocator, "c4_attn_allocator"
):
if prefix_len > 0:
self.tree_cache.dec_lock_ref(decode_req.req.last_node)
raise RuntimeError(
"DSV4 NPU PD disaggregation does not support decode-side "
"prefix cache yet; disable disaggregation decode radix/HiCache "
"for PD + chunked prefill."
)
dst_kv_indices = self._pre_alloc(
decode_req.req,
prefix_indices,
@@ -1176,7 +1165,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
def _swa_payload():
window_size = self.scheduler.sliding_window_size
window_start = max(0, seq_len - window_size)
window_start = max(total_prefix_len, seq_len - window_size)
window_start = page_align_floor(window_start, page_size)
window_kv_indices_full = self.req_to_token_pool.req_to_token[
decode_req.req.req_pool_idx, window_start:seq_len
@@ -1234,15 +1223,6 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
StateType.SWA_RING: _swa_ring_payload,
StateType.C128_STATE: _c128_state_payload,
}
if hasattr(self.req_to_token_pool, "req_to_token_c4"):
# DSV4 on NPU: per-pool dst page indices, produced by the same
# shared builder prefill uses so src/dst line up positionally.
if total_prefix_len != 0:
raise RuntimeError(
"DSV4 NPU PD disaggregation does not support decode-side "
"prefix cache yet; disable disaggregation decode radix/HiCache "
"for PD + chunked prefill."
)
if _is_npu and isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool):
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
dsv4_state_payloads,
@@ -1254,7 +1234,6 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
decode_req.req.req_pool_idx,
seq_len,
self.token_to_kv_pool_allocator.page_size,
self.scheduler.sliding_window_size,
prefix_len=total_prefix_len,
)
)
@@ -1743,7 +1722,7 @@ def alloc_for_decode_prealloc(
)
extra_kwargs = {}
dsv4_unwrap_prealloc = None
if hasattr(allocator, "c4_attn_allocator"):
if hasattr(allocator, "c128_attn_allocator"):
assert req_to_token_pool is not None
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
dsv4_prealloc_kwargs,
+2 -3
View File
@@ -1200,7 +1200,7 @@ class SchedulerDisaggregationPrefillMixin:
def _swa_payload():
window_size = self.sliding_window_size
window_start = max(0, seq_len - window_size)
window_start = max(req.disagg_decode_prefix_len, seq_len - window_size)
window_start = (window_start // page_size) * page_size
window_kv_indices_full = self.req_to_token_pool.req_to_token[
req.req_pool_idx, window_start:seq_len
@@ -1274,8 +1274,7 @@ class SchedulerDisaggregationPrefillMixin:
req.req_pool_idx,
seq_len,
page_size,
self.sliding_window_size,
prefix_len=0,
prefix_len=req.disagg_decode_prefix_len,
)
)
state_indices = [
+17 -17
View File
@@ -1064,17 +1064,7 @@ def setup_state_kv_args(
kv_args.is_hybrid_mla_backend = False
kv_args.state_conv_shard_groups = []
if is_npu() and isinstance(token_to_kv_pool, DSV4NPUTokenToKVPool):
# Pool ships each sub-pool as its own page-indexed component (fixed order
# so prefill and decode register identically); skips get_state_buf_infos.
for (
st,
comp_ptrs,
comp_lens,
comp_item_lens,
) in token_to_kv_pool.get_pd_state_components():
append_state_component(kv_args, st, comp_ptrs, comp_lens, comp_item_lens)
elif isinstance(token_to_kv_pool, MiniMaxSparseKVPool):
if isinstance(token_to_kv_pool, MiniMaxSparseKVPool):
if token_to_kv_pool.index_kv_pool is not None:
raise NotImplementedError(
"PD disaggregation for MiniMax sparse layers with index value "
@@ -1175,15 +1165,25 @@ def setup_state_kv_args(
kv_args, StateType.DSA, data_ptrs, data_lens, item_lens
)
if is_npu() and isinstance(token_to_kv_pool, DSV4NPUTokenToKVPool):
from sglang.srt.disaggregation.ascend.conn import AscendStateType
c128_ptrs, c128_lens, c128_item_lens = token_to_kv_pool.get_c128_kv_buf_infos()
if c128_ptrs:
append_state_component(
kv_args,
AscendStateType.DSV4_C128,
c128_ptrs,
c128_lens,
c128_item_lens,
)
# DSV4 NextN shares the target allocator, so target and draft use the same
# local SWA indices. Keep draft buffers in a separate positional component
# to avoid mixing them into the target's heterogeneous state layout, while
# reusing the existing SWA transport dispatch. NPU has a different paged
# state layout and is intentionally left unchanged.
if (
not is_npu()
and isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
and isinstance(draft_token_to_kv_pool, DeepSeekV4TokenToKVPool)
# reusing the existing SWA transport dispatch on both GPU and NPU.
if isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) and isinstance(
draft_token_to_kv_pool, DeepSeekV4TokenToKVPool
):
if not draft_token_to_kv_pool.compression_ratios or not all(
ratio == 0 for ratio in draft_token_to_kv_pool.compression_ratios
@@ -443,15 +443,14 @@ class AscendAttnBackend(AttentionBackend):
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init the metadata for a forward pass."""
self.forward_metadata = ForwardMetadata()
seq_lens_max = forward_batch.seq_lens.max()
if forward_batch.forward_mode.is_target_verify():
spec_tokens_per_req = int(forward_batch.spec_info.draft_token_num)
# Overlap scheduling can publish the CPU sequence length one step
# ahead of the device tensor. FIA consumes seq_lens_cpu below, so
# derive the block-table width from the same source. Otherwise a
# page-aligned request can expose KV_S=N while asking FIA for N+1.
seq_lens_max = (
forward_batch.seq_lens_cpu.max().item()
+ self.speculative_num_draft_tokens
)
seq_lens_max = forward_batch.seq_lens_cpu.max().item() + spec_tokens_per_req
elif (
forward_batch.forward_mode.is_decode_or_idle()
and forward_batch.spec_info is not None
@@ -499,10 +498,10 @@ class AscendAttnBackend(AttentionBackend):
seq_lens_list_cumsum = np.cumsum(forward_batch.extend_seq_lens_cpu)
self.forward_metadata.seq_lens_list_cumsum = seq_lens_list_cumsum
if forward_batch.forward_mode.is_target_verify() and not _is_dflash_verify(
forward_batch.spec_info
):
self.forward_metadata.seq_lens_cpu_int += self.speculative_num_draft_tokens
if forward_batch.forward_mode.is_target_verify():
spec_algorithm = forward_batch.spec_algorithm
if spec_algorithm is None or not spec_algorithm.is_dspark():
self.forward_metadata.seq_lens_cpu_int += spec_tokens_per_req
elif (
forward_batch.forward_mode.is_decode_or_idle()
and forward_batch.spec_info is not None
@@ -519,11 +518,16 @@ class AscendAttnBackend(AttentionBackend):
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2()
):
spec_tokens_per_req = (
int(forward_batch.spec_info.draft_token_num)
if forward_batch.forward_mode.is_target_verify()
else self.speculative_num_draft_tokens
)
self.forward_metadata.actual_seq_lengths_q = torch.arange(
self.speculative_num_draft_tokens,
self.speculative_num_draft_tokens
+ forward_batch.seq_lens.shape[0] * self.speculative_num_draft_tokens,
self.speculative_num_draft_tokens,
spec_tokens_per_req,
spec_tokens_per_req
+ forward_batch.seq_lens.shape[0] * spec_tokens_per_req,
spec_tokens_per_req,
dtype=torch.int32,
device=self.device,
)
@@ -720,12 +724,15 @@ class AscendAttnBackend(AttentionBackend):
max_seq_pages = (max_len + self.page_size - 1) // self.page_size
if self.is_hybrid_swa:
metadata.block_tables_swa[:bs, :max_seq_pages].copy_(
self.full_to_swa_index_mapping[
self.req_to_token[req_pool_indices[:bs], :max_len]
][:, :: self.page_size]
// self.page_size
full_page_locs = self.req_to_token[
req_pool_indices[:bs],
0 : max_len : self.page_size,
]
swa_page_table = (
self.full_to_swa_index_mapping[full_page_locs] // self.page_size
)
metadata.block_tables_swa[:bs, :max_seq_pages].copy_(swa_page_table)
metadata.block_tables_swa[:bs, max_seq_pages:].fill_(0)
metadata.block_tables_swa[bs:, :].fill_(0)
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,294 @@
"""C128 sidecar ownership for the DSV4 NPU Unified Radix Cache.
The component deliberately exposes only complete physical C128 pages to the
radix tree; partial tail pages remain request-owned.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Callable, Optional
import torch
from sglang.srt.mem_cache.base_prefix_cache import (
InsertParams,
InsertResult,
MatchPrefixParams,
MatchResult,
)
from sglang.srt.mem_cache.unified_cache.cache_action import (
FreeComponentDeviceSlot,
SWARebuild,
)
from sglang.srt.mem_cache.unified_cache.components import (
BASE_COMPONENT_TYPE,
ComponentType,
EvictLayer,
TreeComponent,
)
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.unified_cache.cache_action import (
CacheAction,
ComponentAction,
)
from sglang.srt.mem_cache.unified_radix_cache import (
UnifiedTreeNode,
)
class C128SidecarComponent(TreeComponent):
component_type = ComponentType.C128
@property
def allocator(self):
return self.cache.token_to_kv_pool_allocator
def _adjust_session_path(
self,
leaf: UnifiedTreeNode,
stop: UnifiedTreeNode,
delta: int,
) -> None:
"""Adjust session protection for every C128 boundary on a radix path."""
node = leaf
while node is not stop and node is not self.tree_core.root_node:
cd = node.component_data[self.component_type]
if delta < 0:
assert cd.session_ref > 0
prev_ref = cd.session_ref
cd.session_ref += delta
if (prev_ref == 0) != (cd.session_ref == 0):
self._refresh_session_partition(node)
node = node.parent
def _dec_session_coverage(self, session_id: str, leaf: UnifiedTreeNode) -> None:
self._adjust_session_path(leaf, self.tree_core.root_node, -1)
def _advance_session_coverage(
self,
session_id: str,
leaf: UnifiedTreeNode,
old_ancestor: Optional[UnifiedTreeNode],
) -> None:
stop = old_ancestor or self.tree_core.root_node
self._adjust_session_path(leaf, stop, 1)
def _recede_session_coverage(
self,
session_id: str,
leaf: UnifiedTreeNode,
fallback: Optional[UnifiedTreeNode],
) -> None:
stop = fallback or self.tree_core.root_node
self._adjust_session_path(leaf, stop, -1)
def _attach(self, node: UnifiedTreeNode, pages: torch.Tensor) -> None:
if pages.numel() == 0:
return
ct = self.component_type
cd = node.component_data[ct]
assert cd.value is None
value = pages.clone()
self.tree_core.set_component_device_value(node.id, ct, value)
self.allocator.retain_c128_pages(value)
def create_match_validator(
self, match_device_only: bool = False
) -> Callable[[UnifiedTreeNode], bool]:
# A page is attached only to the node ending its full physical group.
return lambda node: node.component_data[self.component_type].value is not None
def finalize_match_result_in_cache(
self, params: MatchPrefixParams, result: MatchResult
) -> MatchResult:
req = params.req
if req is None:
return result
chunks = []
node = self.tree_core.node_by_id(result.best_match_node)
root = self.tree_core.root_node
while node is not root:
value = node.component_data[self.component_type].value
if value is not None:
chunks.append(value)
node = node.parent
chunks.reverse()
pages = (
torch.cat(chunks)
if chunks
else self.allocator.c128_attn_allocator.free_pages.new_empty((0,))
)
group_tokens = 128 * self.allocator.c128_attn_allocator.page_size
assert pages.numel() == len(result.device_indices) // group_tokens
self.cache.req_to_token_pool.set_c128_prefix_pages(req, pages)
return result
def recover_after_unevict(
self,
node: UnifiedTreeNode,
prefix_len: int,
total_prefix_len: int,
params: InsertParams,
cache_actions: list[CacheAction | ComponentAction],
) -> None:
pages = params.c128_value
assert pages is not None
group_tokens = 128 * self.allocator.c128_attn_allocator.page_size
start = total_prefix_len // group_tokens
end = (total_prefix_len + prefix_len) // group_tokens
self._attach(node, pages[start:end])
@staticmethod
def _node_depth(node: UnifiedTreeNode) -> int:
depth = 0
while node.parent is not None:
depth += len(node.key)
node = node.parent
return depth
@staticmethod
def _split_pending_swa_rebuild(
new_parent: UnifiedTreeNode,
child: UnifiedTreeNode,
cache_actions: list[CacheAction | ComponentAction],
) -> None:
"""Keep a deferred SWA rebuild aligned when C128 splits its source node."""
for i, pending in enumerate(cache_actions):
if isinstance(pending, SWARebuild) and pending.node_id == child.id:
cache_actions[i : i + 1] = [
SWARebuild(
new_parent.id,
new_parent.component_data[BASE_COMPONENT_TYPE].value,
),
SWARebuild(
child.id,
child.component_data[BASE_COMPONENT_TYPE].value,
),
]
return
def _ensure_boundary_node(
self,
tail: UnifiedTreeNode,
boundary: int,
cache_actions: list[CacheAction | ComponentAction],
) -> UnifiedTreeNode:
node = tail
while node.parent is not None and self._node_depth(node.parent) >= boundary:
node = node.parent
node_end = self._node_depth(node)
if node_end == boundary:
return node
node_start = node_end - len(node.key)
assert node_start < boundary < node_end
new_parent, action = self.tree_core._split_node(
node.key, node, boundary - node_start
)
if action is not None:
cache_actions.append(action)
self._split_pending_swa_rebuild(new_parent, node, cache_actions)
return new_parent
def commit_insert_component_data(
self,
node: UnifiedTreeNode,
is_new_leaf: bool,
params: InsertParams,
result: InsertResult,
cache_actions: list[CacheAction | ComponentAction],
) -> None:
if not is_new_leaf:
return
assert params.key is not None
assert params.c128_value is not None
# Full/SWA may initially represent the new suffix as one long leaf.
# Materialize every complete C128 group boundary so a later branch can
# always match the nearest full C128-page prefix instead of falling back to
# the previous, potentially much shorter, Radix node.
group_tokens = 128 * self.allocator.c128_attn_allocator.page_size
first_boundary = (result.prefix_len // group_tokens + 1) * group_tokens
for boundary in range(first_boundary, len(params.key) + 1, group_tokens):
boundary_node = self._ensure_boundary_node(node, boundary, cache_actions)
page_index = boundary // group_tokens - 1
self._attach(boundary_node, params.c128_value[page_index : page_index + 1])
def redistribute_on_node_split(
self, new_parent: UnifiedTreeNode, child: UnifiedTreeNode
) -> None:
# Every stored value belongs to the old child's end boundary. Splitting
# inside that group leaves the page on the child.
ct = self.component_type
new_parent.component_data[ct].session_ref = child.component_data[ct].session_ref
assert new_parent.component_data[ct].session_ids is None
def evict_component(
self,
node: UnifiedTreeNode,
device_frees: dict[ComponentType, list[torch.Tensor]],
host_frees: dict[ComponentType, list[torch.Tensor]],
target: EvictLayer = EvictLayer.DEVICE,
) -> tuple[int, int]:
cd = node.component_data[self.component_type]
if EvictLayer.DEVICE in target and cd.value is not None:
device_frees[self.component_type].append(cd.value)
self.tree_core.component_evictable_size_[self.component_type] -= len(
cd.value
)
cd.value = None
# C128 pages are auxiliary to Full tokens and must not inflate the
# public token-eviction count.
return 0, 0
def prepare_for_caching_req(
self,
req: Req,
insert_params: InsertParams,
token_ids_len: int,
is_finished: bool,
) -> int:
logical_len = token_ids_len
if self.tree_core.is_eagle and logical_len > 0:
logical_len -= 1
group_tokens = 128 * self.allocator.c128_attn_allocator.page_size
cache_len = logical_len // group_tokens * group_tokens
num_pages = cache_len // group_tokens
insert_params.c128_value = self.cache.req_to_token_pool.req_to_c128_sidecar[
int(req.req_pool_idx), :num_pages
].clone()
return cache_len + 1 if self.tree_core.is_eagle and cache_len > 0 else cache_len
def apply_component_action(self, action: ComponentAction) -> None:
if isinstance(action, FreeComponentDeviceSlot):
for page_ids in action.indices:
self.allocator.release_c128_pages(page_ids)
return
raise AssertionError(
f"C128SidecarComponent: unhandled action {type(action).__name__}"
)
def eviction_priority(self, is_leaf: bool) -> int:
return 0 if is_leaf else 2
def _evict_device_start(self, request_cnt):
pass
def _evict_device_next_node(self, tracker, device_frees, host_frees):
return None
def _evict_device_end(self) -> None:
pass
def acquire_component_lock(self, node, result, lock_host=False):
return result
def release_component_lock(self, node, params, lock_host=False) -> None:
pass
def free_host_values(self, host_values) -> None:
pass
@@ -1,30 +1,26 @@
"""DSV4-NPU SWA + c4/c128 paged allocator.
"""DSV4-NPU SWA + c128 paged KV allocator.
Subclasses :class:`SWATokenToKVPoolAllocator` and adds paged allocation for the
c4/c128 compressed-KV pools and their tail-only compress-state pools, alongside
the parent's full + SWA pools.
C128 compressed-KV pool alongside the parent's full + SWA pools. C4 KV slots
are derived from full slots. Compressor state is fixed ring storage owned by
the KV pool and never enters this token allocator.
Per ``alloc_extend`` / ``alloc_decode``:
1. super() allocates the full + SWA slots (``out_full_loc``).
2. Allocate c4/c128 KV slots — one compressed token per ``ratio`` raw tokens
(``seq_len // ratio - prefix_len // ratio``) — via the standard
:class:`NPUPagedTokenToKVPoolAllocator` over the pool's c4/c128 KV buffers.
3. Allocate the c4/c128 compress-state slots the same way, tail-only per req,
using the per-req lens the scheduler packed into ``DSV4StateLens``.
4. Return a :class:`DSV4OutCacheLoc` bundling all five slot families.
2. Derive c4 KV slots from full and allocate c128 KV slots — one compressed
token per ``ratio`` raw tokens
(``seq_len // ratio - prefix_len // ratio``).
3. Return a :class:`DSV4OutCacheLoc` containing only KV slot families.
State slots are paged because the NPU fused compressor runs ``cache_mode=1``; the
base class' ``translate_kv_loc_to_compress_state_loc`` ring-hash is the CUDA-only
path and is unused on NPU. The bundle is the explicit return value:
The bundle is the explicit return value:
mem_cache/common.py unpacks ``out_full_loc`` and stashes the bundle on
``batch.out_cache_loc_dsv4``; ``DSV4NPUReqToTokenPool`` writes the per-req
``req_to_token_c{4,128}[_state]`` tables that :meth:`free` and the last_loc
lookups read back.
``req_to_c128_sidecar`` table that :meth:`free` and the last_loc lookup read back.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, List, Optional
from typing import Optional
import torch
@@ -35,34 +31,29 @@ from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
)
from sglang.srt.mem_cache.allocation import alloc_paged_token_slots_extend
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc, DSV4StateLens
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc
def get_last_loc(
req_to_token: torch.Tensor,
req_to_c128_sidecar: torch.Tensor,
req_pool_indices: torch.Tensor,
prefix_lens: torch.Tensor,
page_size: int,
) -> torch.Tensor:
"""Slot id of each req's last already-allocated token, or -1 when
``prefix_lens[i] == 0`` (fresh req).
Looks up ``req_to_token[req, prefix_lens - 1]`` to anchor the paged
allocator's ``alloc_extend`` on the real previous tail slot, preserving the
intra-page slot continuity the kernel's ``cmp_block_table`` relies on (the
allocator debug-asserts ``(last_loc + 1) % page_size == prefix_lens %
page_size``). Result dtype matches ``prefix_lens``.
Looks up the C128 sidecar page to anchor the paged allocator's
``alloc_extend`` on the real previous tail slot, preserving intra-page slot
continuity. Result dtype matches ``prefix_lens``.
"""
req_pool_indices = req_pool_indices.to(torch.int64)
safe_idx = (prefix_lens.to(torch.int64) - 1).clamp(min=0)
looked_up = req_to_token[req_pool_indices, safe_idx].to(prefix_lens.dtype)
return torch.where(
prefix_lens > 0,
looked_up,
torch.full_like(prefix_lens, -1),
last_pos = (prefix_lens.to(torch.int64) - 1).clamp(min=0)
page_ids = req_to_c128_sidecar[req_pool_indices, last_pos // page_size].to(
prefix_lens.dtype
)
last_loc = page_ids * page_size + last_pos.to(prefix_lens.dtype) % page_size
return torch.where(prefix_lens > 0, last_loc, torch.full_like(prefix_lens, -1))
def alloc_paged_token_slots_extend_npu(*args, batch=None, **kwargs):
@@ -81,20 +72,9 @@ def alloc_paged_token_slots_reserve_extend(
extend_num_tokens: int,
*,
req_pool_indices: Optional[torch.Tensor] = None,
dsv4_state_lens: Optional[DSV4StateLens] = None,
batch=None,
):
"""Allocate reserved draft slots and update DSV4 per-request tables."""
if dsv4_state_lens is None and batch is not None:
allocator = batch.token_to_kv_pool_allocator
dsv4_state_lens = (
allocator.compute_dsv4_state_lens_reserve(
batch.reqs, prefix_lens_cpu, seq_lens_cpu
)
if hasattr(allocator, "compute_dsv4_state_lens_reserve")
else None
)
"""Allocate reserved draft KV slots and update DSV4 KV tables."""
out_cache_loc = alloc_paged_token_slots_extend(
tree_cache,
prefix_lens,
@@ -104,7 +84,6 @@ def alloc_paged_token_slots_reserve_extend(
last_loc,
extend_num_tokens,
req_pool_indices=req_pool_indices,
dsv4_state_lens=dsv4_state_lens,
batch=batch,
)
if batch is not None:
@@ -113,14 +92,12 @@ def alloc_paged_token_slots_reserve_extend(
batch.req_pool_indices_cpu,
prefix_lens_cpu,
seq_lens_cpu,
c4_state_alloc_offsets=prefix_lens_cpu,
c128_state_alloc_offsets=prefix_lens_cpu,
)
return out_cache_loc
class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
"""SWA allocator + c4/c128 KV and compress-state paged allocators for DSV4 on NPU."""
"""SWA allocator + C128 KV allocator and full-derived C4 locations."""
def __init__(
self,
@@ -143,54 +120,29 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
)
def mk(pool_size, pool):
# c4/c128 KV and state sub-pools implement KVCache, so they drop into
# the standard paged allocator. pool_size is in compressed-token units.
# C128 KV sub-pool implements KVCache, so it drops into the standard
# paged allocator. pool_size is in compressed-token units.
return NPUPagedTokenToKVPoolAllocator(
pool_size,
page_size=page_size,
page_size=pool.kernel_page_size,
dtype=dtype,
device=device,
kvcache=pool,
need_sort=need_sort,
)
self.c4_attn_allocator = mk(kvcache.c4_size, kvcache.c4_kv_pool)
self.c128_attn_allocator = mk(kvcache.c128_size, kvcache.c128_kv_pool)
# State allocators (paged, NPU-only). Any layer's pool works as KVCache
# pointer (slot alloc is layer-agnostic); None when no c{ratio} layers or
# zero budget.
self.c4_state_attn_allocator: Optional[NPUPagedTokenToKVPoolAllocator] = None
self.c128_state_attn_allocator: Optional[NPUPagedTokenToKVPoolAllocator] = None
state_pools = getattr(kvcache, "compress_state_pools", None)
if state_pools:
def first_state_pool(want_ratio):
return next(
(
p
for r, p in zip(kvcache.compression_ratios, state_pools)
if r == want_ratio and p is not None
),
None,
)
c4_state_pool = first_state_pool(4)
c128_state_pool = first_state_pool(128)
if c4_state_pool is not None and kvcache.c4_state_pool_size > 0:
self.c4_state_attn_allocator = mk(
kvcache.c4_state_pool_size, c4_state_pool
)
if c128_state_pool is not None and kvcache.c128_state_pool_size > 0:
self.c128_state_attn_allocator = mk(
kvcache.c128_state_pool_size, c128_state_pool
)
self.c128_page_refcount = torch.zeros(
self.c128_attn_allocator.num_pages + 1,
dtype=torch.int32,
device=device,
)
# Returned by the c-pool helpers when a step adds no compressed tokens.
self._empty_loc = torch.empty((0,), dtype=torch.int64, device=device)
# Per-call handle to the DSV4NPUReqToTokenPool, stashed by alloc_extend/
# alloc_decode for last_loc lookups; avoids a permanent allocator->pool ref.
# alloc_decode for the C128 KV last_loc lookup.
self._cur_req_to_token_pool = None
@staticmethod
@@ -206,6 +158,49 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
diff = ((seq_lens_cpu // ratio) - (prefix_lens_cpu // ratio)).clamp(min=0)
return int(diff.sum().item())
@staticmethod
def _derive_c4_loc_from_full(out_full_loc: torch.Tensor) -> torch.Tensor:
"""Map full slots closing a 4-token group to their C4 slots."""
completed_group = (out_full_loc >= 0) & ((out_full_loc % 4) == 3)
return out_full_loc[completed_group] // 4
def retain_c128_pages(self, page_ids: torch.Tensor) -> None:
page_ids = page_ids.to(torch.int64).view(-1)
if page_ids.numel() == 0:
return
self.c128_page_refcount.index_add_(
0,
page_ids,
torch.ones_like(page_ids, dtype=self.c128_page_refcount.dtype),
)
def release_c128_pages(self, page_ids: torch.Tensor) -> None:
page_ids = torch.unique(page_ids.to(torch.int64).view(-1))
page_ids = page_ids[page_ids > 0]
if page_ids.numel() == 0:
return
self.c128_page_refcount.index_add_(
0,
page_ids,
-torch.ones_like(page_ids, dtype=self.c128_page_refcount.dtype),
)
free_pages = page_ids[self.c128_page_refcount[page_ids] == 0]
if free_pages.numel() > 0:
self.c128_attn_allocator.free(
free_pages * self.c128_attn_allocator.page_size
)
def replace_req_c128_prefix(
self, req_pool_idx: int, page_ids: torch.Tensor, req_to_token_pool
) -> None:
table = req_to_token_pool.req_to_c128_sidecar
page_ids = page_ids.to(device=table.device, dtype=table.dtype).view(-1)
old = table[req_pool_idx, : page_ids.numel()].clone()
changed = old != page_ids
self.release_c128_pages(old[changed])
self.retain_c128_pages(page_ids[changed])
table[req_pool_idx, : page_ids.numel()] = page_ids
@staticmethod
def _pool_exhausted(
ratio: int, kind: str, need: int, available: int
@@ -218,60 +213,6 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
f"on req finish."
)
def _alloc_state_extend(
self,
allocator: Optional[NPUPagedTokenToKVPoolAllocator],
raw_prefix_lens: torch.Tensor,
state_prefix_lens: torch.Tensor,
state_prefix_lens_cpu: torch.Tensor,
state_seq_lens: torch.Tensor,
state_seq_lens_cpu: torch.Tensor,
req_pool_indices: torch.Tensor,
last_loc_dtype: torch.dtype,
state_extend_num_tokens: int,
ratio: int,
) -> torch.Tensor:
"""Allocate tail-only state-pool slots for an extend at ``ratio``.
The state pool is a separate paged slot space; each req allocates only
its trailing window (cumulative lens precomputed by
``ScheduleBatch._compute_dsv4_state_lens_*`` and passed via
``DSV4StateLens``). ``state_last_loc`` is looked up from
``req_to_token_c{ratio}_state`` at the RAW position
``raw_prefix_lens - 1`` (the last position the previous extend/decode
populated). Returns ``_empty_loc`` when the allocator is absent (no
c{ratio} layers) or there is nothing to add.
"""
if allocator is None or state_extend_num_tokens == 0:
return self._empty_loc
assert self._cur_req_to_token_pool is not None, (
"alloc_extend/alloc_decode must be called with req_to_token_pool= "
"for the state-pool last_loc lookup."
)
state_table = (
self._cur_req_to_token_pool.req_to_token_c4_state
if ratio == 4
else self._cur_req_to_token_pool.req_to_token_c128_state
)
state_last_loc = get_last_loc(
state_table, req_pool_indices, raw_prefix_lens
).to(last_loc_dtype)
result = allocator.alloc_extend(
state_prefix_lens,
state_prefix_lens_cpu,
state_seq_lens,
state_seq_lens_cpu,
state_last_loc,
state_extend_num_tokens,
)
if result is None:
raise self._pool_exhausted(
ratio, "state", state_extend_num_tokens, allocator.available_size()
)
return result
def _alloc_c_extend(
self,
allocator: NPUPagedTokenToKVPoolAllocator,
@@ -286,7 +227,7 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
"""Allocate compressed-KV slots for an extend at ``ratio``.
Prefix/seq lens are translated to compressed units (``// ratio``); the
c-pool last_loc comes from ``req_to_token_c{ratio}`` via
c-pool last_loc comes from ``req_to_c128_sidecar`` via
:func:`get_last_loc` so the paged allocator continues in-page (or opens
a fresh page at a ratio boundary), keeping the intra-page continuity the
``cmp_block_table`` reader relies on. Returns ``_empty_loc`` when this
@@ -300,16 +241,12 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
"alloc_extend/alloc_decode must be called with req_to_token_pool= "
"for the c-pool last_loc lookup."
)
c_table = (
self._cur_req_to_token_pool.req_to_token_c4
if ratio == 4
else self._cur_req_to_token_pool.req_to_token_c128
)
c_table = self._cur_req_to_token_pool.req_to_c128_sidecar
c_prefix = (prefix_lens // ratio).to(prefix_lens.dtype)
c_seq = (seq_lens // ratio).to(seq_lens.dtype)
c_last_loc = get_last_loc(c_table, req_pool_indices, c_prefix).to(
last_loc_dtype
)
c_last_loc = get_last_loc(
c_table, req_pool_indices, c_prefix, allocator.page_size
).to(last_loc_dtype)
result = allocator.alloc_extend(
c_prefix,
@@ -325,7 +262,17 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
)
return result
def _alloc_c_and_state(
def _has_c128_sidecar_capacity(
self, prefix_lens_cpu: torch.Tensor, seq_lens_cpu: torch.Tensor
) -> bool:
ratio = 128
page_size = self.c128_attn_allocator.page_size
prefix_groups = (prefix_lens_cpu // ratio + page_size - 1) // page_size
seq_groups = (seq_lens_cpu // ratio + page_size - 1) // page_size
need = int((seq_groups - prefix_groups).clamp(min=0).sum().item())
return need <= self.c128_attn_allocator.available_size() // page_size
def _alloc_compressed_kv(
self,
out_full_loc: torch.Tensor,
out_swa_loc: torch.Tensor,
@@ -335,57 +282,13 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
seq_lens_cpu: torch.Tensor,
last_loc_dtype: torch.dtype,
req_pool_indices: Optional[torch.Tensor],
dsv4_state_lens: Optional[DSV4StateLens],
) -> DSV4OutCacheLoc:
"""Allocate c4/c128 KV + state slots and bundle them with full/swa loc.
Shared by alloc_extend / alloc_decode (which differ only in how
prefix_lens is derived). State lens are tail-only, precomputed by
ScheduleBatch._compute_dsv4_state_lens_*; raw prefix_lens drives the
state last_loc lookup.
"""
"""Allocate C128 KV and derive C4 KV, then bundle all KV locations."""
assert req_pool_indices is not None, (
"DSV4NPUTokenToKVPoolAllocator requires req_pool_indices "
"(forwarded from batch.req_pool_indices)."
)
if dsv4_state_lens is not None:
out_c4_state_loc = self._alloc_state_extend(
self.c4_state_attn_allocator,
prefix_lens,
dsv4_state_lens.c4_prefix_lens,
dsv4_state_lens.c4_prefix_lens_cpu,
dsv4_state_lens.c4_seq_lens,
dsv4_state_lens.c4_seq_lens_cpu,
req_pool_indices,
last_loc_dtype,
dsv4_state_lens.c4_extend_num_tokens,
ratio=4,
)
out_c128_state_loc = self._alloc_state_extend(
self.c128_state_attn_allocator,
prefix_lens,
dsv4_state_lens.c128_prefix_lens,
dsv4_state_lens.c128_prefix_lens_cpu,
dsv4_state_lens.c128_seq_lens,
dsv4_state_lens.c128_seq_lens_cpu,
req_pool_indices,
last_loc_dtype,
dsv4_state_lens.c128_extend_num_tokens,
ratio=128,
)
else:
out_c4_state_loc = self._empty_loc
out_c128_state_loc = self._empty_loc
out_c4_loc = self._alloc_c_extend(
self.c4_attn_allocator,
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
req_pool_indices,
last_loc_dtype,
ratio=4,
)
out_c4_loc = self._derive_c4_loc_from_full(out_full_loc)
out_c128_loc = self._alloc_c_extend(
self.c128_attn_allocator,
prefix_lens,
@@ -401,173 +304,6 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
out_swa_loc=out_swa_loc,
out_c4_loc=out_c4_loc,
out_c128_loc=out_c128_loc,
out_c4_state_loc=out_c4_state_loc,
out_c128_state_loc=out_c128_state_loc,
)
def compute_dsv4_state_lens_extend(
self, reqs: List[Req], seq_lens: List[int], prefix_lens: List[int]
) -> Optional[DSV4StateLens]:
"""Per-req c{4,128}_state pool alloc lens for extend (tail-only).
State pool stores only the trailing portion of each sequence (the c{N}
compressor's read/write window); the tail length depends on raw
seq_len's alignment to the SWA page boundary (128)::
c4_alloc_len = tail + 128 if (tail <= 3 and seq_len >= 128) else tail
c128_alloc_len = tail where tail = seq_len % 128
Long prefills allocate only the trailing partial window, not slots for
already-compressed positions, so the small paged state pool (~256
slots/req) stays sufficient even for 28k-token prompts.
Mutates per-req cumulative state via getattr/setattr so the community
``Req`` needs no DSV4 field declarations:
* ``req.c{4,128}_state_kv_len`` — cumulative slot count (prefix for
the paged allocator; never decreases on eviction).
* ``req.c{4,128}_state_alloc_offset`` — low-water raw-position mark
for eviction (see ``dsv4_common_hooks.maybe_evict_dsv4_state``).
Returns None when this model has no paged state pools (CUDA / non-V4 /
zero budget) — callers pass that straight through as ``dsv4_state_lens``.
"""
if self.c4_state_attn_allocator is None:
return None
c4_prefix: List[int] = []
c4_seq: List[int] = []
c128_prefix: List[int] = []
c128_seq: List[int] = []
for req, seq_len, prefix_len in zip(reqs, seq_lens, prefix_lens):
tail = seq_len % 128
c4_alloc_len = tail + 128 if (tail <= 3 and seq_len >= 128) else tail
c128_alloc_len = tail
chunk_len = seq_len - prefix_len
if prefix_len > 0:
c4_count = min(c4_alloc_len, chunk_len)
c128_count = min(c128_alloc_len, chunk_len)
else:
c4_count = c4_alloc_len
c128_count = c128_alloc_len
req.c4_state_alloc_offset = seq_len - c4_alloc_len
req.c128_state_alloc_offset = seq_len - c128_alloc_len
prev_c4 = getattr(req, "c4_state_kv_len", 0)
prev_c128 = getattr(req, "c128_state_kv_len", 0)
new_c4 = prev_c4 + c4_count
new_c128 = prev_c128 + c128_count
c4_prefix.append(prev_c4)
c4_seq.append(new_c4)
c128_prefix.append(prev_c128)
c128_seq.append(new_c128)
req.c4_state_kv_len = new_c4
req.c128_state_kv_len = new_c128
req.c4_state_write_offset = seq_len - c4_count
req.c128_state_write_offset = seq_len - c128_count
return self._pack_state_lens(
c4_prefix,
c4_seq,
c128_prefix,
c128_seq,
c4_extend_num_tokens=int(sum(s - p for s, p in zip(c4_seq, c4_prefix))),
c128_extend_num_tokens=int(
sum(s - p for s, p in zip(c128_seq, c128_prefix))
),
)
def compute_dsv4_state_lens_decode(
self, reqs: List[Req]
) -> Optional[DSV4StateLens]:
"""Per-req c{4,128}_state pool alloc lens for decode: exactly 1 new
state slot per req per pool. ``c{N}_state_alloc_offset`` does NOT
advance here (only eviction advances it). Returns None when there are
no paged state pools."""
if self.c4_state_attn_allocator is None:
return None
c4_prefix: List[int] = []
c4_seq: List[int] = []
c128_prefix: List[int] = []
c128_seq: List[int] = []
for req in reqs:
prev_c4 = getattr(req, "c4_state_kv_len", 0)
prev_c128 = getattr(req, "c128_state_kv_len", 0)
c4_prefix.append(prev_c4)
c4_seq.append(prev_c4 + 1)
c128_prefix.append(prev_c128)
c128_seq.append(prev_c128 + 1)
req.c4_state_kv_len = prev_c4 + 1
req.c128_state_kv_len = prev_c128 + 1
bs = len(reqs)
return self._pack_state_lens(
c4_prefix,
c4_seq,
c128_prefix,
c128_seq,
c4_extend_num_tokens=bs,
c128_extend_num_tokens=bs,
)
def compute_dsv4_state_lens_reserve(
self, reqs: List[Req], prefix_lens: List[int], seq_lens: List[int]
) -> Optional[DSV4StateLens]:
"""Allocate state slots for a speculative pre-reserved raw interval."""
if self.c4_state_attn_allocator is None:
return None
c4_prefix: List[int] = []
c4_seq: List[int] = []
c128_prefix: List[int] = []
c128_seq: List[int] = []
for req, prefix_len, seq_len in zip(reqs, prefix_lens, seq_lens):
reserve = max(0, int(seq_len) - int(prefix_len))
prev_c4 = getattr(req, "c4_state_kv_len", 0)
prev_c128 = getattr(req, "c128_state_kv_len", 0)
c4_prefix.append(prev_c4)
c4_seq.append(prev_c4 + reserve)
c128_prefix.append(prev_c128)
c128_seq.append(prev_c128 + reserve)
req.c4_state_kv_len = prev_c4 + reserve
req.c128_state_kv_len = prev_c128 + reserve
total = sum(max(0, int(s) - int(p)) for p, s in zip(prefix_lens, seq_lens))
return self._pack_state_lens(
c4_prefix,
c4_seq,
c128_prefix,
c128_seq,
c4_extend_num_tokens=total,
c128_extend_num_tokens=total,
)
def _pack_state_lens(
self,
c4_prefix: List[int],
c4_seq: List[int],
c128_prefix: List[int],
c128_seq: List[int],
*,
c4_extend_num_tokens: int,
c128_extend_num_tokens: int,
) -> DSV4StateLens:
c4_prefix_cpu = torch.tensor(c4_prefix, dtype=torch.int64)
c4_seq_cpu = torch.tensor(c4_seq, dtype=torch.int64)
c128_prefix_cpu = torch.tensor(c128_prefix, dtype=torch.int64)
c128_seq_cpu = torch.tensor(c128_seq, dtype=torch.int64)
return DSV4StateLens(
c4_prefix_lens=c4_prefix_cpu.to(self.device, non_blocking=True),
c4_prefix_lens_cpu=c4_prefix_cpu,
c4_seq_lens=c4_seq_cpu.to(self.device, non_blocking=True),
c4_seq_lens_cpu=c4_seq_cpu,
c4_extend_num_tokens=c4_extend_num_tokens,
c128_prefix_lens=c128_prefix_cpu.to(self.device, non_blocking=True),
c128_prefix_lens_cpu=c128_prefix_cpu,
c128_seq_lens=c128_seq_cpu.to(self.device, non_blocking=True),
c128_seq_lens_cpu=c128_seq_cpu,
c128_extend_num_tokens=c128_extend_num_tokens,
)
def alloc_extend(
@@ -580,12 +316,13 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
extend_num_tokens: int,
*,
req_pool_indices: Optional[torch.Tensor] = None,
dsv4_state_lens: Optional[DSV4StateLens] = None,
req_to_token_pool=None,
) -> Optional[DSV4OutCacheLoc]:
# Stash per-req tables for this call's last_loc lookups (read by
# _alloc_c_extend / _alloc_state_extend); no permanent allocator->pool ref.
self._cur_req_to_token_pool = req_to_token_pool
if not self._has_c128_sidecar_capacity(prefix_lens_cpu, seq_lens_cpu):
return None
out_full_loc = super().alloc_extend(
prefix_lens,
prefix_lens_cpu,
@@ -602,7 +339,6 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
seq_lens_cpu,
last_loc.dtype,
req_pool_indices,
dsv4_state_lens,
)
def _wrap_full_alloc(
@@ -614,10 +350,9 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
seq_lens_cpu,
loc_dtype,
req_pool_indices,
dsv4_state_lens,
) -> Optional[DSV4OutCacheLoc]:
# Shared tail of alloc_extend / alloc_extend_swa_tail: translate the full
# loc to swa, then add the c4/c128(+state) pools into a DSV4OutCacheLoc.
# loc to swa, then add the c4/c128 KV pools into a DSV4OutCacheLoc.
if out_full_loc is None:
return None
out_swa_loc = self.translate_loc_from_full_to_swa(out_full_loc)
@@ -625,7 +360,7 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
"translate_loc_from_full_to_swa returned None — "
"full_to_swa_index_mapping not initialized?"
)
return self._alloc_c_and_state(
return self._alloc_compressed_kv(
out_full_loc,
out_swa_loc,
prefix_lens,
@@ -634,7 +369,6 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
seq_lens_cpu,
loc_dtype,
req_pool_indices,
dsv4_state_lens,
)
def alloc_decode(
@@ -644,10 +378,13 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
last_loc: torch.Tensor,
*,
req_pool_indices: Optional[torch.Tensor] = None,
dsv4_state_lens: Optional[DSV4StateLens] = None,
req_to_token_pool=None,
) -> Optional[DSV4OutCacheLoc]:
self._cur_req_to_token_pool = req_to_token_pool
if not self._has_c128_sidecar_capacity(
(seq_lens_cpu - 1).clamp(min=0), seq_lens_cpu
):
return None
out_full_loc = super().alloc_decode(seq_lens, seq_lens_cpu, last_loc)
if out_full_loc is None:
return None
@@ -657,7 +394,7 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
# seq_len//ratio so _alloc_c_extend anchors on the real c-pool last_loc.
prefix_lens = (seq_lens - 1).clamp(min=0)
prefix_lens_cpu = (seq_lens_cpu - 1).clamp(min=0)
return self._alloc_c_and_state(
return self._alloc_compressed_kv(
out_full_loc,
out_swa_loc,
prefix_lens,
@@ -666,7 +403,6 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
seq_lens_cpu,
last_loc.dtype,
req_pool_indices,
dsv4_state_lens,
)
def alloc_extend_swa_tail(
@@ -680,13 +416,14 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
swa_tail_len: int,
*,
req_pool_indices: Optional[torch.Tensor] = None,
dsv4_state_lens: Optional[DSV4StateLens] = None,
req_to_token_pool=None,
) -> Optional[DSV4OutCacheLoc]:
"""Disagg-decode prealloc variant of :meth:`alloc_extend`: super() does
full+swa-tail, then _alloc_c_and_state adds c4/c128(+state) → DSV4OutCacheLoc.
full+swa-tail, then _alloc_compressed_kv adds c4/c128 KV → DSV4OutCacheLoc.
"""
self._cur_req_to_token_pool = req_to_token_pool
if not self._has_c128_sidecar_capacity(prefix_lens_cpu, seq_lens_cpu):
return None
out_full_loc = super().alloc_extend_swa_tail(
prefix_lens,
prefix_lens_cpu,
@@ -704,7 +441,6 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
seq_lens_cpu,
last_loc.dtype,
req_pool_indices,
dsv4_state_lens,
)
def free(
@@ -714,20 +450,15 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
req=None,
req_to_token_pool=None,
):
"""Unified free for full/swa/c4/c128 pools. Two forms (may co-fire):
"""Unified free for full/SWA/C4/C128 KV and C128 request state.
Two forms may co-fire:
* ``free(free_index)`` — full + SWA only (tail/radix eviction; no req
identity, so c-pool free can't run).
* ``free(req=, req_to_token_pool=)`` — from DSV4NPUReqToTokenPool.free
on req finish: reads the per-req slot lists from
``req_to_token_c{4,128}[_state]`` and returns them to the c-pools
(the paged allocator dedupes by page).
KV pools free ``[0, kv_len // ratio)``. State pools are 1-per-raw-token
and free only the tail ``[c{N}_state_alloc_offset, kv_len)`` — the prefix
was already returned by ScheduleBatch._evict_swa (state rides SWA
eviction); freeing it again would double-free (caught by the paged
allocator's debug_mode assert, corrupts the free list otherwise).
on request finish: returns C128 KV pages and clears that request's
fixed C128 state bank before the req_pool_idx can be reused.
"""
if free_index is not None:
super().free(free_index)
@@ -739,53 +470,31 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
if kv_len <= 0 or req_pool_idx is None:
return
# KV pools: free the leading [0, kv_len // ratio) compressed slots.
for ratio, allocator, table_attr in (
(4, self.c4_attn_allocator, "req_to_token_c4"),
(128, self.c128_attn_allocator, "req_to_token_c128"),
):
n = kv_len // ratio
if n > 0 and hasattr(req_to_token_pool, table_attr):
slots = getattr(req_to_token_pool, table_attr)[req_pool_idx, :n]
slots = slots[slots > 0]
# to int64 — paged allocator's free does cpu()//page_size on it.
if slots.numel() > 0:
allocator.free(slots.to(torch.int64))
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))
# State pools: free only the tail [c{N}_state_alloc_offset, kv_len).
for ratio, allocator, table_attr, off_attr in (
(
4,
self.c4_state_attn_allocator,
"req_to_token_c4_state",
"c4_state_alloc_offset",
),
(
128,
self.c128_state_attn_allocator,
"req_to_token_c128_state",
"c128_state_alloc_offset",
),
):
if allocator is None or not hasattr(req_to_token_pool, table_attr):
continue
off = getattr(req, off_attr, 0)
if kv_len > off:
slots = getattr(req_to_token_pool, table_attr)[req_pool_idx, off:kv_len]
slots = slots[slots > 0]
if slots.numel() > 0:
allocator.free(slots.to(torch.int64))
def available_size(self):
return min(
super().available_size(),
self.c128_attn_allocator.available_size() * 128,
)
def resize(self, config) -> None:
self.c128_attn_allocator.size = int(config.c128_max_total_num_tokens)
self.c128_attn_allocator.num_pages = (
self.c128_attn_allocator.size // self.c128_attn_allocator.page_size
)
super().resize(config)
def clear(self):
super().clear()
# super().__init__ calls clear() before our sub-allocators exist;
# getattr(..., None) tolerates that and the always-None state allocators.
for attr in (
"c4_attn_allocator",
"c128_attn_allocator",
"c4_state_attn_allocator",
"c128_state_attn_allocator",
):
# super().__init__ calls clear() before our C128 allocator exists.
for attr in ("c128_attn_allocator",):
allocator = getattr(self, attr, None)
if allocator is not None:
allocator.clear()
refcount = getattr(self, "c128_page_refcount", None)
if refcount is not None:
refcount.zero_()
@@ -1,4 +1,4 @@
"""Helpers used by mem_cache/common.py to wire DSV4-NPU per-req tables.
"""Helpers used by mem_cache/common.py to wire DSV4-NPU KV tables.
mem_cache/common.py runs platform-agnostic alloc flow. When the model is
DSV4 on NPU, ``alloc_paged_token_slots_{extend,decode}`` already stashed the
@@ -7,20 +7,23 @@ DSV4 on NPU, ``alloc_paged_token_slots_{extend,decode}`` already stashed the
these hooks then:
1. Read the bundle from ``batch.out_cache_loc_dsv4``.
2. Write the per-pool slot ids into the per-req tables on the
:class:`DSV4NPUReqToTokenPool`.
2. Write newly allocated C128 page ids into the per-request sidecar.
Compressor state is fixed ring storage and does not participate in this
allocation/write path. PD reuses the public SWA/C128-state payloads and only
builds an NPU-specific payload for the independently addressed C128 KV pool.
Non-DSV4 paths leave ``batch.out_cache_loc_dsv4`` None, so this module is a
no-op for them.
The disagg per-req prealloc path does not build a ``ScheduleBatch`` and so
bypasses the batch hook; it writes the same tables via
bypasses the batch hook; it writes the same sidecar via
``write_dsv4_prealloc_tables`` (driven by ``dsv4_unwrap_prealloc``).
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Sequence
from typing import TYPE_CHECKING
import torch
@@ -33,19 +36,13 @@ def maybe_write_dsv4_extend(
req_pool_indices_cpu: torch.Tensor,
prefix_lens_cpu: torch.Tensor,
seq_lens_cpu: torch.Tensor,
*,
c4_state_alloc_offsets: Sequence[int] | torch.Tensor | None = None,
c128_state_alloc_offsets: Sequence[int] | torch.Tensor | None = None,
) -> None:
"""Post-alloc_extend hook for DSV4. No-op when allocator/pool is not DSV4.
For each compressed pool (c4 / c128), spreads the flat
``out_c{4,128}_loc`` tensor across requests using per-req extend
counts (``seq_lens[i] // ratio - prefix_lens[i] // ratio``) and writes
the resulting slot ids into ``req_to_token_c{4,128}[req, prefix:seq]``.
Spreads the flat ``out_c128_loc`` tensor across requests and writes newly
allocated page ids into ``req_to_c128_sidecar``. C4 locations are derived
from the full-token table.
Also writes ``req_to_token_swa[req, prefix:seq]`` with the swa slots
derived from out_full_loc via the SWA index mapping.
"""
# Bundle stashed on batch.out_cache_loc_dsv4 by mem_cache/common.py;
# None on CUDA / non-V4 paths → no-op.
@@ -54,33 +51,15 @@ def maybe_write_dsv4_extend(
return
req_to_token_pool = batch.req_to_token_pool
if not hasattr(req_to_token_pool, "write_c4"):
if not hasattr(req_to_token_pool, "write_c128"):
return # non-DSV4 pool; skip defensively (shouldn't happen)
# c4_state / c128_state writes: tail-only. Bundle length is
# sum(c{N}_state_alloc_len_i), NOT total raw extend tokens. Normal extend
# uses the per-Req low-water marks; reserve callers can pass explicit raw
# offsets for the pre-reserved interval.
if c4_state_alloc_offsets is None:
c4_state_alloc_offsets = [
getattr(r, "c4_state_write_offset", getattr(r, "c4_state_alloc_offset", 0))
for r in batch.reqs
]
if c128_state_alloc_offsets is None:
c128_state_alloc_offsets = [
getattr(
r, "c128_state_write_offset", getattr(r, "c128_state_alloc_offset", 0)
)
for r in batch.reqs
]
_write_dsv4_tables(
req_to_token_pool,
req_pool_indices_cpu,
prefix_lens_cpu,
seq_lens_cpu,
bundle,
c4_state_offsets=c4_state_alloc_offsets,
c128_state_offsets=c128_state_alloc_offsets,
)
@@ -89,18 +68,10 @@ def dsv4_state_payloads(
req_pool_idx: int,
seq_len: int,
page_size: int,
window_size: int,
*,
prefix_len: int = 0,
):
"""Per-StateType PD-payload builders for DSV4-on-NPU.
For chunked prefill, intermediate chunks can leave old C4/C128 state pages in
the req table. PD only needs the final active tail state; scanning the whole
prompt span would transfer stale state pages and can perturb decode accuracy.
"""
if not hasattr(req_to_token_pool, "req_to_token_c4"):
return {}
"""Build the only NPU-specific DSV4 PD payload: C128 KV pages."""
import numpy as np
@@ -109,103 +80,38 @@ def dsv4_state_payloads(
seq_len = max(0, int(seq_len))
prefix_len = max(0, min(int(prefix_len), seq_len))
def empty_pages():
return np.empty((0,), dtype=np.int32)
def pages(table, lo: int, hi: int, *, drop_zero_pages: bool = False):
def c128_kv_pages():
c128_page_size = req_to_token_pool.c128_page_size
lo = prefix_len // (128 * c128_page_size)
hi = (seq_len // 128 + c128_page_size - 1) // c128_page_size
if hi <= lo:
return empty_pages()
return np.empty((0,), dtype=np.int32)
pages = (
req_to_token_pool.req_to_c128_sidecar[req_pool_idx, lo:hi]
.cpu()
.numpy()
.astype(np.int32)
)
return pages[pages > 0]
lo = max(0, int(lo))
hi = max(lo, int(hi))
page_lo = (lo // page_size) * page_size
page_hi = ((hi + page_size - 1) // page_size) * page_size
if page_hi <= page_lo:
return empty_pages()
slots = table[req_pool_idx, page_lo:page_hi:page_size].cpu().numpy()
if slots.size == 0:
return empty_pages()
page_indices = (slots // page_size).astype(np.int32)
if drop_zero_pages:
page_indices = page_indices[page_indices > 0]
return page_indices
def state_tail_range(compress_ratio: int):
tail_len = seq_len % 128
if compress_ratio == 4:
state_len = tail_len + 128 if tail_len <= 3 and seq_len >= 128 else tail_len
elif compress_ratio == 128:
state_len = tail_len
else:
raise ValueError(f"Unsupported DSV4 state compress ratio: {compress_ratio}")
if state_len == 0:
return None
start = max(prefix_len, seq_len - state_len)
if start >= seq_len:
return None
return start, seq_len
def state_pages(table, compress_ratio: int):
state_range = state_tail_range(compress_ratio)
if state_range is None:
return empty_pages()
lo, hi = state_range
return pages(table, lo, hi, drop_zero_pages=True)
if window_size is None or window_size <= 0:
window_start = prefix_len
else:
window_start = max(prefix_len, seq_len - window_size)
window_start = (window_start // page_size) * page_size
# DSV4_INDEXER shares the c4 slot space (written at the c4 loc).
return {
AscendStateType.DSV4_SWA: lambda: pages(
req_to_token_pool.req_to_token_swa,
window_start,
seq_len,
drop_zero_pages=True,
),
AscendStateType.DSV4_C4: lambda: pages(
req_to_token_pool.req_to_token_c4, prefix_len // 4, seq_len // 4
),
AscendStateType.DSV4_C128: lambda: pages(
req_to_token_pool.req_to_token_c128, prefix_len // 128, seq_len // 128
),
AscendStateType.DSV4_INDEXER: lambda: pages(
req_to_token_pool.req_to_token_c4, prefix_len // 4, seq_len // 4
),
AscendStateType.DSV4_C4_STATE: lambda: state_pages(
req_to_token_pool.req_to_token_c4_state, 4
),
AscendStateType.DSV4_C128_STATE: lambda: state_pages(
req_to_token_pool.req_to_token_c128_state, 128
),
}
return {AscendStateType.DSV4_C128: c128_kv_pages}
def dsv4_prealloc_kwargs(allocator, req, fill_len, req_to_token_pool, *, device):
"""Extra ``alloc_extend(_swa_tail)`` kwargs for the DSV4 allocator; ``{}`` for
non-DSV4 so callers can splat it unconditionally."""
if not hasattr(allocator, "c4_attn_allocator"):
if not hasattr(allocator, "c128_attn_allocator"):
return {}
return dict(
req_pool_indices=torch.tensor(
[req.req_pool_idx], dtype=torch.int64, device=device
),
dsv4_state_lens=allocator.compute_dsv4_state_lens_extend(
[req], [fill_len], [0]
),
req_to_token_pool=req_to_token_pool,
)
def dsv4_unwrap_prealloc(kv_loc, req_to_token_pool, req, prefix_len, fill_len):
"""Unwrap a DSV4OutCacheLoc bundle to its full-pool loc and write the five
"""Unwrap a DSV4OutCacheLoc bundle to its full-pool loc and write the
per-req tables; a plain tensor (non-DSV4) passes through unchanged."""
if kv_loc is None or not hasattr(kv_loc, "out_full_loc"):
return kv_loc
@@ -220,9 +126,9 @@ def write_dsv4_prealloc_tables(
fill_len: int,
bundle,
) -> None:
"""Write the five DSV4 per-req tables for one request on the disagg-decode
"""Write the DSV4 per-req tables for one request on the disagg-decode
prealloc path (no ScheduleBatch); no-op without bundle / DSV4 tables."""
if bundle is None or not hasattr(req_to_token_pool, "write_c4"):
if bundle is None or not hasattr(req_to_token_pool, "write_c128"):
return
rp = torch.tensor([req.req_pool_idx])
pl = torch.tensor([prefix_len])
@@ -234,8 +140,6 @@ def write_dsv4_prealloc_tables(
pl,
sl,
bundle,
c4_state_offsets=[getattr(req, "c4_state_alloc_offset", 0)],
c128_state_offsets=[getattr(req, "c128_state_alloc_offset", 0)],
)
@@ -245,27 +149,8 @@ def _write_dsv4_tables(
prefix_lens_cpu: torch.Tensor,
seq_lens_cpu: torch.Tensor,
bundle,
*,
c4_state_offsets: Sequence[int] | torch.Tensor,
c128_state_offsets: Sequence[int] | torch.Tensor,
) -> None:
"""Write DSV4 SWA, compressed-KV, and compression-state tables."""
_write_per_req_slice(
req_to_token_pool.write_swa,
req_pool_indices_cpu,
prefix_lens_cpu,
seq_lens_cpu,
bundle.out_swa_loc,
ratio=1,
)
_write_per_req_slice(
req_to_token_pool.write_c4,
req_pool_indices_cpu,
prefix_lens_cpu,
seq_lens_cpu,
bundle.out_c4_loc,
ratio=4,
)
"""Write newly allocated C128 page ids into the request sidecar."""
_write_per_req_slice(
req_to_token_pool.write_c128,
req_pool_indices_cpu,
@@ -275,36 +160,14 @@ def _write_dsv4_tables(
ratio=128,
)
if bundle.out_c4_state_loc is not None and hasattr(
req_to_token_pool, "write_c4_state"
):
_write_state_tail_per_req(
req_to_token_pool.write_c4_state,
req_pool_indices_cpu,
c4_state_offsets,
seq_lens_cpu,
bundle.out_c4_state_loc,
)
if bundle.out_c128_state_loc is not None and hasattr(
req_to_token_pool, "write_c128_state"
):
_write_state_tail_per_req(
req_to_token_pool.write_c128_state,
req_pool_indices_cpu,
c128_state_offsets,
seq_lens_cpu,
bundle.out_c128_state_loc,
)
def maybe_write_dsv4_decode(
batch: ScheduleBatch,
seq_lens_cpu: torch.Tensor,
token_per_req: int,
) -> None:
"""Post-alloc_decode hook for DSV4. Spreads the new token slot ids
(one per req for swa, gated by ratio boundary for c4/c128) into the
per-req tables on DSV4NPUReqToTokenPool.
"""Post-alloc_decode hook for DSV4. Spreads new C128 KV slot ids into
the per-req sidecar on DSV4NPUReqToTokenPool.
``seq_lens_cpu`` is the POST-decode seq len (already incremented by
``token_per_req``); the new compressed tokens go at positions
@@ -317,28 +180,12 @@ def maybe_write_dsv4_decode(
return
req_to_token_pool = batch.req_to_token_pool
if not hasattr(req_to_token_pool, "write_c4"):
if not hasattr(req_to_token_pool, "write_c128"):
return
prefix_lens_cpu = (seq_lens_cpu - token_per_req).clamp(min=0)
req_pool_indices_cpu = batch.req_pool_indices.cpu()
_write_per_req_slice(
req_to_token_pool.write_swa,
req_pool_indices_cpu,
prefix_lens_cpu,
seq_lens_cpu,
bundle.out_swa_loc,
ratio=1,
)
_write_per_req_slice(
req_to_token_pool.write_c4,
req_pool_indices_cpu,
prefix_lens_cpu,
seq_lens_cpu,
bundle.out_c4_loc,
ratio=4,
)
_write_per_req_slice(
req_to_token_pool.write_c128,
req_pool_indices_cpu,
@@ -348,32 +195,13 @@ def maybe_write_dsv4_decode(
ratio=128,
)
# State table decode writes: one slot per raw decode token (ratio=1).
if bundle.out_c4_state_loc is not None and hasattr(
req_to_token_pool, "write_c4_state"
):
_write_per_req_slice(
req_to_token_pool.write_c4_state,
req_pool_indices_cpu,
prefix_lens_cpu,
seq_lens_cpu,
bundle.out_c4_state_loc,
ratio=1,
)
if bundle.out_c128_state_loc is not None and hasattr(
req_to_token_pool, "write_c128_state"
):
_write_per_req_slice(
req_to_token_pool.write_c128_state,
req_pool_indices_cpu,
prefix_lens_cpu,
seq_lens_cpu,
bundle.out_c128_state_loc,
ratio=1,
)
def maybe_build_dsv4_verify_bundle(batch: ScheduleBatch, draft_token_num: int):
def maybe_build_dsv4_verify_bundle(
batch: ScheduleBatch,
draft_token_num: int,
*,
live_seq_lens_cpu: torch.Tensor | None = None,
):
"""Build the DSV4 cache-location view for one target-verify pass.
Spec-v2 reserves cache ahead of time, so target verify must select only the
@@ -381,31 +209,45 @@ def maybe_build_dsv4_verify_bundle(batch: ScheduleBatch, draft_token_num: int):
the larger allocation bundle produced during decode preparation.
"""
pool = batch.req_to_token_pool
if not hasattr(pool, "req_to_token_c4"):
if not hasattr(pool, "req_to_c128_sidecar"):
return None
reserve_bundle = batch.out_cache_loc_dsv4
if reserve_bundle is None:
return None
req_indices = batch.req_pool_indices_cpu.tolist()
seq_lens = batch.seq_lens_cpu.tolist()
if live_seq_lens_cpu is None:
live_seq_lens_cpu = batch.seq_lens_cpu
if live_seq_lens_cpu is None:
live_seq_lens_cpu = batch.seq_lens[: len(req_indices)].cpu()
live_seq_lens = live_seq_lens_cpu[: len(req_indices)].tolist()
verify_lens = [int(draft_token_num)] * len(req_indices)
def flatten_interval(table: torch.Tensor, ratio: int) -> torch.Tensor:
page_size = pool.c128_page_size
chunks = []
for req_idx, seq_len in zip(req_indices, seq_lens):
start = int(seq_len) // ratio
end = (int(seq_len) + draft_token_num) // ratio
for req_idx, live_seq_len, verify_len in zip(
req_indices, live_seq_lens, verify_lens
):
start = int(live_seq_len) // ratio
end = (int(live_seq_len) + int(verify_len)) // ratio
if end > start:
chunks.append(table[int(req_idx), start:end])
positions = torch.arange(start, end, device=table.device)
pages = table[int(req_idx), positions // page_size]
chunks.append(pages * page_size + positions % page_size)
return torch.cat(chunks) if chunks else table.new_empty((0,))
out_full_loc = batch.out_cache_loc
out_c4_loc = out_full_loc[(out_full_loc >= 0) & ((out_full_loc % 4) == 3)] // 4
return type(reserve_bundle)(
out_full_loc=batch.out_cache_loc,
out_swa_loc=flatten_interval(pool.req_to_token_swa, 1),
out_c4_loc=flatten_interval(pool.req_to_token_c4, 4),
out_c128_loc=flatten_interval(pool.req_to_token_c128, 128),
out_c4_state_loc=flatten_interval(pool.req_to_token_c4_state, 1),
out_c128_state_loc=flatten_interval(pool.req_to_token_c128_state, 1),
out_full_loc=out_full_loc,
out_swa_loc=batch.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
out_full_loc
),
out_c4_loc=out_c4_loc,
out_c128_loc=flatten_interval(pool.req_to_c128_sidecar, 128),
)
@@ -437,23 +279,6 @@ def _write_per_req(
pt += alloc_len
def _write_state_tail_per_req(
write_fn,
req_pool_indices_cpu: torch.Tensor,
state_alloc_offsets: list,
seq_lens_cpu: torch.Tensor,
flat_loc: torch.Tensor,
) -> None:
"""Tail-only state write: req i's slots go at ``[state_alloc_offsets[i],
seq_lens[i])`` in ``req_to_token_c{N}_state``."""
_write_per_req(
write_fn,
req_pool_indices_cpu,
flat_loc,
lambda i: (int(state_alloc_offsets[i]), int(seq_lens_cpu[i].item())),
)
def _write_per_req_slice(
write_fn,
req_pool_indices_cpu: torch.Tensor,
@@ -473,108 +298,3 @@ def _write_per_req_slice(
int(seq_lens_cpu[i].item()) // ratio,
),
)
def maybe_evict_dsv4_state(batch: ScheduleBatch, req: Req, pre_len: int) -> None:
"""Per-decode evict for the DSV4-NPU compress-state pools, independent of
SWA evict cadence. Called every decode step from ``ScheduleBatch``.
The state pool is small (~2 pages c4 / ~3 pages c128 of raw positions per
req) — with a large sliding_window (SWA evict fires every
``eviction_interval`` and needs ``pre_len > sliding_window + page_size`` to
free anything) the pool exhausts before the first SWA frontier advance, so
we drain it here on its own cadence.
Retention windows (kernel read window + decode lookahead margin):
c4 = 8 + 16, c128 = 128 + 64 raw positions — intentionally smaller than one
SWA page so the first eviction fires before the small pool fills. Watermarks
are page-aligned so freed slots are whole pages reclaimable by the paged
allocator. ``req.c{4,128}_state_alloc_offset`` (read/written via getattr/
setattr) is the low-water mark. No-op on non-DSV4-NPU paths.
"""
allocator = batch.token_to_kv_pool_allocator
pool = batch.req_to_token_pool
if not hasattr(allocator, "c4_state_attn_allocator") or (
allocator.c4_state_attn_allocator is None
and allocator.c128_state_attn_allocator is None
):
return
page_size = batch.tree_cache.page_size
c4_watermark = ((max(0, pre_len - (8 + 16))) // page_size) * page_size
c128_watermark = ((max(0, pre_len - (128 + 64))) // page_size) * page_size
_free_state_range(
allocator.c4_state_attn_allocator,
pool,
"req_to_token_c4_state",
req,
"c4_state_alloc_offset",
c4_watermark,
)
_free_state_range(
allocator.c128_state_attn_allocator,
pool,
"req_to_token_c128_state",
req,
"c128_state_alloc_offset",
c128_watermark,
)
def maybe_evict_dsv4_state_on_swa(
allocator, pool, req: Req, new_swa_evicted_seqlen: int
) -> None:
"""Free compress-state slots that ride along with SWA eviction.
State at raw positions < ``swa_evicted_seqlen`` is no longer readable (the
compressor only reads the trailing ``2*ratio`` window) and is returned to
its paged allocator to keep the small state pool from exhausting on long
generations. No-op when the DSV4-NPU state allocators are absent.
This path is needed for small-sliding-window models where
``sliding_window < retention`` (e.g. c128 retention 192 > window 128):
in that case the watermark-based eviction alone may not free slots
fast enough, and the SWA-ride eviction is the primary reclaim mechanism.
For typical large-window models (DS-V4 with window >> 192), the
watermark eviction always runs first, making this path a no-op.
"""
if not hasattr(allocator, "c4_state_attn_allocator"):
return
_free_state_range(
allocator.c4_state_attn_allocator,
pool,
"req_to_token_c4_state",
req,
"c4_state_alloc_offset",
new_swa_evicted_seqlen,
)
_free_state_range(
allocator.c128_state_attn_allocator,
pool,
"req_to_token_c128_state",
req,
"c128_state_alloc_offset",
new_swa_evicted_seqlen,
)
def _free_state_range(
state_allocator,
pool,
table_attr: str,
req: Req,
offset_attr: str,
watermark: int,
) -> None:
"""Free ``[alloc_offset, watermark)`` raw-position state slots for ``req``
and advance its low-water mark. No-op when the allocator/table is absent or
the watermark hasn't advanced past the current offset."""
offset = getattr(req, offset_attr, 0)
if state_allocator is None or not hasattr(pool, table_attr) or watermark <= offset:
return
free_slots = getattr(pool, table_attr)[req.req_pool_idx, offset:watermark]
free_slots = free_slots[free_slots > 0]
if free_slots.numel() > 0:
state_allocator.free(free_slots.to(torch.int64))
setattr(req, offset_attr, watermark)
@@ -1,34 +1,19 @@
"""NPU-only KV pool variant for DeepSeek-V4.
Subclasses :class:`DeepSeekV4TokenToKVPool` to swap the ring-buffered
:class:`CompressStatePool` for the paged :class:`NPUCompressStatePool` that
the on-NPU fused compressor kernel (``torch.ops.custom.compressor`` with
``cache_mode=1``) requires. Atlas A3 rejects ``cache_mode=2`` (ring) entirely,
so this is the only valid layout on that hardware.
The full/SWA/C4/C128 KV buffers keep their Ascend-specific PA_ND layout. The
Compressor state buffers, however, use the same ownership and flat ``state_loc``
rules as the GPU implementation:
Selected at pool construction time by
:meth:`ModelRunnerKVCacheMixin._init_pools` when the model is DSV4 AND the
device is NPU. CUDA continues to use the unchanged base class.
* C4A/C4Li state follows SWA physical pages.
* C128A state follows ``req_pool_idx`` and absolute position.
The subclass overrides only:
* ``_make_attn_state_pool`` / ``_make_indexer_state_pool`` — the per-ratio
state-pool factories the base ``_init_paged_compress_states`` loop calls.
Both return :class:`NPUCompressStatePool` (paged, ``cache_mode=1``)
instead of the base's ring-buffered :class:`CompressStatePool`.
* ``translate_kv_loc_to_compress_state_loc`` — raise loudly. The ring
hash this method implements is meaningless on the paged kernel; callers
must consume ``out_cache_loc_dsv4.out_c{4,128}_state_loc`` from the
allocator bundle instead. Currently the only NPU caller that still
invokes translate is the unfused Python compressor decode path
(``layers/attention/dsv4/compressor.py``); with USE_FUSED_COMPRESSOR=1
that path is dead. If someone disables the fused compressor, they hit
the raise with a clear message.
``NPUCompressStatePool`` only adds the contiguous 3-D view and positive dummy
location required by the Atlas A3 ``cache_mode=2`` operator. There is no paged
state allocator or ``cache_mode=1`` compatibility storage.
"""
from __future__ import annotations
import math
from typing import List, Optional, Tuple
import torch
@@ -42,6 +27,7 @@ from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
DeepSeekV4SingleKVPool,
DeepSeekV4TokenToKVPool,
)
from sglang.srt.runtime_context import get_schedule
class NPUDeepSeekV4SingleKVPool(DeepSeekV4SingleKVPool):
@@ -49,10 +35,9 @@ class NPUDeepSeekV4SingleKVPool(DeepSeekV4SingleKVPool):
``npu_sparse_attn_sharedkv`` reads KV in PA_ND layout
``(num_pages, kernel_page_size, num_kv_heads=1, dim)`` with ``dim`` packing
K_nope + K_rope as bf16, and requires ``cmp_kv.shape[1] == ori_kv.shape[1]``.
So the c4/c128 pools (whose token-level page_size is ``page_size // ratio``)
are allocated at the GLOBAL ``kernel_page_size`` rather than their own
per-ratio page_size; the SWA pool uses ``kernel_page_size == page_size``.
K_nope + K_rope as bf16. C4 uses its native page so its physical page id can
be shared with the corresponding full page. C128 uses its independently
configured physical page size; Full/SWA use the global page size.
The CUDA fp8-packed-bytes layout (the base ``create_buffer``) is untouched.
"""
@@ -68,8 +53,8 @@ class NPUDeepSeekV4SingleKVPool(DeepSeekV4SingleKVPool):
return super().create_buffer(num_pages=num_pages)
kv_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
self.kv_cache_total_dim = kv_dim
# GLOBAL kernel_page_size keeps cmp_kv.shape[1] == ori_kv.shape[1]; writes
# are flat-indexed by loc, so page granularity affects shape not location.
# Writes are flat-indexed by loc; kernel_page_size controls the physical
# page layout exposed to the NPU operators.
npu_num_pages = (self.size + self.kernel_page_size + 1) // self.kernel_page_size
return torch.zeros(
npu_num_pages,
@@ -81,59 +66,17 @@ class NPUDeepSeekV4SingleKVPool(DeepSeekV4SingleKVPool):
)
def npu_state_pool_size(
*,
ratio: int,
page_size: int,
max_num_reqs: int,
) -> int:
"""Per-pool state slot count for the NPU paged state pool's
:class:`NPUPagedTokenToKVPoolAllocator`.
Sizing formula::
max(2, ceil(1.8 * ratio / page_size) + 1) * max_num_reqs * page_size
Sized for steady-state during decode: each req keeps roughly the trailing
``sliding_window_size`` worth of state slots live at any one time (SWA
eviction in :meth:`ScheduleBatch._evict_swa` frees state slots as it
advances), and the 1.8x factor adds headroom for the tail-only allocation
pattern across page boundaries.
Prefill no longer drives sizing because allocation is tail-only — long
prompts only allocate ``c{ratio}_alloc_len`` slots (``≤ tail + 128`` for
c4, ``≤ tail`` for c128, where ``tail = seq_len % 128``), not the full raw
seqlen. See :meth:`ScheduleBatch._compute_dsv4_state_lens_extend` for the
per-req formula.
Result is in TOKEN units (matches the SGLang allocator
``PagedTokenToKVPoolAllocator(size, ...)`` convention where
``num_pages = size // page_size`` is the count of USABLE pages handed out
by ``free_pages = arange(1, num_pages+1)``). The BUFFER allocates one extra
page (see :class:`NPUCompressStatePool`, sized ``(num_pages + 1) *
page_size`` — page 0 is the kernel's skip-sentinel).
"""
blocks_per_req = max(2, math.ceil(1.8 * ratio / page_size) + 1)
num_usable_pages = blocks_per_req * max_num_reqs
return num_usable_pages * page_size
class NPUCompressStatePool(CompressStatePool):
"""Paged compress-state pool for the NPU fused compressor kernel.
"""Thin A3 adapter over the shared GPU-style ring state pool.
``torch.ops.custom.compressor`` (cache_mode=1) reads/writes the compress
state via ``state_cache`` shape ``(block_num, page_size, 2*coff*head_dim)``
indexed by a paged ``state_block_table`` (block ids from 1; value 0 means
"skip this slot"). The CUDA :class:`CompressStatePool` sizes itself
ring-style, which misaddresses slots under cache_mode=1 (ring is also
unsupported on Atlas A3). This subclass keeps the parent's buffer layout
(``(self._size, 2*coff*head_dim)`` flat; ``state_cache_3d`` reshapes to
``(num_blocks, page_size, 2*coff*head_dim)``) but replaces the size formula
with a paged one derived from ``max_num_reqs``. Block 0 is reserved as the
kernel's skip-sentinel (zero kv / -inf score) so any ``state_block_table``
entry defaulting to 0 lands in a deterministic, attention-neutral place.
Allocation, sizing, ring ownership and address translation are inherited
from :class:`CompressStatePool`. NPU only requests a contiguous 3-D view,
enforces the A3 FP32 contract and replaces invalid locations with a cleared
positive dummy row.
NPU-only; CUDA keeps using the unchanged :class:`CompressStatePool`.
Location 0 is valid in explicit mode. Invalid/history-padding locations map
to the final cleared row instead of ``-1`` because the A3 kernel consumes
unsigned offsets.
"""
def __init__(
@@ -146,56 +89,66 @@ class NPUCompressStatePool(CompressStatePool):
device: str,
enable_memory_saver: bool,
ratio: int,
page_size: int,
ring_size: int,
swa_page_size: int,
):
# Bypass parent __init__ — its ring-based sizing is incompatible with the
# kernel's paged block-id contract. We redo buffer alloc and set the same
# fields so the parent API (state_cache_3d, kv_score_buffer) stays intact.
assert ratio in (
4,
128,
), f"NPUCompressStatePool only supports ratio in (4, 128); got {ratio}"
assert page_size > 1, (
"NPUCompressStatePool requires page_size>1 (kernel's "
"state_cache_3d view is (block_num, page_size, slot_dim)). "
"Got page_size=%d." % page_size
assert dtype == torch.float32, (
"Atlas A3 custom.compressor requires FP32 state_cache, "
f"but NPUCompressStatePool got {dtype}."
)
assert ring_size > 0, f"ring_size must be positive, got {ring_size}"
super().__init__(
size=size,
ring_size=ring_size,
overlap=overlap,
head_dim=head_dim,
dtype=dtype,
device=device,
enable_memory_saver=enable_memory_saver,
ratio=ratio,
online=False,
swa_page_size=swa_page_size,
state_cache_page_size=ring_size,
)
self.dummy_state_loc = self._size - 1
# The shared pool initializes its dummy row. A cold C128 request bank
# additionally needs every row initialized before its first partial use.
if ratio == 128:
self.kv_score_buffer.clear()
def _replace_invalid_with_dummy(self, state_loc: torch.Tensor) -> torch.Tensor:
return torch.where(
state_loc < 0,
torch.full_like(state_loc, self.dummy_state_loc),
state_loc,
)
# ``size`` is the ALLOCATOR's size (npu_state_pool_size output). Buffer
# needs one EXTRA page so the free list arange(1, num_pages+1) indexes it
# without OOB (page 0 = skip sentinel; pages 1..num_pages handed out).
num_usable_pages = (size + page_size - 1) // page_size
num_buffer_pages = num_usable_pages + 1
self._size = num_buffer_pages * page_size
self.ratio = ratio
self.page_size = page_size
# ring_size=0 marks "not ring-buffered" (paged allocator replaces the
# parent's ring hashing); kept so downstream hasattr probes don't break.
self.ring_size = 0
# online compress is a CUDA-only opt with no NPU fused-compressor support;
# force off so layout matches kernel expectations.
self.online = False
# Slot dim = 2 * coff * head_dim = [kv | score]; coff = 1 (no overlap) or
# 2 (overlap). Matches CompressStatePool non-online layout.
self.last_dim = 2 * (1 + int(overlap)) * head_dim
# Reuse parent's buffer-alloc helper; only self._size differs from the
# ring-based parent path.
self._alloc_kv_score_buffer(
dtype=dtype, device=device, enable_memory_saver=enable_memory_saver
def translate_from_swa_loc_to_state_loc(
self, swa_loc: torch.Tensor
) -> torch.Tensor:
return self._replace_invalid_with_dummy(
super().translate_from_swa_loc_to_state_loc(swa_loc)
)
# Block 0 = kernel skip-sentinel: kv zeroed, score -inf (softmax → 0).
# The free list excludes it; only stale state_block_table entries land here.
self.kv_score_buffer.kv[:page_size].zero_()
self.kv_score_buffer.score[:page_size].fill_(float("-inf"))
def translate_from_req_position_to_state_loc(
self, req_pool_indices: torch.Tensor, positions: torch.Tensor
) -> torch.Tensor:
return self._replace_invalid_with_dummy(
super().translate_from_req_position_to_state_loc(
req_pool_indices, positions
)
)
class NPUDeepSeekV4IndexerPool(DeepSeekV4IndexerPool):
"""NPU c4-indexer pool. Keeps the base packed CUDA buffer (read by
get_contiguous_buf_infos / NSA) and ADDS dedicated int8 K + float16 scale
buffers in PA_ND layout at the global ``kernel_page_size``, written by
buffers in PA_ND layout at the native C4 ``kernel_page_size``, written by
``torch_npu.npu_scatter_nd_update_`` and read by
``torch.ops.custom.npu_quant_lightning_indexer``.
"""
@@ -270,12 +223,12 @@ class NPUDeepSeekV4IndexerPool(DeepSeekV4IndexerPool):
class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
"""NPU-only DSV4 KV pool with paged compress-state buffers.
"""NPU-only DSV4 KV pool with explicit-location ring state buffers.
The full / SWA / c4 / c128 KV pools use the NPU bf16 PA_ND layout
(:class:`NPUDeepSeekV4SingleKVPool`); the compress-state pool is paged
(:class:`NPUCompressStatePool`) rather than ring-buffered; and the indexer
pool adds dedicated int8 K + fp16 scale buffers
(:class:`NPUDeepSeekV4SingleKVPool`); :class:`NPUCompressStatePool`
exposes the explicit-location ring view required by A3;
and the indexer pool adds dedicated int8 K + fp16 scale buffers
(:class:`NPUDeepSeekV4IndexerPool`). The generic-accessor / port-hook
methods at the bottom of this class are the NPU equivalents of the CUDA
DSV4 store-cache chain — kept here, not in the community base, which raises
@@ -283,6 +236,16 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
accessors instead).
"""
def __init__(self, *args, **kwargs):
c128_page_size = get_schedule().c128_page_size
if c128_page_size <= 0 or c128_page_size % 16 != 0:
raise ValueError(
"c128_page_size must be a positive multiple of 16 for the NPU "
f"sparse-attention operator, got {c128_page_size}"
)
self.c128_page_size = c128_page_size
super().__init__(*args, **kwargs)
def _make_kv_pool(
self,
*,
@@ -301,6 +264,16 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
"enable_hisparse is not supported on the NPU DSV4 KV pool "
f"(got c4 pool class {cls.__name__})."
)
# Full/SWA use the global page size, C4 uses its native compressed page,
# and C128 has an independent physical page size.
is_c4_pool = page_size * 4 == global_page_size
is_c128_pool = page_size * 128 == global_page_size
if is_c4_pool:
kernel_page_size = page_size
elif is_c128_pool:
kernel_page_size = self.c128_page_size
else:
kernel_page_size = global_page_size
return NPUDeepSeekV4SingleKVPool(
size,
page_size,
@@ -310,7 +283,7 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
layer_num,
device,
enable_memory_saver,
kernel_page_size=global_page_size,
kernel_page_size=kernel_page_size,
)
def _get_state_pool(self, layer_id: int, from_indexer: bool) -> CompressStatePool:
@@ -332,13 +305,14 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
)
return NPUCompressStatePool(
size=self._state_pool_size(ratio),
ring_size=self.get_ring_size(ratio),
overlap=ratio == 4,
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim,
dtype=self.c4_state_dtype if ratio == 4 else self.c128_state_dtype,
device=self.device,
enable_memory_saver=enable_memory_saver,
ratio=ratio,
page_size=self.swa_page_size,
swa_page_size=self.swa_page_size,
)
def _make_indexer_state_pool(
@@ -348,24 +322,16 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
# slot_dim (indexer_head_dim vs attention head_dim).
return NPUCompressStatePool(
size=self.c4_state_pool_size,
ring_size=self.get_ring_size(ratio),
overlap=ratio == 4,
head_dim=self.indexer_head_dim,
device=self.device,
dtype=self.c4_state_dtype,
enable_memory_saver=enable_memory_saver,
ratio=ratio,
page_size=self.swa_page_size,
swa_page_size=self.swa_page_size,
)
def clear_unaccepted_c128_draft_states(
self,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
accept_lens: torch.Tensor,
num_draft_tokens: int,
) -> None:
pass
def _make_indexer_pool(
self,
size: int,
@@ -376,8 +342,7 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
device: str,
enable_memory_saver: bool,
) -> NPUDeepSeekV4IndexerPool:
# NPU dedicated int8 K + fp16 scale buffers use the GLOBAL page_size
# (= self.page_size) as kernel_page_size, matching ori_kv for the kernel.
# Indexer shares C4 addresses and therefore uses the same native page.
return NPUDeepSeekV4IndexerPool(
size,
page_size,
@@ -386,86 +351,58 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
layer_num,
device,
enable_memory_saver,
kernel_page_size=self.page_size,
kernel_page_size=page_size,
)
def get_contiguous_buf_infos(self) -> Tuple[List[int], List[int], List[int]]:
# No full-token contiguous space on NPU; everything ships per-pool via
# get_pd_state_components(), so the contiguous path is empty.
return [], [], []
def get_pd_state_components(
self,
) -> List[Tuple[str, List[int], List[int], List[int]]]:
"""Ordered ``(AscendStateType, data_ptrs, data_lens, item_lens)`` per pool, in a
fixed order so prefill and decode register identically (empty pools skipped)."""
from sglang.srt.disaggregation.ascend.conn import AscendStateType
components: List[Tuple[str, List[int], List[int], List[int]]] = []
def kv_entry(bufs):
return (
[b.data_ptr() for b in bufs],
[b.nbytes for b in bufs],
[b[0].nbytes for b in bufs],
)
def state_entry(want_ratio: int, include_indexer: bool):
ptrs: List[int] = []
lens: List[int] = []
ilens: List[int] = []
def add(pool):
t = pool.kv_score_buffer.kv_score
ptrs.append(t.data_ptr())
lens.append(t.nbytes)
ilens.append(t[0].nbytes * pool.page_size)
for ratio, pool in zip(self.compression_ratios, self.compress_state_pools):
if pool is not None and ratio == want_ratio:
add(pool)
if include_indexer:
# indexer compress-state pools are all ratio 4 and share the
# c4_state slot space.
for pool in self.indexer_compress_state_pools:
if pool is not None:
add(pool)
return ptrs, lens, ilens
# KV pools (4D PA_ND).
if self.swa_kv_pool is not None:
components.append(
(AscendStateType.DSV4_SWA, *kv_entry(self.swa_kv_pool.kv_buffer))
)
if self.c4_kv_pool is not None:
components.append(
(AscendStateType.DSV4_C4, *kv_entry(self.c4_kv_pool.kv_buffer))
)
if self.c128_kv_pool is not None:
components.append(
(AscendStateType.DSV4_C128, *kv_entry(self.c128_kv_pool.kv_buffer))
)
if self.c4_indexer_kv_pool is not None:
idx_bufs = list(self.c4_indexer_kv_pool.index_k_buffer) + list(
self.c4_indexer_kv_pool.index_scale_buffer
)
components.append((AscendStateType.DSV4_INDEXER, *kv_entry(idx_bufs)))
# Compress-state pools (paged, flat 2D). c4_state bundles attn-c4-state +
# indexer-c4-state (same req_to_token_c4_state slot space).
components.append(
(AscendStateType.DSV4_C4_STATE, *state_entry(4, include_indexer=True))
"""Main PD buffers addressed by the full KV page id."""
buffers = (
self.c4_kv_pool.kv_buffer
+ self.c4_indexer_kv_pool.index_k_buffer
+ self.c4_indexer_kv_pool.index_scale_buffer
)
components.append(
(AscendStateType.DSV4_C128_STATE, *state_entry(128, include_indexer=False))
return (
[buf.data_ptr() for buf in buffers],
[buf.nbytes for buf in buffers],
[buf[0].nbytes for buf in buffers],
)
# Drop empty components (e.g. a ratio with no layers) so every shipped
# component has non-zero item_lens; the set is identical on both sides.
return [c for c in components if c[1]]
def get_state_buf_infos(self) -> Tuple[List[int], List[int], List[int]]:
"""GPU-compatible ``StateType.SWA`` component.
SWA KV, C4 attention state and C4 indexer state retain separate buffers
but share the same SWA page/state index.
"""
data_ptrs: List[int] = []
data_lens: List[int] = []
item_lens: List[int] = []
for buf in self.swa_kv_pool.kv_buffer:
data_ptrs.append(buf.data_ptr())
data_lens.append(buf.nbytes)
item_lens.append(buf[0].nbytes)
for pools in (self.compress_state_pools, self.indexer_compress_state_pools):
for pool in pools:
if pool is None or pool.ratio != 4:
continue
state = pool.kv_score_buffer.kv_score
data_ptrs.append(state.data_ptr())
data_lens.append(state.nbytes)
item_lens.append(state[0].nbytes * pool.ring_size)
return data_ptrs, data_lens, item_lens
def get_c128_kv_buf_infos(self) -> Tuple[List[int], List[int], List[int]]:
buffers = self.c128_kv_pool.kv_buffer
return (
[buf.data_ptr() for buf in buffers],
[buf.nbytes for buf in buffers],
[buf[0].nbytes for buf in buffers],
)
def get_state_cache(self, layer_id: int, from_indexer: bool) -> torch.Tensor:
"""fp32 ``[block_num, page_size, 2*coff*D]`` view of this layer's
"""FP32 ``[block_num, ring_size, 2*coff*D]`` view of this layer's
kv+score buffer — the fused compressor op
(``torch.ops.custom.compressor``)'s ``state_cache`` argument."""
return self._get_state_pool(layer_id, from_indexer).state_cache_3d
@@ -559,42 +496,43 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
cache = cache.unsqueeze(1)
buf_flat[loc] = cache.to(buf_flat.dtype)
# ------------------------------------------------------------------
# NPU port hooks — used by dsv4/{compressor,indexer}.py forward_npu.
# CompressStatePool stores a fused [kv | score] tensor; split is a last-dim slice.
# ------------------------------------------------------------------
def set_state_buffer(
def set_swa_key_buffer_radix_fused_norm_rope(
self,
layer_id: int,
loc: torch.Tensor,
swa_loc: torch.Tensor,
kv: torch.Tensor,
score: torch.Tensor,
from_indexer: bool,
kv_weight: torch.Tensor,
eps: float,
freqs_cis: torch.Tensor,
positions: torch.Tensor,
) -> None:
# KVAndScore.kv_score is [..., 2*coff*head_dim] = [kv | score].
kv_score = self._get_state_pool(layer_id, from_indexer).kv_score_buffer.kv_score
last_dim = kv_score.shape[-1]
half = last_dim // 2
kv_view = kv.reshape(-1, half).to(kv_score.dtype)
score_view = score.reshape(-1, half).to(kv_score.dtype)
kv_score[loc, :half] = kv_view
kv_score[loc, half:] = score_view
kv_out = torch_npu.npu_rms_norm(kv, kv_weight, eps)[0]
def get_state_buffer(
self,
layer_id: int,
from_indexer: bool,
kv_indices: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
kv_score = self._get_state_pool(layer_id, from_indexer).kv_score_buffer.kv_score
if kv_indices is not None:
kv_score = kv_score[kv_indices]
last_dim = kv_score.shape[-1]
half = last_dim // 2
kv = kv_score[..., :half].unsqueeze(-2) # add num_kv_heads=1 axis
score = kv_score[..., half:].unsqueeze(-2)
return kv, score
rope_dim = freqs_cis.shape[-1] * 2
from sglang.srt.hardware_backend.npu.dsv4.dsv4_rope import Dsv4NpuRoPE
cos, sin = Dsv4NpuRoPE.for_freqs(freqs_cis).get_cos_sin(
positions,
kv_out.dtype,
view_4d=True,
allow_build=True,
cache_dtype=torch.float32,
)
Dsv4NpuRoPE.apply_rotary_mul_inplace(
kv_out.reshape(kv_out.shape[0], -1, kv_out.shape[-1]),
None,
cos,
sin,
qk_nope_dim=kv_out.shape[-1] - rope_dim,
)
safe_swa_loc = swa_loc.clamp_min(0).to(torch.int64)
self.set_swa_buffer(
layer_id,
safe_swa_loc,
kv_out,
)
def set_compress_buffer(
self,
@@ -648,22 +586,3 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
assert from_indexer, "only indexer compress pool has dequant scale"
compress_layer_id = self.layer_mapping[layer_id].compress_layer_id
return self.c4_indexer_kv_pool.get_index_scale(compress_layer_id)
def translate_kv_loc_to_compress_state_loc(
self,
kv_loc: torch.Tensor,
compress_ratio: int,
) -> torch.Tensor:
# Parent's ring-buffer hash is meaningless under the paged cache_mode=1
# contract; returning a stale value would silently corrupt state. Fail loud.
raise RuntimeError(
"DSV4NPUTokenToKVPool.translate_kv_loc_to_compress_state_loc was "
"called, but the NPU fused compressor kernel uses a paged state "
"pool (cache_mode=1) and does not support ring-buffer state "
"addressing (cache_mode=2 is explicitly unsupported on Atlas A3). "
"Callers must consume out_cache_loc_dsv4.out_c{4,128}_state_loc "
"from the allocator bundle (set during alloc_extend/alloc_decode) "
"and read state_page_table from req_to_token_c{4,128}_state on "
"the DSV4NPUReqToTokenPool instead. See "
"hardware_backend/npu/dsv4_memory_pool.py for the rationale."
)
@@ -1,33 +1,13 @@
"""DSV4-NPU per-request mapping pool.
Subclass of ``ReqToTokenPool`` that adds five auxiliary per-request tables
Subclass of ``ReqToTokenPool`` that adds the one auxiliary per-request table
needed by the DSV4 attention backend:
* ``req_to_token_swa`` — slot ids in the SWA full-pool view
* ``req_to_token_c4`` — slot ids in the c4 compressed-KV pool
* ``req_to_token_c128`` — slot ids in the c128 compressed-KV pool
* ``req_to_token_c4_state`` — c4 state-pool slot ids, 1 per raw token
* ``req_to_token_c128_state`` — c128 state-pool slot ids, 1 per raw token
* ``req_to_c128_sidecar`` — one page id per C128 physical page
Compressed KV pools store 1 slot per ``ratio`` raw tokens, so their per-req
table column count is ``max_context_len // ratio``. swa mirrors the raw
token count. Elements are token-level slot ids; the attention backend
converts to page ids via ``// page_size`` when constructing PA_ND block
tables.
The c4/c128 STATE pools also have per-req tables here: the NPU fused
compressor uses a paged state pool (``cache_mode=1``), so each raw token's
state slot id is recorded (1 column per raw token) and the backend builds
``state_block_table = req_to_token_c{N}_state[req, ::page_size] // page_size``
to feed the kernel. (The base class' ``translate_kv_loc_to_compress_state_loc``
ring-hash is the CUDA-only path; it is disabled on NPU.)
Memory cost example (size=64, max_context_len=32K): swa 8MB + c4 2MB +
c128 64KB ≈ 10MB extra on top of the base req_to_token (8MB).
The tables are populated by the ``dsv4_common_hooks`` writers (driven from
``mem_cache/common.py``) immediately after a successful alloc_extend /
alloc_decode, using the per-pool slot indices returned in ``DSV4OutCacheLoc``.
C4 locations are derived from the base full-token table, and SWA locations use
the existing full-to-SWA mapping. The sidecar is populated by the existing
``dsv4_common_hooks`` flow from the slot indices in ``DSV4OutCacheLoc``.
"""
from __future__ import annotations
@@ -37,6 +17,7 @@ import torch
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
from sglang.srt.disaggregation.decode import DecodeReqToTokenPool
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.runtime_context import get_schedule
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
@@ -51,7 +32,11 @@ class DSV4ReqToTokenTablesMixin:
"""
def _init_dsv4_tables(
self, max_context_len: int, device: str, enable_memory_saver: bool
self,
max_context_len: int,
device: str,
enable_memory_saver: bool,
c128_page_size: int,
) -> None:
memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=enable_memory_saver
@@ -59,70 +44,83 @@ class DSV4ReqToTokenTablesMixin:
# Back-ref to DSV4NPUTokenToKVPoolAllocator, wired via
# register_dsv4_allocator after both exist, so free(req) can release
# c4/c128 pages. None at construction so base clear() runs safely.
# c128 pages. None at construction so base clear() runs safely.
self._dsv4_allocator = None
self.c128_page_size = c128_page_size
# (name, columns). swa + state tables: 1 slot per raw token; c4/c128:
# 1 slot per `ratio` raw tokens. Init zero so unallocated columns map to
# block 0 (kernel skip sentinel cleared by NPUCompressStatePool).
group_tokens = 128 * c128_page_size
with memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
for name, cols in (
("req_to_token_swa", max_context_len),
("req_to_token_c4", max(1, max_context_len // 4)),
("req_to_token_c128", max(1, max_context_len // 128)),
("req_to_token_c4_state", max_context_len),
("req_to_token_c128_state", max_context_len),
):
setattr(
self,
name,
torch.zeros(
(self._alloc_size, cols),
dtype=torch.int32,
device=device,
),
)
# Per-pool write helpers, called by mem_cache/common.py after alloc, using
# slot indices from DSV4OutCacheLoc. Args: (req_pool_idx, token_offset), slot.
def write_swa(self, indices, values: torch.Tensor) -> None:
self.req_to_token_swa[indices] = values
def write_c4(self, indices, values: torch.Tensor) -> None:
self.req_to_token_c4[indices] = values
self.req_to_c128_sidecar = torch.zeros(
(
self._alloc_size,
max(1, (max_context_len + group_tokens - 1) // group_tokens),
),
dtype=torch.int32,
device=device,
)
def write_c128(self, indices, values: torch.Tensor) -> None:
self.req_to_token_c128[indices] = values
def write_c4_state(self, indices, values: torch.Tensor) -> None:
self.req_to_token_c4_state[indices] = values
def write_c128_state(self, indices, values: torch.Tensor) -> None:
self.req_to_token_c128_state[indices] = values
req_pool_idx, token_slice = indices
page_size = self.c128_page_size
first_group = (token_slice.start + page_size - 1) // page_size
end_group = (token_slice.stop + page_size - 1) // page_size
if first_group == end_group:
return
groups = torch.arange(first_group, end_group, device=values.device)
pages = values[groups * page_size - token_slice.start] // page_size
prefix_pages = self.req_to_c128_sidecar[req_pool_idx, :end_group].clone()
prefix_pages[groups] = pages
self._dsv4_allocator.replace_req_c128_prefix(req_pool_idx, prefix_pages, self)
def register_dsv4_allocator(self, allocator) -> None:
"""Wire the DSV4NPUTokenToKVPoolAllocator ref so ``free(req)`` can
release c4/c128 pool pages alongside the req_pool_idx slot."""
release C128 KV pages."""
self._dsv4_allocator = allocator
def set_c128_prefix_pages(self, req, page_ids: torch.Tensor) -> None:
"""Install pages returned by a Radix match.
Prefix matching can happen before a request slot is allocated, so the
page ids are temporarily carried by ``Req`` and installed by ``alloc``.
"""
if req.req_pool_idx is None:
req.c128_prefix_page_ids = page_ids
return
self._dsv4_allocator.replace_req_c128_prefix(
int(req.req_pool_idx), page_ids, self
)
def alloc(self, reqs):
fresh = [req.req_pool_idx is None for req in reqs]
indices = super().alloc(reqs)
if indices is None:
return None
for is_fresh, req, req_pool_idx in zip(fresh, reqs, indices):
if is_fresh:
self.req_to_c128_sidecar[int(req_pool_idx)].zero_()
pages = getattr(req, "c128_prefix_page_ids", None)
if pages is not None:
self._dsv4_allocator.replace_req_c128_prefix(
int(req_pool_idx), pages, self
)
req.c128_prefix_page_ids = None
return indices
def _dsv4_free(self, req) -> None:
# Trigger c4/c128 free via the allocator's unified free path. May be None
# Trigger C128 KV free/state clear via the allocator's unified path. May be None
# between __init__ and register_dsv4_allocator — defensive None check.
if self._dsv4_allocator is not None:
self._dsv4_allocator.free(req=req, req_to_token_pool=self)
class DSV4NPUReqToTokenPool(DSV4ReqToTokenTablesMixin, ReqToTokenPool):
"""ReqToTokenPool extended with DSV4 SWA + c4/c128 per-req tables.
"""ReqToTokenPool extended with the DSV4 C128 group sidecar mapping.
Drop-in replacement for ReqToTokenPool when the model is DeepSeek-V4 on
NPU. Selected by ``model_runner_kv_cache_mixin`` based on model arch +
device. Non-DSV4 and non-NPU paths continue to use the base class.
The auxiliary tables are intentionally NOT zeroed on ``clear()``: they are
indexed only by active rows (via req_pool_idx) and only each row's
``[:seq_len]`` prefix is read, so stale entries past kv_committed_len are
unreachable by the attention metadata builder.
Each freshly allocated request row is cleared before use.
"""
def __init__(
@@ -133,7 +131,12 @@ class DSV4NPUReqToTokenPool(DSV4ReqToTokenTablesMixin, ReqToTokenPool):
enable_memory_saver: bool,
):
super().__init__(size, max_context_len, device, enable_memory_saver)
self._init_dsv4_tables(max_context_len, device, enable_memory_saver)
self._init_dsv4_tables(
max_context_len,
device,
enable_memory_saver,
get_schedule().c128_page_size,
)
def free(self, req):
self._dsv4_free(req)
@@ -141,7 +144,7 @@ class DSV4NPUReqToTokenPool(DSV4ReqToTokenTablesMixin, ReqToTokenPool):
class DSV4NPUDecodeReqToTokenPool(DSV4ReqToTokenTablesMixin, DecodeReqToTokenPool):
"""DecodeReqToTokenPool with the DSV4 swa/c4/c128(+state) per-req tables.
"""DecodeReqToTokenPool with the C128 group sidecar mapping.
The disagg-decode counterpart of DSV4NPUReqToTokenPool; DecodeReqToTokenPool
pre-allocates extra req slots for in-flight prefill transfers.
@@ -162,7 +165,12 @@ class DSV4NPUDecodeReqToTokenPool(DSV4ReqToTokenTablesMixin, DecodeReqToTokenPoo
enable_memory_saver=enable_memory_saver,
pre_alloc_size=pre_alloc_size,
)
self._init_dsv4_tables(max_context_len, device, enable_memory_saver)
self._init_dsv4_tables(
max_context_len,
device,
enable_memory_saver,
get_schedule().c128_page_size,
)
def free(self, req):
self._dsv4_free(req)
@@ -0,0 +1,122 @@
from __future__ import annotations
import logging
import os
import re
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Optional
import torch
logger = logging.getLogger(__name__)
@dataclass
class OpLibSpec:
"""Configuration for a standalone operator library loaded into ``torch.ops``."""
name: str # human-readable id, used in logs/errors
so_env: str # env var that points to the standalone .so path
namespace: str # torch.ops.<namespace> the operators register into
required_ops: tuple[str, ...]
pre_load_imports: tuple[str, ...] = () # modules to import before loading
class TorchOpLoader:
"""Load a standalone .so into ``torch.ops`` and validate its operators."""
def __init__(self, spec: OpLibSpec) -> None:
self._spec = spec
self._loaded_library: Optional[Path] = None
def _missing_ops(self) -> list[str]:
namespace = getattr(torch.ops, self._spec.namespace, None)
if namespace is None:
return list(self._spec.required_ops)
return [op for op in self._spec.required_ops if not hasattr(namespace, op)]
def registered(self) -> bool:
"""Return whether the required operators are already registered."""
return not self._missing_ops()
def _resolve_so_path(self) -> Path:
explicit = os.environ.get(self._spec.so_env)
if not explicit:
raise RuntimeError(
f"The {self._spec.name} operators are not registered. Set "
f"{self._spec.so_env} to the standalone .so library path."
)
path = Path(explicit).expanduser().resolve()
if not path.is_file():
raise RuntimeError(f"{self._spec.so_env} points to a missing file: {path}")
return path
def _validate_python_abi(self, library_path: Path) -> None:
abi_match = re.search(r"\.cpython-(\d+)-", library_path.name)
current_abi = f"{sys.version_info.major}{sys.version_info.minor}"
if abi_match is not None and abi_match.group(1) != current_abi:
raise RuntimeError(
f"{library_path} was built for CPython {abi_match.group(1)}, "
f"but SGLang is running CPython {current_abi}. Rebuild the "
"extension with the SGLang Python/Torch/torch-npu environment."
)
def initialize(self) -> Optional[Path]:
"""Register the operators before backend execution.
Idempotent: returns ``None`` if the operators are already registered
(e.g. by another package). Otherwise loads the standalone .so pointed
to by ``so_env`` into ``torch.ops`` and validates the required operators.
Returns the loaded library path when this call loaded it, else ``None``.
"""
if self.registered():
return None
if self._loaded_library is not None:
missing = self._missing_ops()
raise RuntimeError(
f"Loaded {self._loaded_library}, but required "
f"{self._spec.namespace} operators are missing: {missing}."
)
for module in self._spec.pre_load_imports:
__import__(module) # noqa: F401 side-effect imports (e.g. torch_npu)
library_path = self._resolve_so_path()
self._validate_python_abi(library_path)
try:
torch.ops.load_library(str(library_path))
except Exception as exc:
raise RuntimeError(
f"Failed to load the {self._spec.name} operator library "
f"{library_path}. Ensure its dependent CANN/custom-op libraries "
"are visible through LD_LIBRARY_PATH and the Ascend OPP setup."
) from exc
missing = self._missing_ops()
if missing:
raise RuntimeError(
f"Loaded {library_path}, but required "
f"{self._spec.namespace} operators are missing: {missing}."
)
self._loaded_library = library_path
logger.info("Registered %s operators from %s", self._spec.name, library_path)
return library_path
def initialize_dspark_sparse_attn_ops() -> Optional[Path]:
"""Register the DSpark sparse-attention ops before backend execution."""
spec = OpLibSpec(
name="DSpark sparse-attention",
so_env="SGLANG_DSPARK_EXTRA_OPS_SO",
namespace="_C_ascend",
required_ops=(
"npu_sparse_attn_sharedkv_metadata",
"npu_sparse_attn_sharedkv",
),
pre_load_imports=("torch_npu",),
)
return TorchOpLoader(spec).initialize()
@@ -138,6 +138,66 @@ class ModelSlimConfig(QuantizationConfig):
"forward_npu",
[npu_wrapper_rmsnorm_forward],
)
# DSpark checkpoint weights use mtp.<stage>.*, while the runtime draft
# model constructs canonical modules under stages.<stage>.*. Keep
# this transformation in sync with
# DeepseekV4ForCausalLMDSpark._remap_dspark_weight_name. Merely
# replacing ``mtp`` with ``stages`` is insufficient: it silently
# misses ModelSlim lookups such as stages.0.mlp.experts and
# stages.0.self_attn.
dspark_quant_aliases = {}
for name, scheme in quant_config.items():
if not isinstance(name, str) or not name.startswith("mtp."):
continue
parts = name.split(".", 2)
if len(parts) != 3:
continue
stage_id, rest = parts[1], parts[2]
if not stage_id.isdigit():
continue
# The draft model attaches the target model's shared embedding and
# LM head; mtp-local copies are not runtime draft modules.
if rest.startswith(("embed.", "embed_tokens.", "head.", "lm_head.")):
continue
if rest.startswith("markov_head."):
alias = f"markov_head.{rest[len('markov_head.'):]}"
elif rest.startswith("confidence_head."):
alias = f"confidence_head.{rest[len('confidence_head.'):]}"
else:
mapped_rest = rest
if mapped_rest.startswith("attn."):
mapped_rest = "self_attn." + mapped_rest.removeprefix("attn.")
elif mapped_rest.startswith("ffn."):
mapped_rest = "mlp." + mapped_rest.removeprefix("ffn.")
elif mapped_rest.startswith("attn_norm."):
mapped_rest = "input_layernorm." + mapped_rest.removeprefix(
"attn_norm."
)
elif mapped_rest.startswith("ffn_norm."):
mapped_rest = (
"post_attention_layernorm."
+ mapped_rest.removeprefix("ffn_norm.")
)
mapped_rest = mapped_rest.replace(".w1.", ".gate_proj.")
mapped_rest = mapped_rest.replace(".w2.", ".down_proj.")
mapped_rest = mapped_rest.replace(".w3.", ".up_proj.")
mapped_rest = mapped_rest.replace(".gate.tid2eid", ".topk.tid2eid")
mapped_rest = mapped_rest.replace(
".gate.bias", ".gate.e_score_correction_bias"
)
alias = f"stages.{stage_id}.{mapped_rest}"
dspark_quant_aliases[alias] = scheme
quant_config = {
**dspark_quant_aliases,
**quant_config,
}
self.quant_description = quant_config
def update_packed_modules_mapping(self, mapping: Dict[str, List[str]]) -> None:
self.packed_modules_mapping.update(mapping)
+3 -9
View File
@@ -81,9 +81,6 @@ from sglang.srt.disaggregation.decode_schedule_batch_mixin import (
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode
from sglang.srt.dllm.mixin.req import ReqDllmMixin
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
maybe_evict_dsv4_state,
)
from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import (
NewTokenRatioTracker,
@@ -1164,6 +1161,7 @@ class Req(ReqDllmMixin):
# kv_send(req.input_ids[req.start_send_idx:req.extend_range.end])
# start_send_idx = req.extend_range.end
self.start_send_idx: int = 0
self.disagg_decode_prefix_len: int = 0
# For overlap schedule, we delay the kv transfer until `process_batch_result_disagg_prefill` rather than `process_prefill_chunk` in non-overlap
# This is because kv is not ready in `process_prefill_chunk`.
@@ -2073,8 +2071,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# The output locations of the KV cache
out_cache_loc: torch.Tensor = None # shape: [b], int64
# DSV4-NPU: per-pool slot bundle from DSV4NPUTokenToKVPoolAllocator (None
# elsewhere); c4/c128 state lens ride on ``batch.dsv4_state_lens``.
# DSV4-NPU: KV-only per-pool slot bundle from
# DSV4NPUTokenToKVPoolAllocator (None elsewhere).
out_cache_loc_dsv4: Optional[Any] = None
# For hybrid GDN prefix cache
@@ -3330,10 +3328,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
):
self._evict_swa(req, req.seqlen - 1)
# DSV4-NPU only (no-op elsewhere): the small paged compress-state
# pool must drain every decode step, independent of SWA cadence.
maybe_evict_dsv4_state(self, req, req.seqlen - 1)
# Once the decode position has moved past the sliding window,
# the SWA portion of the prefill-time tree lock is no longer
# needed by this request. Convert it from protected to
+5 -39
View File
@@ -47,7 +47,6 @@ if _is_cpu:
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.model_executor.forward_batch_info import DSV4StateLens
logger = logging.getLogger(__name__)
@@ -172,30 +171,6 @@ def alloc_token_slots(
return out_cache_loc
def _compute_dsv4_state_lens(batch, *, is_decode: bool):
"""Per-req c{4,128}_state pool alloc lens (``DSV4StateLens``) for this step.
None on CUDA / non-V4 paths (allocator has no ``compute_dsv4_state_lens_*``).
"""
allocator = batch.token_to_kv_pool_allocator
if not hasattr(allocator, "compute_dsv4_state_lens_extend"):
return None
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
maybe_evict_dsv4_state,
)
if is_decode:
for req in batch.reqs:
maybe_evict_dsv4_state(batch, req, req.seqlen - 1)
return allocator.compute_dsv4_state_lens_decode(batch.reqs)
prefix_lens = batch.prefix_lens
for req, prefix_len in zip(batch.reqs, prefix_lens):
if prefix_len > 0:
maybe_evict_dsv4_state(batch, req, prefix_len)
return allocator.compute_dsv4_state_lens_extend(
batch.reqs, batch.seq_lens_cpu.tolist(), prefix_lens
)
def alloc_paged_token_slots_extend(
tree_cache: BasePrefixCache,
prefix_lens: torch.Tensor,
@@ -205,7 +180,6 @@ def alloc_paged_token_slots_extend(
last_loc: torch.Tensor,
extend_num_tokens: int,
req_pool_indices: Optional[torch.Tensor] = None,
dsv4_state_lens: Optional[DSV4StateLens] = None,
batch=None,
):
# Over estimate the number of tokens: assume each request needs a new page.
@@ -213,15 +187,13 @@ def alloc_paged_token_slots_extend(
num_tokens = extend_num_tokens + len(seq_lens_cpu) * allocator.page_size
evict_from_tree_cache(tree_cache, num_tokens)
is_dsv4 = req_pool_indices is not None and hasattr(allocator, "c4_attn_allocator")
is_dsv4 = req_pool_indices is not None and hasattr(allocator, "c128_attn_allocator")
extra_alloc_kwargs = {}
if is_dsv4:
extra_alloc_kwargs["req_pool_indices"] = req_pool_indices
# Per-call per-req tables for the c-pool / state last_loc lookup.
# Per-call per-req table for the C128 KV last_loc lookup.
if batch is not None:
extra_alloc_kwargs["req_to_token_pool"] = batch.req_to_token_pool
if dsv4_state_lens is not None:
extra_alloc_kwargs["dsv4_state_lens"] = dsv4_state_lens
out = allocator.alloc_extend(
prefix_lens,
@@ -370,7 +342,6 @@ def alloc_for_extend(
last_loc=torch.cat(last_loc),
extend_num_tokens=batch.extend_num_tokens,
req_pool_indices=req_pool_indices_device,
dsv4_state_lens=_compute_dsv4_state_lens(batch, is_decode=False),
batch=batch,
)
@@ -466,7 +437,6 @@ def _alloc_extend_loc_with_kv_reuse(
last_loc=torch.cat(last_loc),
extend_num_tokens=alloc_extend_num_tokens,
req_pool_indices=req_pool_indices_device,
dsv4_state_lens=_compute_dsv4_state_lens(batch, is_decode=False),
batch=batch,
)
@@ -497,7 +467,6 @@ def alloc_paged_token_slots_decode(
last_loc: torch.Tensor,
token_per_req: int = 1,
req_pool_indices: Optional[torch.Tensor] = None,
dsv4_state_lens: Optional[DSV4StateLens] = None,
batch=None,
) -> torch.Tensor:
"""Allocate paged KV cache for decode batch."""
@@ -506,17 +475,15 @@ def alloc_paged_token_slots_decode(
num_tokens = len(seq_lens) * allocator.page_size
evict_from_tree_cache(tree_cache, num_tokens)
# DSV4-NPU allocator also needs req_pool_indices + per-req state lens and
# DSV4-NPU allocator also needs req_pool_indices for C128 KV allocation and
# returns a DSV4OutCacheLoc bundle; hasattr-gated so others stay unchanged.
is_dsv4 = req_pool_indices is not None and hasattr(allocator, "c4_attn_allocator")
is_dsv4 = req_pool_indices is not None and hasattr(allocator, "c128_attn_allocator")
extra_alloc_kwargs = {}
if is_dsv4:
extra_alloc_kwargs["req_pool_indices"] = req_pool_indices
# Per-call per-req tables for the last_loc lookup.
# Per-call per-req C128 table for the last_loc lookup.
if batch is not None:
extra_alloc_kwargs["req_to_token_pool"] = batch.req_to_token_pool
if dsv4_state_lens is not None:
extra_alloc_kwargs["dsv4_state_lens"] = dsv4_state_lens
out = allocator.alloc_decode(seq_lens, seq_lens_cpu, last_loc, **extra_alloc_kwargs)
@@ -571,7 +538,6 @@ def alloc_for_decode(batch: ScheduleBatch, token_per_req: int) -> torch.Tensor:
last_loc=last_loc,
token_per_req=token_per_req,
req_pool_indices=batch.req_pool_indices,
dsv4_state_lens=_compute_dsv4_state_lens(batch, is_decode=True),
batch=batch,
)
@@ -67,6 +67,9 @@ class InsertParams:
# Mamba specific
mamba_value: Optional[torch.Tensor] = None
# DSV4 NPU C128 sidecar pages, one page id per physical C128 page group.
c128_value: Optional[torch.Tensor] = None
# SWA specific
prev_prefix_len: int = 0
swa_evicted_seqlen: int = 0
-6
View File
@@ -10,9 +10,6 @@ from sglang.kernels.ops.memory.common import (
_get_last_loc_safe_kernel as _get_last_loc_safe_kernel,
)
from sglang.kernels.ops.memory.common import get_last_loc_kernel as get_last_loc_kernel
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
maybe_evict_dsv4_state_on_swa,
)
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
@@ -96,9 +93,6 @@ def free_swa_out_of_window_slots(
req.req_pool_idx, req.kv.swa_evicted_seqlen : new_swa_evicted_seqlen
]
token_to_kv_pool_allocator.free_swa(free_slots)
maybe_evict_dsv4_state_on_swa(
token_to_kv_pool_allocator, req_to_token_pool, req, new_swa_evicted_seqlen
)
req.kv.swa_evicted_seqlen = new_swa_evicted_seqlen
@@ -8,11 +8,10 @@ import torch
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
from sglang.srt.mem_cache.utils import maybe_init_custom_mem_pool
from sglang.srt.utils import is_hip, is_npu
from sglang.srt.utils import is_hip
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
_is_hip = is_hip()
_is_npu = is_npu()
def _lcm(a: int, b: int) -> int:
@@ -95,11 +94,14 @@ class CompressStatePool:
online: bool = False,
swa_page_size: int = 0,
online_mtp_max_draft_tokens: int = 0,
state_cache_page_size: int = 1,
):
self.ratio = ratio
self.ring_size = ring_size
self.swa_page_size = swa_page_size
self.page_size = state_cache_page_size
self.enable_memory_saver = enable_memory_saver
self.online = online
self.online_mtp_state_slot_offset = 0
self.online_mtp_max_draft_tokens = 0
@@ -115,11 +117,12 @@ class CompressStatePool:
last_dim = 3 * head_dim
else:
self._size = size + self.ring_size + 1
# Pad to lcm(ratio, page_size) so the flat buffer reshapes cleanly into
# [block_num, page_size, last_dim] for the fused compressor op; page_size=1 falls back to ratio-only padding.
pad_to = (
_lcm(ratio, swa_page_size) if (swa_page_size > 1 and _is_npu) else ratio
)
# The common GPU pool is flat by default. A backend that also needs
# a physical 3-D cache view can request its second-axis page size;
# allocation and ring ownership still stay in this shared class.
pad_to = ratio
if state_cache_page_size > 1:
pad_to = _lcm(pad_to, state_cache_page_size)
self._size = (self._size + pad_to - 1) // pad_to * pad_to
self._logical_size = self._size
last_dim = 2 * (1 + overlap) * head_dim
@@ -148,10 +151,9 @@ class CompressStatePool:
:class:`KVAndScore`. Sets ``self.memory_saver_adapter``,
``self.custom_mem_pool`` and ``self.kv_score_buffer``.
Subclasses (e.g. :class:`NPUCompressStatePool`) that compute a
different ``self._size`` reuse this instead of duplicating the
allocation boilerplate. Requires ``self._size`` and ``self.last_dim``
to be set already.
The shared constructor computes ``self._size`` and ``self.last_dim``
before entering this helper. Backend subclasses should normally call
that constructor instead of duplicating this allocation path.
"""
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=enable_memory_saver
@@ -1082,37 +1082,18 @@ class KVCacheConfigurator:
else:
compression_ratios = self.model_config.compress_ratios
# NPU + DSV4 → paged-state subclass: the fused compressor kernel
# needs cache_mode=1 (paged); Atlas A3 rejects cache_mode=2 (ring),
# so the CUDA ring-buffer state path can't be shared. CUDA keeps
# DeepSeekV4TokenToKVPool unchanged; NPU recomputes state sizes below.
# NPU keeps its PA_ND KV-pool subclass, while Compressor state sizing
# follows the same fixed ring ownership as GPU. Do not replace the
# configurator's C4-SWA/C128-request budgets with a paged allocator
# estimate: Atlas A3 cache_mode=2 consumes explicit flat state_locs.
if _is_npu:
from sglang.srt.hardware_backend.npu.dsv4.dsv4_memory_pool import (
DSV4NPUTokenToKVPool,
npu_state_pool_size,
)
pool_cls = DSV4NPUTokenToKVPool
# Recompute state pool sizes for the NPU paged formula (CUDA's
# ring sizes are dropped here). Tail-only allocation keeps the
# per-req-budget formula sufficient at any prefill length: long
# prompts allocate only ``tail+128`` (c4) / ``tail`` (c128)
# slots (tail = seq_len % 128), and decode is drained by
# sliding eviction in ``ScheduleBatch._evict_swa``.
c4_state_pool_size = npu_state_pool_size(
ratio=4,
page_size=get_schedule().page_size,
max_num_reqs=max_running_requests,
)
c128_state_pool_size = npu_state_pool_size(
ratio=128,
page_size=get_schedule().page_size,
max_num_reqs=max_running_requests,
)
else:
pool_cls = DeepSeekV4TokenToKVPool
c4_state_pool_size = c4_state_pool_size
c128_state_pool_size = c128_state_pool_size
token_to_kv_pool = pool_cls(
max_num_reqs=max_running_requests,
+11
View File
@@ -176,6 +176,17 @@ def _create_unified_radix_cache(
if ctx.is_hybrid_ssm:
tree_components.append(ComponentType.MAMBA)
if hasattr(params.req_to_token_pool, "req_to_c128_sidecar"):
from sglang.srt.hardware_backend.npu.dsv4.c128_sidecar_component import (
C128SidecarComponent,
)
tree_components.append(ComponentType.C128)
params.component_registry_override = {
**(params.component_registry_override or {}),
ComponentType.C128: C128SidecarComponent,
}
params.tree_components = tuple(tree_components)
if use_mlx() and ctx.is_hybrid_ssm:
from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import (
@@ -9,6 +9,7 @@ class ComponentType(int, Enum):
FULL = 0
SWA = 1
MAMBA = 2
C128 = 3
def __str__(self) -> str: # keep human-readable logging
return self.name.lower()
@@ -454,6 +454,7 @@ class UnifiedRadixCache(BasePrefixCache):
ComponentType.FULL: params.num_tokens,
ComponentType.SWA: params.swa_num_tokens,
ComponentType.MAMBA: params.mamba_num,
ComponentType.C128: 0,
}
self._evict_components(request_by_type, tracker)
@@ -808,7 +809,7 @@ class UnifiedRadixCache(BasePrefixCache):
result = self.insert(insert_params)
# Match prefix
match_result = self.match_prefix(MatchPrefixParams(key=radix_key))
match_result = self.match_prefix(MatchPrefixParams(key=radix_key, req=req))
new_indices = match_result.device_indices
new_last_node = match_result.last_device_node
new_prefix_len = result.prefix_len
@@ -301,59 +301,23 @@ def compute_local_num_token_non_padded_cpu(
class DSV4OutCacheLoc:
"""Per-forward-pass KV cache allocation for DeepSeek-V4 on NPU.
Bundles slot indices for full/SWA pools, the two compressed-KV pools
(c4/c128), and the two compressed-state pools (c4_state/c128_state).
Bundles slot indices for full/SWA pools and the two compressed-KV pools
(C4/C128). Compressor state uses fixed ring storage and explicit
``state_loc`` metadata, so it is not part of the token-allocation bundle.
Populated by the NPU V4 allocator (DSV4NPUTokenToKVPoolAllocator) when
the model is DeepSeek-V4 on NPU; left as ``None`` on ForwardBatch
otherwise. CUDA's DSV4 path doesn't construct this bundle (state is
derived via translate_kv_loc_to_compress_state_loc there).
otherwise.
All fields are token-level slot ids in their respective pools (NOT page
ids). Attention backends convert to page ids via ``// page_size`` when
constructing PA_ND block tables.
State fields default to ``None`` so the bundle is constructible from
paths that allocate KV but not state (or vice versa); the NPU allocator
fills all six on real alloc, CUDA paths leave state ones None and use
the ring-hash translation instead.
"""
out_full_loc: torch.Tensor
out_swa_loc: torch.Tensor
out_c4_loc: torch.Tensor
out_c128_loc: torch.Tensor
out_c4_state_loc: Optional[torch.Tensor] = None
out_c128_state_loc: Optional[torch.Tensor] = None
@dataclass
class DSV4StateLens:
"""Per-extend/decode c4/c128 compress-state pool allocation lens (DSV4-NPU).
Built by ``ScheduleBatch._compute_dsv4_state_lens_{extend,decode}`` and
threaded through ``mem_cache/common.py`` to
``DSV4NPUTokenToKVPoolAllocator.alloc_{extend,decode}``, which consumes:
* ``c{4,128}_prefix_lens`` / ``..._cpu`` — per-req prev cumulative
state-slot count (the paged allocator's ``prefix`` contract).
* ``c{4,128}_seq_lens`` / ``..._cpu`` — per-req new cumulative count.
* ``c{4,128}_extend_num_tokens`` — total new state slots this step.
Replaces the 10 loose ``c{4,128}_state_*`` kwargs the allocator used to
take: scheduler only produces this object, common only forwards it, the
allocator only consumes it.
"""
c4_prefix_lens: torch.Tensor
c4_prefix_lens_cpu: torch.Tensor
c4_seq_lens: torch.Tensor
c4_seq_lens_cpu: torch.Tensor
c4_extend_num_tokens: int
c128_prefix_lens: torch.Tensor
c128_prefix_lens_cpu: torch.Tensor
c128_seq_lens: torch.Tensor
c128_seq_lens_cpu: torch.Tensor
c128_extend_num_tokens: int
@dataclass
+123 -3
View File
@@ -734,7 +734,10 @@ class MQALayer(MqaAttentionBase):
self.register_buffer("cos_cache", cos_cache, persistent=False)
self.register_buffer("sin_cache", sin_cache, persistent=False)
if envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get() and alt_streams is not None:
if alt_streams is not None and (
(_is_cuda and envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get())
or (_is_npu and envs.SGLANG_NPU_USE_MULTI_STREAM.get())
):
self.alt_streams = alt_streams[:3]
self.alt_streams_indexer = alt_streams[-2:]
else:
@@ -957,6 +960,108 @@ class MQALayer(MqaAttentionBase):
return q
def _forward_prepare_multi_stream_npu(
self,
x: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
attn_backend,
q_out: Optional[torch.Tensor] = None,
x_quant=None,
) -> torch.Tensor:
# NPU multi-stream: KV on stream_kv, Q on stream_q, overlapped with
# indexer/compressor on current. rope is split; the kv-only call passes
# kv.unsqueeze(1) as q_rope so the op sees [T,1,1,head_dim] like the
# fused path.
assert self.alt_streams is not None
current_stream = torch.npu.current_stream()
stream_kv = self.alt_streams[0]
stream_q = self.alt_streams[1]
stream_kv.wait_stream(current_stream)
stream_q.wait_stream(current_stream)
x_linear = x_quant if x_quant is not None else x
qkv_a: Optional[torch.Tensor] = None
qkv_a_ready = None
if self.fuse_wqa_wkv:
qkv_a, _ = self.wqkv_a(x_linear)
qkv_a_ready = current_stream.record_event()
if qkv_a is not None:
q_lora = qkv_a[..., : self.q_lora_rank]
else:
q_lora, _ = self.wq_a(x_linear)
q_lora = self.q_norm(q_lora)
q_lora_ready = current_stream.record_event()
# KV block on stream_kv.
with torch.npu.stream(stream_kv):
if qkv_a_ready is not None:
stream_kv.wait_event(qkv_a_ready)
if qkv_a is not None:
kv = qkv_a[..., self.q_lora_rank :]
else:
kv, _ = self.wkv(x)
kv = self.kv_norm(kv)
cos4_k, sin4_k = self._get_npu_rope_position_cache(
positions, kv.dtype, inverse=False
)
Dsv4NpuRoPE.apply_rotary_mul_inplace(
kv.unsqueeze(1),
None,
cos4_k,
sin4_k,
qk_nope_dim=self.qk_nope_head_dim,
)
attn_backend.store_cache(
layer_id=self.layer_id,
swa_k=kv,
forward_batch=forward_batch,
)
# Q block on stream_q (needs only q_lora).
with torch.npu.stream(stream_q):
stream_q.wait_event(q_lora_ready)
q, _ = self.wq_b(q_lora)
q = q.view(-1, self.n_local_heads, self.head_dim)
_dummy = q.new_ones(q.shape[-1])
q = torch_npu.npu_rms_norm(q, _dummy, self.eps)[0]
cos4_q, sin4_q = self._get_npu_rope_position_cache(
positions, q.dtype, inverse=False
)
Dsv4NpuRoPE.apply_rotary_mul_inplace(
q,
None,
cos4_q,
sin4_q,
qk_nope_dim=self.qk_nope_head_dim,
)
if q_out is not None:
q_out.copy_(q)
q.record_stream(stream_q)
del qkv_a
# Indexer + compressor: serial on current.
if self.indexer is not None:
self.indexer(
x=x,
q_lora=q_lora,
forward_batch=forward_batch,
attn_backend=attn_backend,
)
if self.compressor is not None:
attn_backend.forward_core_compressor(
x,
forward_batch,
self.layer_id,
self.compressor,
)
# Join stream_kv + stream_q before downstream attention.
current_stream.wait_stream(stream_kv)
current_stream.wait_stream(stream_q)
return q
def _forward_prepare_multi_stream_hip(
self,
x: torch.Tensor,
@@ -1306,6 +1411,12 @@ class MQALayer(MqaAttentionBase):
)
and not (self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch))
and not (_is_hip and self.compressor is None)
) or (
_is_npu
and envs.SGLANG_NPU_USE_MULTI_STREAM.get()
and self.alt_streams is not None
and x.shape[0] <= self._multi_stream_bs_limit
and not forward_batch.forward_mode.is_extend_or_draft_extend_or_mixed()
)
tp_slice, q_padded, q_out = slice(None), None, None
@@ -1339,6 +1450,15 @@ class MQALayer(MqaAttentionBase):
q_out,
x_quant=x_quant,
)
elif _is_npu:
q = self._forward_prepare_multi_stream_npu(
x,
positions,
forward_batch,
attn_backend,
q_out,
x_quant=x_quant,
)
else:
q = self._forward_prepare_multi_stream(
x,
@@ -1506,7 +1626,7 @@ class DeepseekV4DecoderLayer(nn.Module):
layer_id=layer_id,
quant_config=quant_config,
prefix=add_prefix("self_attn", prefix),
alt_streams=None if _is_npu else alt_streams,
alt_streams=alt_streams,
compress_ratio_override=compress_ratio_override,
)
moe_alt_stream = (
@@ -2353,7 +2473,7 @@ class DeepseekV4Model(nn.Module):
or (_is_npu and envs.SGLANG_NPU_USE_MULTI_STREAM.get())
)
device_module = torch.get_device_module()
num_alt_streams = 5 if _is_cuda else 2
num_alt_streams = 5 if (_is_cuda or _is_npu) else 2
self.alt_streams = (
[device_module.Stream() for _ in range(num_alt_streams)]
if use_stream_pool
+171 -12
View File
@@ -18,12 +18,16 @@ from sglang.kernels.ops.speculative.dspark.dspark_draft_model import (
)
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.moe.utils import is_shared_experts_fusion_disabled
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.layers.vocab_parallel_embedding import (
ParallelLMHead,
VocabParallelEmbedding,
)
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
@@ -53,7 +57,7 @@ from sglang.srt.speculative.ragged_verify import (
RaggedVerifyMode,
read_ragged_verify_mode,
)
from sglang.srt.utils import add_prefix, is_blackwell_supported
from sglang.srt.utils import add_prefix, is_blackwell_supported, is_npu
from sglang.srt.utils.invariants import Bucket, InClosedRange, Invariant, expect
logger = logging.getLogger(__name__)
@@ -64,6 +68,7 @@ _PAD_NUM_HEADS = 64
_CONFIDENCE = Invariant(
"dspark.model.confidence", Bucket.GUARD, InClosedRange(0.0, 1.0)
)
_is_npu = is_npu()
def apply_rotary_emb(
@@ -126,6 +131,12 @@ class DSparkAttention(MqaAttentionBase):
self._use_fast_kernel = envs.SGLANG_DSPARK_FAST_KERNEL.get()
self.alt_streams = alt_streams
self._multi_stream_bs_limit = 128 if is_blackwell_supported() else 64
if _is_npu:
self.register_buffer(
"_q_post_norm_weight",
torch.ones(self.head_dim),
persistent=False,
)
def kv_proj_only(self, x: torch.Tensor) -> torch.Tensor:
kv, _ = self.wkv(x)
@@ -180,10 +191,36 @@ class DSparkAttention(MqaAttentionBase):
fused_q_norm_rope(q, q_out, self.eps, self.freqs_cis, positions)
return q_out
else:
q = q * torch.rsqrt(
q.float().square().mean(-1, keepdim=True) + self.eps
).to(q.dtype)
apply_rotary_emb(q[..., -self.rope_head_dim :], self.freqs_cis[positions])
if _is_npu:
import torch_npu
from sglang.srt.hardware_backend.npu.dsv4.dsv4_rope import (
Dsv4NpuRoPE,
)
q = torch_npu.npu_rms_norm(q, self._q_post_norm_weight, self.eps)[0]
cos4, sin4 = Dsv4NpuRoPE.for_freqs(self.freqs_cis).get_cos_sin(
positions,
q.dtype,
view_4d=True,
inverse=False,
allow_build=True,
cache_dtype=torch.float32,
)
Dsv4NpuRoPE.apply_rotary_mul_inplace(
q,
None,
cos4,
sin4,
qk_nope_dim=q.shape[-1] - self.rope_head_dim,
)
else:
q = q * torch.rsqrt(
q.float().square().mean(-1, keepdim=True) + self.eps
).to(q.dtype)
apply_rotary_emb(
q[..., -self.rope_head_dim :], self.freqs_cis[positions]
)
if q_out is not None:
q_out.copy_(q)
return q_out
@@ -195,6 +232,10 @@ class DSparkAttention(MqaAttentionBase):
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
if _is_npu and forward_batch.forward_mode.is_idle():
return torch.zeros_like(hidden_states)
from sglang.srt.model_executor.forward_context import get_attn_backend
pool = _resolve_dspark_pool()
@@ -255,6 +296,7 @@ class DSparkAttention(MqaAttentionBase):
attn_sink=attn_sink,
save_kv_cache=False,
)
if o.shape[1] != self.n_local_heads:
o = o[:, : self.n_local_heads, :]
@@ -262,6 +304,24 @@ class DSparkAttention(MqaAttentionBase):
fused_rope_inplace(
o[..., -rd:], None, self.freqs_cis, positions=positions, inverse=True
)
elif _is_npu:
from sglang.srt.hardware_backend.npu.dsv4.dsv4_rope import Dsv4NpuRoPE
cos4, sin4 = Dsv4NpuRoPE.for_freqs(self.freqs_cis).get_cos_sin(
positions,
o.dtype,
view_4d=True,
inverse=True,
allow_build=True,
cache_dtype=torch.float32,
)
Dsv4NpuRoPE.apply_rotary_mul_inplace(
o,
None,
cos4,
sin4,
qk_nope_dim=o.shape[-1] - rd,
)
else:
apply_rotary_emb(o[..., -rd:], self.freqs_cis[positions], inverse=True)
@@ -581,6 +641,26 @@ class DSparkV4Stage(DeepseekV4DecoderLayer):
class DeepseekV4ForCausalLMDSpark(nn.Module):
# ModelSlim NPU checkpoints carry QuaRot-aligned, MTP-local
# embedding/head weights. The native CUDA path keeps the original DSpark
# behavior and shares the target model's vocabulary modules.
uses_own_vocab_modules = _is_npu
@classmethod
def shared_experts_fusion_disable_reason(
cls,
hf_config,
quant_config,
):
if _is_npu:
return (
"NPU DSpark ModelSlim weight loading does not support mapping "
"shared experts into fused expert slots."
)
return DeepseekV4ForCausalLM.shared_experts_fusion_disable_reason(
hf_config, quant_config
)
@classmethod
def shared_experts_fusion_disable_reason(cls, hf_config, quant_config):
@@ -662,10 +742,26 @@ class DeepseekV4ForCausalLMDSpark(nn.Module):
self.norm_eps = float(config.rms_norm_eps)
self.hc_eps = float(config.hc_eps)
self.embed_tokens: Optional[nn.Module] = None
self.lm_head: Optional[nn.Module] = None
if self.uses_own_vocab_modules:
self.embed_tokens = VocabParallelEmbedding(
config.vocab_size,
config.hidden_size,
prefix=add_prefix("embed_tokens", prefix),
enable_tp=not is_dp_attention_enabled(),
)
self.lm_head = ParallelLMHead(
config.vocab_size,
config.hidden_size,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_parallel().enable_dp_lm_head,
)
else:
self.embed_tokens: Optional[nn.Module] = None
self.lm_head: Optional[nn.Module] = None
self._use_fp32_lm_head = envs.SGLANG_DSPARK_FP32_LM_HEAD.get()
self._opt_markov_w2_tp_shard = envs.SGLANG_DSPARK_OPT_MARKOV_W2_TP_SHARD.get()
if self.lm_head is not None:
self.markov_head.configure_tp_shard(lm_head=self.lm_head)
@property
def enable_confidence_head(self) -> bool:
@@ -674,9 +770,10 @@ class DeepseekV4ForCausalLMDSpark(nn.Module):
def attach_shared_modules(
self, *, embed_tokens: nn.Module, lm_head: nn.Module
) -> None:
self.embed_tokens = embed_tokens
self.lm_head = lm_head
self.markov_head.configure_tp_shard(lm_head=lm_head)
if not self.uses_own_vocab_modules:
self.embed_tokens = embed_tokens
self.lm_head = lm_head
self.markov_head.configure_tp_shard(lm_head=self.lm_head)
def project_target_hidden(self, main_hidden: torch.Tensor) -> torch.Tensor:
stage0 = self.stages[0]
@@ -820,7 +917,11 @@ class DeepseekV4ForCausalLMDSpark(nn.Module):
)
for name, loaded_weight in weights:
mapped = self._remap_dspark_weight_name(name)
mapped = (
self._remap_dspark_weight_name_npu(name)
if _is_npu
else self._remap_dspark_weight_name(name)
)
if mapped is None:
continue
if self.num_fused_shared_experts > 0 and ".mlp.shared_experts." in mapped:
@@ -929,5 +1030,63 @@ class DeepseekV4ForCausalLMDSpark(nn.Module):
mapped_rest = mapped_rest.replace(".scale", ".weight_scale_inv")
return f"stages.{stage_id}.{mapped_rest}"
def _remap_dspark_weight_name_npu(self, name: str) -> Optional[str]:
if name.startswith(("embed.", "embed_tokens.", "head.", "lm_head.")):
return None
if "rotary_emb.inv_freq" in name:
return None
if not name.startswith("mtp."):
return None
parts = name.split(".", 2)
if len(parts) < 3:
return None
stage_id, rest = parts[1], parts[2]
if not stage_id.isdigit():
return None
stage_idx = int(stage_id)
if rest in ("embed.weight", "embed_tokens.weight"):
return (
"embed_tokens.weight"
if self.uses_own_vocab_modules and stage_idx == 0
else None
)
if rest in ("head.weight", "lm_head.weight"):
return (
"lm_head.weight"
if self.uses_own_vocab_modules and stage_idx == self.num_stages - 1
else None
)
if rest.startswith(("embed.", "embed_tokens.", "head.", "lm_head.")):
return None
if rest.startswith("markov_head."):
return f"markov_head.{rest[len('markov_head.'):]}"
if rest.startswith("confidence_head."):
if self.confidence_head is None:
return None
return f"confidence_head.{rest[len('confidence_head.'):]}"
mapped_rest = rest
if mapped_rest.startswith("attn."):
mapped_rest = "self_attn." + mapped_rest.removeprefix("attn.")
elif mapped_rest.startswith("ffn."):
mapped_rest = "mlp." + mapped_rest.removeprefix("ffn.")
elif mapped_rest.startswith("attn_norm."):
mapped_rest = "input_layernorm." + mapped_rest.removeprefix("attn_norm.")
elif mapped_rest.startswith("ffn_norm."):
mapped_rest = "post_attention_layernorm." + mapped_rest.removeprefix(
"ffn_norm."
)
mapped_rest = mapped_rest.replace(".w1.", ".gate_proj.")
mapped_rest = mapped_rest.replace(".w2.", ".down_proj.")
mapped_rest = mapped_rest.replace(".w3.", ".up_proj.")
mapped_rest = mapped_rest.replace(".gate.tid2eid", ".topk.tid2eid")
mapped_rest = mapped_rest.replace(".gate.bias", ".gate.e_score_correction_bias")
if mapped_rest.endswith(".scale"):
mapped_rest = mapped_rest.removesuffix(".scale") + ".weight_scale_inv"
return f"stages.{stage_id}.{mapped_rest}"
EntryClass = [DeepseekV4ForCausalLMDSpark]
+5
View File
@@ -893,6 +893,11 @@ class ServerArgs:
Arg(help="The number of tokens in a page.", resolvable=True),
NS("schedule"),
] = None
c128_page_size: A[
int,
"The physical page size of the NPU DSV4 C128 KV cache. Must be a positive multiple of 16.",
NS("schedule"),
] = 16
swa_full_tokens_ratio: A[
float,
Arg(
@@ -0,0 +1,32 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from sglang.srt.managers.overlap_utils import RelayPayload
from sglang.srt.speculative.draft_worker_common import make_draft_input_v2
if TYPE_CHECKING:
from sglang.srt.managers.overlap_utils import FutureMap
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
def build_dflash_family_disagg_draft_input(
batch: ScheduleBatch,
last_tokens_tensor: torch.Tensor,
future_map: FutureMap,
) -> DFlashDraftInputV2:
spec_info = make_draft_input_v2(
bonus_tokens=last_tokens_tensor,
new_seq_lens=batch.seq_lens,
)
if batch.enable_overlap:
spec_info.future_indices = batch.req_pool_indices
future_map.publish(spec_info.future_indices, batch.seq_lens)
future_map.stash(
spec_info.future_indices,
RelayPayload(bonus_tokens=last_tokens_tensor),
)
return spec_info
+27 -4
View File
@@ -13,11 +13,14 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode,
)
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
from sglang.srt.utils import is_npu
if TYPE_CHECKING:
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
_is_npu = is_npu()
@dataclass
class DFlashVerifyInput(SpecInput):
@@ -43,6 +46,9 @@ class DFlashVerifyInput(SpecInput):
num_tokens_per_req: int = -1
ragged_verify_layout: Optional[RaggedVerifyLayout] = None
# Committed/live lengths before the verify caller temporarily expands
# batch.seq_lens_cpu to the target-attention KV lengths.
live_seq_lens_cpu: Optional[torch.Tensor] = None
def __post_init__(self):
super().__init__(spec_input_type=SpecInputType.DFLASH_VERIFY)
@@ -58,14 +64,25 @@ class DFlashVerifyInput(SpecInput):
"""Prepare a DFLASH verify forward batch for overlap scheduling.
The caller computes and stores `batch.out_cache_loc` before this
method is called. This helper only packages the verify forward and pre-initializes either CUDA-graph replay
metadata or eager attention metadata so the actual forward can run with
`skip_attn_backend_init=True`.
method is called. GPU keeps the original pre-planning path. NPU leaves
attention/graph metadata initialization to ModelRunner because DP/EP
padding can still change the compressor's runtime shapes.
"""
from sglang.srt.speculative.spec_utils import prepare_mamba_track_for_verify
batch.input_ids = self.draft_token
batch.spec_info = self
if _is_npu and not batch.forward_mode.is_idle():
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
maybe_build_dsv4_verify_bundle,
)
batch.out_cache_loc_dsv4 = maybe_build_dsv4_verify_bundle(
batch,
self.draft_token_num,
live_seq_lens_cpu=self.live_seq_lens_cpu,
)
batch.forward_mode = (
ForwardMode.IDLE
if batch.forward_mode.is_idle()
@@ -90,7 +107,13 @@ class DFlashVerifyInput(SpecInput):
verify_forward_batch
)
)
if can_run_cuda_graph:
if _is_npu:
# Do not pre-plan target verify on NPU. DP/EP padding can change
# the compressor's logical batch without changing ForwardBatch's
# stale-plan shape fields. Let ModelRunner select graph/eager and
# initialize metadata after final batch preparation.
return verify_forward_batch, can_run_cuda_graph
elif can_run_cuda_graph:
target_worker.model_runner.decode_cuda_graph_runner.load_batch(
verify_forward_batch
)
@@ -124,6 +124,9 @@ class DFlashDraftInputV2(SpecInput):
bs = batch.batch_size()
if bs == 0:
return
batch.maybe_evict_swa()
self._ensure_prepare_length_buffers(bs, batch.device)
assert self._prepare_batch_seq_lens_cpu_buf is not None
assert self._prepare_cur_kv_lens_cpu_buf is not None
@@ -205,7 +208,8 @@ class DFlashDraftInputV2(SpecInput):
# plan-stream context, so forward work cannot observe partially
# prepared req_to_token / KV allocation state.
caller_stream.wait_stream(plan_stream)
for req in batch.reqs:
req.decode_batch_idx += 1
# Seed committed; overlap's resolve overwrites it with the published value.
batch.seq_lens_cpu = batch_seq_lens_cpu_t
batch.seq_lens_sum = committed_seq_lens_sum
@@ -396,6 +396,10 @@ class DraftBlockProposer:
forward_batch.can_run_dp_cuda_graph = batch.can_run_dp_cuda_graph
if not self._dp_moe_sync or batch.global_num_tokens is None:
return
# Graph bucket selection uses the raw per-rank request counts. Keep
# them separate from global_num_tokens_cpu below, which is scaled into
# draft-token units for DP/MoE synchronization.
forward_batch.original_global_num_tokens_cpu = batch.global_num_tokens
gnt, gnt_logprob = spec_scale_global_num_tokens(
self._draft_block_spec_info,
batch.global_num_tokens,
@@ -146,7 +146,9 @@ def _resolve_folded_sampling(*, model, gamma, max_bs, device, tp_rank) -> bool:
noise_bytes = max_bs * vocab * 4
logits_bytes = max_bs * gamma * vocab * model.lm_head.weight.dtype.itemsize
need_gb = (noise_bytes + logits_bytes) / (1 << 30)
available_gb = get_available_gpu_memory(device, torch.cuda.current_device())
available_gb = get_available_gpu_memory(
device, torch.get_device_module().current_device()
)
if available_gb - need_gb >= _CAPTURE_HEADROOM_GB:
return True
if tp_rank == 0:
@@ -45,8 +45,11 @@ from sglang.srt.speculative.spec_utils import (
SIMULATE_ACC_METHOD,
sample_simulated_acc_len,
)
from sglang.srt.utils import is_npu
from sglang.srt.utils.invariants import Bucket, Invariant, NotNaN, expect
_is_npu = is_npu()
# Draft proposal probs feeding rejection sampling; the data layer is the
# in-kernel NaN-q guard in reject_sampling.py, so this is signal-only.
_VERIFY_DRAFT_PROBS = Invariant("dspark.verify.draft_probs", Bucket.GUARD, NotNaN())
@@ -214,6 +217,7 @@ class TargetVerifyExecutor:
batch.seq_lens_cpu = torch.ones((num_dummy_slots,), dtype=torch.int64)
batch.seq_lens_sum = num_dummy_slots
batch.forward_mode = ForwardMode.TARGET_VERIFY
verify_input.live_seq_lens_cpu = batch.seq_lens_cpu
verify_forward_batch, _ = verify_input.prepare_for_verify(
batch, self.target_worker
)
@@ -221,7 +225,7 @@ class TargetVerifyExecutor:
batch=None,
forward_batch=verify_forward_batch,
is_verify=True,
skip_attn_backend_init=True,
skip_attn_backend_init=True if not _is_npu else None,
)
def run_non_compact(
@@ -243,6 +247,7 @@ class TargetVerifyExecutor:
draft_token_num=verify_w,
custom_mask=None,
capture_hidden_mode=CaptureHiddenMode.FULL,
live_seq_lens_cpu=batch.seq_lens_cpu,
)
batch.out_cache_loc = verify_cache_loc
seq_lens_cpu_backup = batch.seq_lens_cpu
@@ -289,7 +294,7 @@ class TargetVerifyExecutor:
batch=None,
forward_batch=verify_forward_batch,
is_verify=True,
skip_attn_backend_init=True,
skip_attn_backend_init=True if not _is_npu else None,
)
return TargetVerifyResult(
logits_output=target_out.logits_output,
@@ -355,6 +360,7 @@ class TargetVerifyExecutor:
custom_mask=None,
capture_hidden_mode=CaptureHiddenMode.FULL,
ragged_verify_layout=layout,
live_seq_lens_cpu=batch.seq_lens_cpu,
)
batch.out_cache_loc = ragged_window.verify_cache_loc
seq_lens_cpu_backup = batch.seq_lens_cpu
@@ -66,10 +66,12 @@ from sglang.srt.speculative.spec_utils import (
draft_tp_context,
prepare_mamba_track_for_verify,
)
from sglang.srt.utils import get_available_gpu_memory, is_cuda
from sglang.srt.utils import get_available_gpu_memory, is_cuda, is_npu
logger = logging.getLogger(__name__)
_is_npu = is_npu()
class DSparkWorkerV2(BaseSpecWorker):
@@ -173,16 +175,22 @@ class DSparkWorkerV2(BaseSpecWorker):
draft_token_num=int(self.gamma), device=self.device
)
target_model = self.target_worker.model_runner.model
lm_head = getattr(target_model, "lm_head", None)
if lm_head is None or not hasattr(lm_head, "weight"):
raise RuntimeError(
"DSpark requires the target model to expose `lm_head` with `weight`."
if getattr(self.draft_model, "uses_own_vocab_modules", False):
if self.ps.tp_rank == 0:
logger.info(
"DSpark draft uses its checkpoint-local embedding and LM head."
)
else:
target_model = self.target_worker.model_runner.model
lm_head = getattr(target_model, "lm_head", None)
if lm_head is None or not hasattr(lm_head, "weight"):
raise RuntimeError(
"DSpark requires the target model to expose `lm_head` with `weight`."
)
self.draft_model.attach_shared_modules(
embed_tokens=self._resolve_target_embed_tokens(target_model),
lm_head=lm_head,
)
self.draft_model.attach_shared_modules(
embed_tokens=self._resolve_target_embed_tokens(target_model),
lm_head=lm_head,
)
self._verify_planner = DSparkVerifyPlanner(
draft_model=self.draft_model,
@@ -329,6 +337,12 @@ class DSparkWorkerV2(BaseSpecWorker):
def init_attention_backends(self):
with self._draft_context():
if _is_npu:
from sglang.srt.hardware_backend.npu.extra_ops_loader import (
initialize_dspark_sparse_attn_ops,
)
initialize_dspark_sparse_attn_ops()
self._draft_worker.init_attention_backends()
self._need_mamba_verify_commit = mambaish_config(
self.model_runner.model_config
@@ -570,7 +584,6 @@ class DSparkWorkerV2(BaseSpecWorker):
self._observers.begin_step()
target_model = self.target_worker.model_runner.model
verify_window = alloc_verify_window(
batch=batch,
bs=bs,
@@ -675,7 +688,6 @@ class DSparkWorkerV2(BaseSpecWorker):
hidden_strided = None
logits_output = target_verify.logits_output
can_run_cuda_graph = target_verify.can_run_cuda_graph
if batch.has_grammar:
# run_compact scatters its rows back to (bs * chain_len), so the mask
# lines up with the logits on both verify paths.
@@ -59,7 +59,6 @@ from sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend import (
DeepseekV4AscendMultiStepDraftBackend,
_apply_hadamard,
_get_kv_indices,
_overlap_transform,
_walsh_hadamard_matrix,
)
@@ -182,78 +181,6 @@ class TestApplyHadamard(unittest.TestCase):
self.assertTrue(torch.equal(out, expected))
class TestOverlapTransform(unittest.TestCase):
def test_shape(self):
# (n_chunks, ratio, 2*d) -> (n_chunks, 2*ratio, d)
n_chunks, r, d = 3, 2, 4
tensor = torch.randn(n_chunks, r, 2 * d)
out = _overlap_transform(tensor, value=0.0, head_dim=d)
self.assertEqual(out.shape, (n_chunks, 2 * r, d))
def test_first_chunk_left_half_filled_with_value(self):
n_chunks, r, d = 3, 2, 4
tensor = torch.randn(n_chunks, r, 2 * d)
fill = float("-inf")
out = _overlap_transform(tensor, value=fill, head_dim=d)
self.assertTrue(torch.equal(out[0, :r], torch.full((r, d), fill)))
def test_first_chunk_left_half_filled_with_zero(self):
n_chunks, r, d = 2, 2, 4
tensor = torch.randn(n_chunks, r, 2 * d)
out = _overlap_transform(tensor, value=0.0, head_dim=d)
self.assertTrue(torch.equal(out[0, :r], torch.zeros(r, d)))
def test_right_half_mirrors_tensor_second_half(self):
n_chunks, r, d = 3, 2, 4
tensor = torch.randn(n_chunks, r, 2 * d)
out = _overlap_transform(tensor, value=0.0, head_dim=d)
self.assertTrue(torch.equal(out[:, r:], tensor[..., d:]))
def test_previous_chunk_left_half(self):
n_chunks, r, d = 3, 2, 4
tensor = torch.randn(n_chunks, r, 2 * d)
out = _overlap_transform(tensor, value=0.0, head_dim=d)
self.assertTrue(torch.equal(out[1:, :r], tensor[:-1, :, :d]))
def test_single_chunk(self):
n_chunks, r, d = 1, 2, 4
tensor = torch.randn(n_chunks, r, 2 * d)
fill = 7.0
out = _overlap_transform(tensor, value=fill, head_dim=d)
self.assertEqual(out.shape, (1, 2 * r, d))
self.assertTrue(torch.equal(out[0, :r], torch.full((r, d), fill)))
self.assertTrue(torch.equal(out[0, r:], tensor[0, :, d:]))
def test_full_element_mapping(self):
n_chunks, r, d = 2, 2, 3
tensor = torch.arange(n_chunks * r * 2 * d, dtype=torch.float32).reshape(
n_chunks, r, 2 * d
)
fill = -1.0
out = _overlap_transform(tensor, value=fill, head_dim=d)
for c in range(n_chunks):
for row in range(2 * r):
for col in range(d):
if c == 0 and row < r:
expected = fill
elif row >= r:
expected = tensor[c, row - r, d + col].item()
else:
expected = tensor[c - 1, row, col].item()
self.assertEqual(
out[c, row, col].item(),
expected,
f"mismatch at (c={c}, row={row}, col={col})",
)
def test_preserves_input_dtype(self):
n_chunks, r, d = 2, 2, 4
tensor = torch.randn(n_chunks, r, 2 * d, dtype=torch.bfloat16)
out = _overlap_transform(tensor, value=0.0, head_dim=d)
self.assertEqual(out.dtype, torch.bfloat16)
class TestGetKvIndices(unittest.TestCase):
_PATCH_TARGET = (
"sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend.get_attn_backend"
@@ -41,6 +41,10 @@ _EVICT_METHOD = "maybe_evict_swa"
# Any added/removed/recounted site fails until reviewed here.
_SB = "managers/schedule_batch.py"
_EAGLE_DECODE = ("speculative/eagle_utils.py", "eagle_prepare_for_decode")
_DFLASH_DECODE = (
"speculative/dflash_info_v2.py",
"DFlashDraftInputV2.prepare_for_decode",
)
_RESOLVE = (
"managers/scheduler_components/batch_result_processor.py",
"SchedulerBatchResultProcessor._resolve_spec_v2_tokens",
@@ -62,6 +66,11 @@ _OWNER_SITES = {
# inside the owned-kv alloc_for_spec_decode function (op42).
(*_EAGLE_DECODE, "decode_batch_idx"): 1,
(*_EAGLE_DECODE, "evict"): 1,
# DFlash uses its stateful scheduler-side preparation instead of
# eagle_prepare_for_decode. spec_prepare_for_decode dispatches to exactly
# one of these two owners for each speculative decode iteration.
(*_DFLASH_DECODE, "decode_batch_idx"): 1,
(*_DFLASH_DECODE, "evict"): 1,
(
"mem_cache/allocation.py",
"alloc_for_spec_decode",