[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:
co-authored by
JiaruiChang5268
Kelon
unknown
Talantan1102
Talantan1102
parent
e03c53fc13
commit
b83d507cd7
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 = [
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user