[NPU] Add Ascend NPU support for DeepSeek-V4 (#25144)
Co-authored-by: khalil2ji3mp6 <khalilzhk@gmail.com> Co-authored-by: randgun <kelonlu@163.com> Co-authored-by: t00937989 <tanlei33@huawei.com>
This commit is contained in:
co-authored by
khalil2ji3mp6
randgun
t00937989
parent
3f66873304
commit
9b10821c8e
@@ -15,9 +15,29 @@ def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None
|
||||
|
||||
server_args.attention_backend = "dsv4"
|
||||
server_args.page_size = 256
|
||||
if server_args.kv_cache_dtype == "auto":
|
||||
server_args.kv_cache_dtype = "fp8_e4m3"
|
||||
logger.warning(
|
||||
f"Setting KV cache dtype to {server_args.kv_cache_dtype} for {model_arch}."
|
||||
)
|
||||
|
||||
if server_args.device == "npu":
|
||||
# NPU keeps the device-aware "dsv4" backend (the registry routes it to
|
||||
# the Ascend V4 subclass); only the pool geometry / dtype differ.
|
||||
# set_default_server_args() pins all three backends to "ascend" for
|
||||
# generic NPU models; undo that here so V4 stays consistently on dsv4.
|
||||
server_args.prefill_attention_backend = "dsv4"
|
||||
server_args.decode_attention_backend = "dsv4"
|
||||
server_args.page_size = 128
|
||||
server_args.kv_cache_dtype = "bfloat16"
|
||||
|
||||
logger.info(
|
||||
f"Use dsv4 attention backend for {model_arch}, setting page_size to 256."
|
||||
f"Use dsv4 attention backend for {model_arch}, setting page_size to {server_args.page_size}."
|
||||
)
|
||||
assert server_args.kv_cache_dtype in [
|
||||
"fp8_e4m3",
|
||||
"bfloat16",
|
||||
], f"{server_args.kv_cache_dtype} is not supported for {model_arch}"
|
||||
|
||||
if server_args.max_running_requests is None:
|
||||
server_args.max_running_requests = 256
|
||||
@@ -25,15 +45,6 @@ def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None
|
||||
f"Setting max_running_requests to {server_args.max_running_requests} for {model_arch}."
|
||||
)
|
||||
|
||||
if server_args.kv_cache_dtype == "auto":
|
||||
server_args.kv_cache_dtype = "fp8_e4m3"
|
||||
logger.warning(
|
||||
f"Setting KV cache dtype to {server_args.kv_cache_dtype} for {model_arch}."
|
||||
)
|
||||
assert server_args.kv_cache_dtype in [
|
||||
"fp8_e4m3"
|
||||
], f"{server_args.kv_cache_dtype} is not supported for {model_arch}"
|
||||
|
||||
if server_args.speculative_algorithm is not None:
|
||||
assert (
|
||||
server_args.speculative_algorithm == "EAGLE"
|
||||
|
||||
@@ -550,6 +550,16 @@ class AscendAttnBackend(AttentionBackend):
|
||||
dtype=torch.int64,
|
||||
device=self.device,
|
||||
)
|
||||
# V4-specific extra graph buffers. Default no-op on the base class;
|
||||
# DeepseekV4AscendAttnBackend overrides.
|
||||
self._init_dsv4_graph_buffers(max_bs=max_bs, max_num_tokens=max_num_tokens)
|
||||
|
||||
def _init_dsv4_graph_buffers(self, *, max_bs: int, max_num_tokens: int) -> None:
|
||||
"""Hook for V4-Flash to preallocate dsv4-specific graph buffers.
|
||||
|
||||
Default no-op. Overridden by DeepseekV4AscendAttnBackend.
|
||||
"""
|
||||
pass
|
||||
|
||||
def _init_cuda_graph_metadata(
|
||||
self,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,625 @@
|
||||
"""DSV4-NPU SWA + c4/c128 paged 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.
|
||||
|
||||
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.
|
||||
|
||||
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:
|
||||
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.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.hardware_backend.npu.allocator_npu import NPUPagedTokenToKVPoolAllocator
|
||||
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
|
||||
|
||||
|
||||
def get_last_loc(
|
||||
req_to_token: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
prefix_lens: torch.Tensor,
|
||||
) -> 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``.
|
||||
"""
|
||||
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),
|
||||
)
|
||||
|
||||
|
||||
class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
"""SWA allocator + c4/c128 KV and compress-state paged allocators for DSV4 on NPU."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size: int,
|
||||
size_swa: int,
|
||||
page_size: int,
|
||||
dtype: torch.dtype,
|
||||
device: str,
|
||||
kvcache,
|
||||
need_sort: bool,
|
||||
):
|
||||
super().__init__(
|
||||
size=size,
|
||||
size_swa=size_swa,
|
||||
page_size=page_size,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
kvcache=kvcache,
|
||||
need_sort=need_sort,
|
||||
)
|
||||
|
||||
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.
|
||||
return NPUPagedTokenToKVPoolAllocator(
|
||||
pool_size,
|
||||
page_size=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
|
||||
)
|
||||
|
||||
# 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.
|
||||
self._cur_req_to_token_pool = None
|
||||
|
||||
@staticmethod
|
||||
def _compute_c_extend_counts(
|
||||
prefix_lens_cpu: torch.Tensor,
|
||||
seq_lens_cpu: torch.Tensor,
|
||||
ratio: int,
|
||||
) -> int:
|
||||
"""New compressed-K tokens this extend produces across the batch:
|
||||
``sum_i (seq_lens[i] // ratio - prefix_lens[i] // ratio)``."""
|
||||
if prefix_lens_cpu is None or seq_lens_cpu is None:
|
||||
return 0
|
||||
diff = ((seq_lens_cpu // ratio) - (prefix_lens_cpu // ratio)).clamp(min=0)
|
||||
return int(diff.sum().item())
|
||||
|
||||
@staticmethod
|
||||
def _pool_exhausted(
|
||||
ratio: int, kind: str, need: int, available: int
|
||||
) -> RuntimeError:
|
||||
return RuntimeError(
|
||||
f"DSV4 c{ratio} {kind} pool exhausted: need {need} new slots, "
|
||||
f"available={available}. Raise --mem-fraction-static, lower "
|
||||
f"--max-running-requests, or check that "
|
||||
f"DSV4NPUTokenToKVPoolAllocator.free(req=...) releases {kind} slots "
|
||||
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,
|
||||
prefix_lens: torch.Tensor,
|
||||
prefix_lens_cpu: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
seq_lens_cpu: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
last_loc_dtype: torch.dtype,
|
||||
ratio: int,
|
||||
) -> torch.Tensor:
|
||||
"""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
|
||||
: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
|
||||
step closes no compressed token.
|
||||
"""
|
||||
c_extend = self._compute_c_extend_counts(prefix_lens_cpu, seq_lens_cpu, ratio)
|
||||
if c_extend == 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 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_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
|
||||
)
|
||||
|
||||
result = allocator.alloc_extend(
|
||||
c_prefix,
|
||||
prefix_lens_cpu // ratio,
|
||||
c_seq,
|
||||
seq_lens_cpu // ratio,
|
||||
c_last_loc,
|
||||
c_extend,
|
||||
)
|
||||
if result is None:
|
||||
raise self._pool_exhausted(
|
||||
ratio, "KV", c_extend, allocator.available_size()
|
||||
)
|
||||
return result
|
||||
|
||||
def _alloc_c_and_state(
|
||||
self,
|
||||
out_full_loc: torch.Tensor,
|
||||
out_swa_loc: torch.Tensor,
|
||||
prefix_lens: torch.Tensor,
|
||||
prefix_lens_cpu: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
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.
|
||||
"""
|
||||
assert req_pool_indices is not None, (
|
||||
"DSV4NPUTokenToKVPoolAllocator requires req_pool_indices "
|
||||
"(forwarded from batch.req_pool_indices)."
|
||||
)
|
||||
assert dsv4_state_lens is not None, (
|
||||
"DSV4NPUTokenToKVPoolAllocator requires dsv4_state_lens "
|
||||
"(ScheduleBatch._compute_dsv4_state_lens_*)."
|
||||
)
|
||||
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_c128_loc = self._alloc_c_extend(
|
||||
self.c128_attn_allocator,
|
||||
prefix_lens,
|
||||
prefix_lens_cpu,
|
||||
seq_lens,
|
||||
seq_lens_cpu,
|
||||
req_pool_indices,
|
||||
last_loc_dtype,
|
||||
ratio=128,
|
||||
)
|
||||
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,
|
||||
)
|
||||
return DSV4OutCacheLoc(
|
||||
out_full_loc=out_full_loc,
|
||||
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]
|
||||
) -> 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 in zip(reqs, seq_lens):
|
||||
tail = seq_len % 128
|
||||
c4_alloc_len = tail + 128 if (tail <= 3 and seq_len >= 128) else tail
|
||||
c128_alloc_len = tail
|
||||
|
||||
prev_c4 = getattr(req, "c4_state_kv_len", 0)
|
||||
prev_c128 = getattr(req, "c128_state_kv_len", 0)
|
||||
new_c4 = prev_c4 + c4_alloc_len
|
||||
new_c128 = prev_c128 + c128_alloc_len
|
||||
|
||||
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_alloc_offset = seq_len - c4_alloc_len
|
||||
req.c128_state_alloc_offset = seq_len - c128_alloc_len
|
||||
|
||||
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 _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(
|
||||
self,
|
||||
prefix_lens: torch.Tensor,
|
||||
prefix_lens_cpu: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
seq_lens_cpu: torch.Tensor,
|
||||
last_loc: torch.Tensor,
|
||||
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
|
||||
out_full_loc = super().alloc_extend(
|
||||
prefix_lens,
|
||||
prefix_lens_cpu,
|
||||
seq_lens,
|
||||
seq_lens_cpu,
|
||||
last_loc,
|
||||
extend_num_tokens,
|
||||
)
|
||||
if out_full_loc is None:
|
||||
return None
|
||||
|
||||
out_swa_loc = self.translate_loc_from_full_to_swa(out_full_loc)
|
||||
assert out_swa_loc is not None, (
|
||||
"translate_loc_from_full_to_swa returned None — "
|
||||
"full_to_swa_index_mapping not initialized?"
|
||||
)
|
||||
return self._alloc_c_and_state(
|
||||
out_full_loc,
|
||||
out_swa_loc,
|
||||
prefix_lens,
|
||||
prefix_lens_cpu,
|
||||
seq_lens,
|
||||
seq_lens_cpu,
|
||||
last_loc.dtype,
|
||||
req_pool_indices,
|
||||
dsv4_state_lens,
|
||||
)
|
||||
|
||||
def alloc_decode(
|
||||
self,
|
||||
seq_lens: torch.Tensor,
|
||||
seq_lens_cpu: torch.Tensor,
|
||||
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
|
||||
out_full_loc = super().alloc_decode(seq_lens, seq_lens_cpu, last_loc)
|
||||
if out_full_loc is None:
|
||||
return None
|
||||
|
||||
out_swa_loc = self.translate_loc_from_full_to_swa(out_full_loc)
|
||||
# One new token per req. Model as an extend from (seq_len-1)//ratio to
|
||||
# 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(
|
||||
out_full_loc,
|
||||
out_swa_loc,
|
||||
prefix_lens,
|
||||
prefix_lens_cpu,
|
||||
seq_lens,
|
||||
seq_lens_cpu,
|
||||
last_loc.dtype,
|
||||
req_pool_indices,
|
||||
dsv4_state_lens,
|
||||
)
|
||||
|
||||
def free(
|
||||
self,
|
||||
free_index: Optional[torch.Tensor] = None,
|
||||
*,
|
||||
req=None,
|
||||
req_to_token_pool=None,
|
||||
):
|
||||
"""Unified free for full/swa/c4/c128 pools. 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).
|
||||
"""
|
||||
if free_index is not None:
|
||||
super().free(free_index)
|
||||
|
||||
if req is None or req_to_token_pool is None:
|
||||
return
|
||||
kv_len = req.kv_committed_len
|
||||
req_pool_idx = req.req_pool_idx
|
||||
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]
|
||||
# to int64 — paged allocator's free does cpu()//page_size on it.
|
||||
allocator.free(slots.to(torch.int64))
|
||||
|
||||
# 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]
|
||||
allocator.free(slots.to(torch.int64))
|
||||
|
||||
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",
|
||||
):
|
||||
allocator = getattr(self, attr, None)
|
||||
if allocator is not None:
|
||||
allocator.clear()
|
||||
@@ -0,0 +1,356 @@
|
||||
"""Helpers used by mem_cache/common.py to wire DSV4-NPU per-req 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
|
||||
:class:`DSV4OutCacheLoc` the allocator returned onto
|
||||
``batch.out_cache_loc_dsv4``. After each ``alloc_extend`` / ``alloc_decode``
|
||||
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`.
|
||||
|
||||
Non-DSV4 paths leave ``batch.out_cache_loc_dsv4`` None, so this module is a
|
||||
no-op for them.
|
||||
|
||||
TODO: the disagg DSV4 path bypasses these hooks — it calls
|
||||
``allocator.alloc_extend`` directly then ``req_to_token_pool.write`` without
|
||||
going through ``mem_cache/common.py`` (see ``disaggregation/decode.py``). The
|
||||
DSV4OutCacheLoc bundle is still produced but never written into the per-req
|
||||
tables, so disagg + DSV4 is unsupported here (c-pages leak). Fixing requires
|
||||
calling these hooks from disagg's per-req alloc loop, or moving the write
|
||||
into the allocator itself.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||
|
||||
|
||||
def maybe_write_dsv4_extend(
|
||||
batch: ScheduleBatch,
|
||||
req_pool_indices_cpu: torch.Tensor,
|
||||
prefix_lens_cpu: torch.Tensor,
|
||||
seq_lens_cpu: torch.Tensor,
|
||||
) -> 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]``.
|
||||
|
||||
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.
|
||||
bundle = batch.out_cache_loc_dsv4
|
||||
if bundle is None:
|
||||
return
|
||||
|
||||
req_to_token_pool = batch.req_to_token_pool
|
||||
if not hasattr(req_to_token_pool, "write_c4"):
|
||||
return # non-DSV4 pool; skip defensively (shouldn't happen)
|
||||
|
||||
# SWA writes: prefix..seq token positions, one slot per raw token.
|
||||
_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,
|
||||
)
|
||||
|
||||
# c4 / c128 writes: prefix//ratio .. seq//ratio compressed positions.
|
||||
_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,
|
||||
prefix_lens_cpu,
|
||||
seq_lens_cpu,
|
||||
bundle.out_c128_loc,
|
||||
ratio=128,
|
||||
)
|
||||
|
||||
# c4_state / c128_state writes: tail-only. Bundle length is
|
||||
# sum(c{N}_state_alloc_len_i), NOT total raw extend tokens; each req's slots
|
||||
# go at raw positions [req.c{N}_state_alloc_offset, seq_len).
|
||||
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,
|
||||
[getattr(r, "c4_state_alloc_offset", 0) for r in batch.reqs],
|
||||
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,
|
||||
[getattr(r, "c128_state_alloc_offset", 0) for r in batch.reqs],
|
||||
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.
|
||||
|
||||
``seq_lens_cpu`` is the POST-decode seq len (already incremented by
|
||||
``token_per_req``); the new compressed tokens go at positions
|
||||
``[(old_seq) // ratio, (new_seq) // ratio)``.
|
||||
"""
|
||||
# Bundle stashed on batch.out_cache_loc_dsv4 by mem_cache/common.py;
|
||||
# None on CUDA / non-V4 paths → no-op.
|
||||
bundle = batch.out_cache_loc_dsv4
|
||||
if bundle is None:
|
||||
return
|
||||
|
||||
req_to_token_pool = batch.req_to_token_pool
|
||||
if not hasattr(req_to_token_pool, "write_c4"):
|
||||
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,
|
||||
prefix_lens_cpu,
|
||||
seq_lens_cpu,
|
||||
bundle.out_c128_loc,
|
||||
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 _write_per_req(
|
||||
write_fn,
|
||||
req_pool_indices_cpu: torch.Tensor,
|
||||
flat_loc: torch.Tensor,
|
||||
bounds_fn,
|
||||
) -> None:
|
||||
"""Distribute a flat ``[total_alloc]`` slot tensor across reqs.
|
||||
|
||||
``bounds_fn(i) -> (lo, hi)`` gives req i's write window; the matching
|
||||
``hi - lo`` slots are sliced off ``flat_loc`` in order and written via
|
||||
``write_fn((req_idx, slice(lo, hi)), values)``. flat_loc may be None /
|
||||
empty when the alloc path bypassed DSV4NPUTokenToKVPoolAllocator (e.g.
|
||||
page_size=1 or HiSparse wrapper); skip then.
|
||||
"""
|
||||
if flat_loc is None or flat_loc.numel() == 0:
|
||||
return
|
||||
pt = 0
|
||||
for i in range(req_pool_indices_cpu.shape[0]):
|
||||
lo, hi = bounds_fn(i)
|
||||
alloc_len = max(0, hi - lo)
|
||||
if alloc_len == 0:
|
||||
continue
|
||||
req_idx = int(req_pool_indices_cpu[i].item())
|
||||
chunk = flat_loc[pt : pt + alloc_len].to(torch.int32)
|
||||
write_fn((req_idx, slice(lo, hi)), chunk)
|
||||
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,
|
||||
prefix_lens_cpu: torch.Tensor,
|
||||
seq_lens_cpu: torch.Tensor,
|
||||
flat_loc: torch.Tensor,
|
||||
ratio: int,
|
||||
) -> None:
|
||||
"""Compressed-position write: req i's slots go at
|
||||
``[prefix_lens[i] // ratio, seq_lens[i] // ratio)``."""
|
||||
_write_per_req(
|
||||
write_fn,
|
||||
req_pool_indices_cpu,
|
||||
flat_loc,
|
||||
lambda i: (
|
||||
int(prefix_lens_cpu[i].item()) // ratio,
|
||||
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]
|
||||
state_allocator.free(free_slots.to(torch.int64))
|
||||
setattr(req, offset_attr, watermark)
|
||||
@@ -0,0 +1,584 @@
|
||||
"""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.
|
||||
|
||||
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.
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch_npu
|
||||
|
||||
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
|
||||
from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
|
||||
ONLINE_C128,
|
||||
DeepSeekV4IndexerPool,
|
||||
DeepSeekV4SingleKVPool,
|
||||
DeepSeekV4TokenToKVPool,
|
||||
)
|
||||
|
||||
|
||||
class NPUDeepSeekV4SingleKVPool(DeepSeekV4SingleKVPool):
|
||||
"""NPU bf16 variant of the full / SWA / c4 / c128 single-KV pool.
|
||||
|
||||
``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``.
|
||||
The CUDA fp8-packed-bytes layout (the base ``create_buffer``) is untouched.
|
||||
"""
|
||||
|
||||
def __init__(self, *args, kernel_page_size: int, **kwargs):
|
||||
# Set before super().__init__ — it calls _create_buffers() ->
|
||||
# create_buffer(), which reads self.kernel_page_size.
|
||||
self.kernel_page_size = kernel_page_size
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def create_buffer(self, *, num_pages: int):
|
||||
# Non-bf16 store dtype (shouldn't happen here) falls back to base layout.
|
||||
if self.store_dtype != torch.bfloat16:
|
||||
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.
|
||||
npu_num_pages = (self.size + self.kernel_page_size + 1) // self.kernel_page_size
|
||||
return torch.zeros(
|
||||
npu_num_pages,
|
||||
self.kernel_page_size,
|
||||
1,
|
||||
kv_dim,
|
||||
dtype=torch.bfloat16,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
|
||||
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.
|
||||
|
||||
``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.
|
||||
|
||||
NPU-only; CUDA keeps using the unchanged :class:`CompressStatePool`.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
size: int,
|
||||
overlap: bool,
|
||||
head_dim: int,
|
||||
dtype: torch.dtype,
|
||||
device: str,
|
||||
enable_memory_saver: bool,
|
||||
ratio: int,
|
||||
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
|
||||
)
|
||||
|
||||
# ``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.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
|
||||
)
|
||||
|
||||
# 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"))
|
||||
|
||||
|
||||
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
|
||||
``torch_npu.npu_scatter_nd_update_`` and read by
|
||||
``torch.ops.custom.npu_quant_lightning_indexer``.
|
||||
"""
|
||||
|
||||
def __init__(self, *args, kernel_page_size: int, **kwargs):
|
||||
# Set before super().__init__ — it calls _create_buffer().
|
||||
self._kernel_page_size = kernel_page_size
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def _create_buffer(self):
|
||||
# Base allocates the packed CUDA index_k_with_scale_buffer (kept for
|
||||
# get_contiguous_buf_infos / NSA compat); then add the NPU buffers.
|
||||
super()._create_buffer()
|
||||
kp = self._kernel_page_size
|
||||
npu_num_pages = (self.size + kp + 1) // kp
|
||||
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||||
self.index_k_buffer = [
|
||||
torch.zeros(
|
||||
npu_num_pages,
|
||||
kp,
|
||||
1,
|
||||
self.index_head_dim,
|
||||
dtype=torch.int8,
|
||||
device=self.device,
|
||||
)
|
||||
for _ in range(self.layer_num)
|
||||
]
|
||||
self.index_scale_buffer = [
|
||||
torch.zeros(
|
||||
npu_num_pages,
|
||||
kp,
|
||||
1,
|
||||
1,
|
||||
dtype=torch.float16,
|
||||
device=self.device,
|
||||
)
|
||||
for _ in range(self.layer_num)
|
||||
]
|
||||
|
||||
@property
|
||||
def has_npu_storage(self) -> bool:
|
||||
return True
|
||||
|
||||
def get_index_k(self, layer_id: int) -> torch.Tensor:
|
||||
return self.index_k_buffer[layer_id]
|
||||
|
||||
def get_index_scale(self, layer_id: int) -> torch.Tensor:
|
||||
return self.index_scale_buffer[layer_id]
|
||||
|
||||
def set_index_k_scale(
|
||||
self,
|
||||
layer_id: int,
|
||||
loc: torch.Tensor,
|
||||
index_k: torch.Tensor,
|
||||
index_k_scale: Optional[torch.Tensor],
|
||||
) -> None:
|
||||
# int8 K + fp16 scale come from _compressor_epilog_npu's npu_dynamic_quant
|
||||
# output (index_k: int8 [T, D], index_k_scale: fp16 [T, 1]).
|
||||
d = self.index_head_dim
|
||||
loc_long = loc.view(-1, 1).long()
|
||||
torch_npu.npu_scatter_nd_update_(
|
||||
self.index_k_buffer[layer_id].view(-1, 1, d),
|
||||
loc_long,
|
||||
index_k.to(torch.int8).view(-1, 1, d),
|
||||
)
|
||||
if index_k_scale is not None:
|
||||
torch_npu.npu_scatter_nd_update_(
|
||||
self.index_scale_buffer[layer_id].view(-1, 1, 1),
|
||||
loc_long,
|
||||
index_k_scale.to(torch.float16).view(-1, 1, 1),
|
||||
)
|
||||
|
||||
|
||||
class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
|
||||
"""NPU-only DSV4 KV pool with paged compress-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:`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
|
||||
``NotImplementedError`` for them (CUDA goes through the radix / store_cache
|
||||
accessors instead).
|
||||
"""
|
||||
|
||||
def _make_kv_pool(
|
||||
self,
|
||||
*,
|
||||
size: int,
|
||||
page_size: int,
|
||||
dtype: torch.dtype,
|
||||
layer_num: int,
|
||||
device: str,
|
||||
enable_memory_saver: bool,
|
||||
global_page_size: int,
|
||||
cls: type = DeepSeekV4SingleKVPool,
|
||||
) -> NPUDeepSeekV4SingleKVPool:
|
||||
# NPU does not use the HiSparse c4 device pool; fail loud if someone
|
||||
# enables it so the silent layout mismatch surfaces at init.
|
||||
assert cls is DeepSeekV4SingleKVPool, (
|
||||
"enable_hisparse is not supported on the NPU DSV4 KV pool "
|
||||
f"(got c4 pool class {cls.__name__})."
|
||||
)
|
||||
return NPUDeepSeekV4SingleKVPool(
|
||||
size,
|
||||
page_size,
|
||||
dtype,
|
||||
self.qk_nope_head_dim,
|
||||
self.qk_rope_head_dim,
|
||||
layer_num,
|
||||
device,
|
||||
enable_memory_saver,
|
||||
kernel_page_size=global_page_size,
|
||||
)
|
||||
|
||||
def _get_state_pool(self, layer_id: int, from_indexer: bool) -> CompressStatePool:
|
||||
"""Select this layer's attention vs c4-indexer compress-state pool.
|
||||
Wraps the community getters so the NPU port hooks below don't index the
|
||||
pool lists directly."""
|
||||
if from_indexer:
|
||||
return self.get_indexer_compress_states(layer_id)
|
||||
return self.get_attention_compress_states(layer_id)
|
||||
|
||||
def _make_attn_state_pool(
|
||||
self, ratio: int, enable_memory_saver: bool
|
||||
) -> NPUCompressStatePool:
|
||||
# ONLINE_C128 (CUDA-only) collapses the c128 ring to size 1; the NPU fused
|
||||
# compressor has no online mode, so assert the config mismatch early.
|
||||
assert not (ratio == 128 and ONLINE_C128), (
|
||||
"SGLANG_OPT_USE_ONLINE_COMPRESS is incompatible with the "
|
||||
"NPU fused compressor (no online mode in the kernel)."
|
||||
)
|
||||
return NPUCompressStatePool(
|
||||
size=self._state_pool_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,
|
||||
)
|
||||
|
||||
def _make_indexer_state_pool(
|
||||
self, ratio: int, enable_memory_saver: bool
|
||||
) -> NPUCompressStatePool:
|
||||
# c4 indexer shares the c4 state pool size budget but has its own
|
||||
# slot_dim (indexer_head_dim vs attention head_dim).
|
||||
return NPUCompressStatePool(
|
||||
size=self.c4_state_pool_size,
|
||||
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,
|
||||
)
|
||||
|
||||
def _make_indexer_pool(
|
||||
self,
|
||||
size: int,
|
||||
page_size: int,
|
||||
dtype: torch.dtype,
|
||||
index_head_dim: int,
|
||||
layer_num: int,
|
||||
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.
|
||||
return NPUDeepSeekV4IndexerPool(
|
||||
size,
|
||||
page_size,
|
||||
dtype,
|
||||
index_head_dim,
|
||||
layer_num,
|
||||
device,
|
||||
enable_memory_saver,
|
||||
kernel_page_size=self.page_size,
|
||||
)
|
||||
|
||||
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
|
||||
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
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Generic KV accessors (community base raises NotImplementedError; CUDA uses
|
||||
# store_cache). AscendAttnBackend reads KV through these, routed to the right
|
||||
# sub-pool by compression ratio.
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def get_key_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
item = self.layer_mapping[layer_id]
|
||||
ratio = item.compress_ratio
|
||||
if ratio == 0:
|
||||
return self.swa_kv_pool.kv_buffer[item.compress_layer_id]
|
||||
if ratio == 4:
|
||||
return self.c4_kv_pool.kv_buffer[item.compress_layer_id]
|
||||
if ratio == 128:
|
||||
return self.c128_kv_pool.kv_buffer[item.compress_layer_id]
|
||||
raise ValueError(f"unsupported compress_ratio={ratio} for get_key_buffer")
|
||||
|
||||
def get_value_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
# V4 uses MQA / latent attention — the K buffer doubles as V.
|
||||
return self.get_key_buffer(layer_id)
|
||||
|
||||
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
buf = self.get_key_buffer(layer_id)
|
||||
return buf, buf
|
||||
|
||||
def get_swa_buffer(
|
||||
self, layer_id: int, loc: Optional[torch.Tensor] = None
|
||||
) -> torch.Tensor:
|
||||
"""Return the SWA layer's KV cache in PA_ND layout
|
||||
(num_pages, page_size, num_kv_heads=1, dim). When ``loc`` is given,
|
||||
flatten across (num_pages, page_size) and gather the matching tokens —
|
||||
shape becomes (num_tokens, 1, dim).
|
||||
"""
|
||||
# Index by RAW layer_id, not compress_layer_id (a per-bucket counter that
|
||||
# would collide across ratios). swa_kv_pool is sized layer_num=total_layers.
|
||||
kv = self.swa_kv_pool.kv_buffer[layer_id]
|
||||
if loc is not None:
|
||||
kv = kv.flatten(0, 1)[loc]
|
||||
return kv
|
||||
|
||||
def get_compress_buffer(
|
||||
self,
|
||||
layer_id: int,
|
||||
from_indexer: bool = False,
|
||||
loc: Optional[torch.Tensor] = None,
|
||||
) -> Optional[torch.Tensor]:
|
||||
"""Return the compressed KV buffer for a c4 / c128 layer.
|
||||
|
||||
Routes to c4 / c128 kv_pool by layer compression ratio. Returns
|
||||
``None`` for ratio == 0 (no compress KV exists). The
|
||||
from_indexer=True branch returns the dedicated int8 K buffer that
|
||||
``torch.ops.custom.npu_quant_lightning_indexer`` consumes.
|
||||
"""
|
||||
item = self.layer_mapping[layer_id]
|
||||
if item.compress_ratio == 4:
|
||||
if from_indexer:
|
||||
kv = self.c4_indexer_kv_pool.get_index_k(item.compress_layer_id)
|
||||
else:
|
||||
kv = self.c4_kv_pool.kv_buffer[item.compress_layer_id]
|
||||
elif item.compress_ratio == 128:
|
||||
assert not from_indexer, "c128 has no indexer pool"
|
||||
kv = self.c128_kv_pool.kv_buffer[item.compress_layer_id]
|
||||
else:
|
||||
return None
|
||||
if loc is not None:
|
||||
kv = kv.flatten(0, 1)[loc]
|
||||
return kv
|
||||
|
||||
def set_swa_buffer(
|
||||
self,
|
||||
layer_id: int,
|
||||
loc: torch.Tensor,
|
||||
cache: torch.Tensor,
|
||||
) -> None:
|
||||
"""Write ``cache`` into the SWA pool at flat token positions ``loc``.
|
||||
|
||||
``cache`` shape: (num_tokens, num_kv_heads=1, dim). The buffer view is
|
||||
(num_pages, page_size, 1, dim) so we flatten the first two dims and
|
||||
index_put.
|
||||
"""
|
||||
# Index by raw layer_id (see get_swa_buffer) to avoid bucket collision.
|
||||
buf = self.swa_kv_pool.kv_buffer[layer_id]
|
||||
buf_flat = buf.flatten(0, 1) # (num_pages * page_size, 1, dim)
|
||||
# Caller (V4 MQALayer) may hand us cache shaped (T, dim); the buffer has
|
||||
# an explicit num_kv_heads=1 axis, so insert it.
|
||||
if cache.ndim == buf_flat.ndim - 1:
|
||||
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(
|
||||
self,
|
||||
layer_id: int,
|
||||
loc: torch.Tensor,
|
||||
kv: torch.Tensor,
|
||||
score: torch.Tensor,
|
||||
from_indexer: bool,
|
||||
) -> 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
|
||||
|
||||
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
|
||||
|
||||
def set_compress_buffer(
|
||||
self,
|
||||
layer_id: int,
|
||||
loc: torch.Tensor,
|
||||
kv: torch.Tensor,
|
||||
kv_scale: Optional[torch.Tensor],
|
||||
from_indexer: bool,
|
||||
) -> None:
|
||||
# Routes to c4_indexer (from_indexer) / c4_kv (ratio 4) / c128_kv (ratio
|
||||
# 128). NPU bypasses CUDA fused_store_cache with direct bf16 writes.
|
||||
ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||
device_type = kv.device.type
|
||||
if from_indexer:
|
||||
assert ratio == 4, f"indexer only on c4 layers, got ratio={ratio}"
|
||||
if device_type == "npu":
|
||||
assert (
|
||||
self.c4_indexer_kv_pool.has_npu_storage
|
||||
), "NPU index buffers not allocated — pool was init'd on CUDA?"
|
||||
self.c4_indexer_kv_pool.set_index_k_scale(
|
||||
compress_layer_id, loc, kv, kv_scale
|
||||
)
|
||||
return
|
||||
if kv_scale is None:
|
||||
self.c4_indexer_kv_pool.set_index_fused(compress_layer_id, loc, kv)
|
||||
return
|
||||
self.c4_indexer_kv_pool.set_index_k_scale_buffer(
|
||||
compress_layer_id, loc, kv, kv_scale
|
||||
)
|
||||
return
|
||||
compress_pool = self.c4_kv_pool if ratio == 4 else self.c128_kv_pool
|
||||
if device_type == "npu":
|
||||
# PA_ND layout: kv_buffer[layer_id] shape = (num_pages, page_size,
|
||||
# 1, kv_dim). Flatten (num_pages, page_size) and index by `loc`.
|
||||
buf = compress_pool.kv_buffer[compress_layer_id]
|
||||
buf_flat = buf.flatten(0, 1)
|
||||
kv_view = kv.to(buf_flat.dtype)
|
||||
if kv_view.ndim == buf_flat.ndim - 1:
|
||||
kv_view = kv_view.unsqueeze(1)
|
||||
buf_flat[loc] = kv_view
|
||||
return
|
||||
compress_pool.set_key_buffer_fused(compress_layer_id, loc, kv)
|
||||
|
||||
def get_compress_dequant_scale_buffer(
|
||||
self,
|
||||
layer_id: int,
|
||||
from_indexer: bool,
|
||||
) -> torch.Tensor:
|
||||
# Returns the float16 dequant scale buffer (NPU indexer pool's dedicated
|
||||
# scale buffer alongside the int8 K buffer).
|
||||
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."
|
||||
)
|
||||
@@ -0,0 +1,128 @@
|
||||
"""DSV4-NPU per-request mapping pool.
|
||||
|
||||
Subclass of ``ReqToTokenPool`` that adds five auxiliary per-request tables
|
||||
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
|
||||
|
||||
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``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
|
||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||
|
||||
|
||||
class DSV4NPUReqToTokenPool(ReqToTokenPool):
|
||||
"""ReqToTokenPool extended with DSV4 SWA + c4/c128 per-req tables.
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size: int,
|
||||
max_context_len: int,
|
||||
device: str,
|
||||
enable_memory_saver: bool,
|
||||
):
|
||||
super().__init__(size, max_context_len, device, enable_memory_saver)
|
||||
|
||||
memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
enable=enable_memory_saver
|
||||
)
|
||||
|
||||
# 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.
|
||||
self._dsv4_allocator = None
|
||||
|
||||
# (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).
|
||||
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
|
||||
|
||||
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
|
||||
|
||||
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. This is a
|
||||
one-way ref (pool -> allocator). The reverse direction (the allocator
|
||||
reading these per-req tables for its c-pool / state last_loc lookup) is
|
||||
no longer a stored back-ref: mem_cache/common.py passes this pool into
|
||||
``alloc_extend`` / ``alloc_decode`` per call instead."""
|
||||
self._dsv4_allocator = allocator
|
||||
|
||||
def free(self, req):
|
||||
# Trigger c4/c128 free via the allocator's unified free 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)
|
||||
super().free(req)
|
||||
@@ -35,7 +35,11 @@ from typing import TYPE_CHECKING, Dict, Optional, Union
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from sglang.srt.configs.model_config import AttentionArch, is_deepseek_dsa
|
||||
from sglang.srt.configs.model_config import (
|
||||
AttentionArch,
|
||||
is_deepseek_dsa,
|
||||
is_deepseek_v4,
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.model_executor.runner import DecodeCudaGraphRunner
|
||||
@@ -231,7 +235,10 @@ class NPUGraphRunner(DecodeCudaGraphRunner):
|
||||
|
||||
graph_key = self._make_graph_key(self.bs)
|
||||
|
||||
if not is_deepseek_dsa(self.model_runner.model_config.hf_config):
|
||||
if not (
|
||||
is_deepseek_dsa(self.model_runner.model_config.hf_config)
|
||||
or is_deepseek_v4(self.model_runner.model_config.hf_config)
|
||||
):
|
||||
if forward_batch.forward_mode.is_target_verify():
|
||||
seq_lens_cpu = forward_batch.seq_lens.cpu() + self.num_tokens_per_bs
|
||||
seq_lens = seq_lens_cpu.tolist() + [0] * (self.bs - self.raw_bs)
|
||||
|
||||
@@ -41,6 +41,26 @@ def fused_topk_npu(
|
||||
)
|
||||
topk_weights = topk_weights.to(torch.float32)
|
||||
|
||||
# sqrtsoftplus (DSV4 noaux_tc): the NPU op only scores sigmoid/softmax, so use
|
||||
# a torch path. top-k over (scores + bias); weights from un-biased scores.
|
||||
elif topk_config.scoring_func == "sqrtsoftplus":
|
||||
scores = torch.nn.functional.softplus(router_logits.float()).sqrt()
|
||||
scores_for_choice = (
|
||||
scores + correction_bias.unsqueeze(0).float()
|
||||
if correction_bias is not None
|
||||
else scores
|
||||
)
|
||||
_, topk_ids = torch.topk(
|
||||
scores_for_choice, k=topk_config.top_k, dim=-1, sorted=False
|
||||
)
|
||||
topk_ids = topk_ids.to(torch.int32)
|
||||
topk_weights = scores.gather(1, topk_ids)
|
||||
if renormalize:
|
||||
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||
else:
|
||||
topk_weights = topk_weights * topk_config.routed_scaling_factor
|
||||
topk_weights = topk_weights.to(torch.float32)
|
||||
|
||||
# Support grouped top-k or correction bias or sigmoid or routed_scaling_factor
|
||||
elif (
|
||||
correction_bias is not None
|
||||
@@ -66,6 +86,7 @@ def fused_topk_npu(
|
||||
),
|
||||
eps=float(1e-20),
|
||||
)
|
||||
topk_weights = topk_weights.to(torch.float32)
|
||||
|
||||
# torch native is not yet supported num_token_non_padded
|
||||
# Fallback to torch native implementation
|
||||
|
||||
@@ -99,7 +99,12 @@ def init_npu_backend():
|
||||
|
||||
assert _is_npu, "NPU backend initialization called on non-NPU device."
|
||||
|
||||
import sgl_kernel_npu # noqa: F401
|
||||
try:
|
||||
import custom_ops # noqa: F401
|
||||
import sgl_kernel_npu # noqa: F401
|
||||
except ImportError as e:
|
||||
logger.warning("NPU custom kernel packages unavailable: %s", e)
|
||||
|
||||
import torch_npu
|
||||
from torch_npu.contrib import transfer_to_npu # noqa: F401
|
||||
|
||||
|
||||
@@ -6,9 +6,11 @@ from sglang.srt.configs.linear_attn_model_registry import (
|
||||
get_linear_attn_config,
|
||||
import_backend_class,
|
||||
)
|
||||
from sglang.srt.utils import get_device_capability, is_musa
|
||||
from sglang.srt.utils import get_device_capability, is_hip, is_musa, is_npu
|
||||
|
||||
_is_musa = is_musa()
|
||||
_is_npu = is_npu()
|
||||
_is_hip = is_hip()
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -126,9 +128,13 @@ def _create_nsa_compat(runner):
|
||||
|
||||
@register_attention_backend("dsv4")
|
||||
def create_dsv4_backend(runner):
|
||||
from sglang.srt.utils import is_hip
|
||||
if _is_npu:
|
||||
from sglang.srt.hardware_backend.npu.attention.ascend_dsv4_backend import (
|
||||
DeepseekV4AscendAttnBackend,
|
||||
)
|
||||
|
||||
if is_hip():
|
||||
return DeepseekV4AscendAttnBackend(runner)
|
||||
elif _is_hip:
|
||||
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
|
||||
DeepseekV4HipRadixBackend,
|
||||
)
|
||||
|
||||
@@ -19,17 +19,21 @@ from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp
|
||||
from sglang.srt.layers.attention.dsv4.quant_k_cache import (
|
||||
quant_to_nope_fp8_rope_bf16_pack_triton,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import get_attention_cp_size
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import ReplicatedLinear
|
||||
from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output
|
||||
from sglang.srt.layers.utils.multi_platform import MultiPlatformOp
|
||||
from sglang.srt.mem_cache.deepseek_v4_compress_state import (
|
||||
CompressStatePool,
|
||||
)
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||
from sglang.srt.models.deepseek_v2 import _is_hip
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix, get_bool_env_var, set_weight_attrs
|
||||
from sglang.srt.utils import add_prefix, get_bool_env_var, is_npu, set_weight_attrs
|
||||
|
||||
_is_npu = is_npu()
|
||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||
_tgemm = None
|
||||
if _use_aiter:
|
||||
@@ -341,7 +345,7 @@ def create_paged_compressor_data(
|
||||
return FusedCompressMetadata(write_loc=write_loc, extra_data=extra_data, plan=plan)
|
||||
|
||||
|
||||
class Compressor(nn.Module):
|
||||
class Compressor(MultiPlatformOp):
|
||||
def __init__(
|
||||
self,
|
||||
config: DeepSeekV4Config,
|
||||
@@ -390,6 +394,9 @@ class Compressor(nn.Module):
|
||||
def _apply_ape_hotfix(self):
|
||||
self.ape_converted = True
|
||||
|
||||
if _is_npu:
|
||||
return
|
||||
|
||||
if self.overlap:
|
||||
ape = torch.chunk(self.ape.data, 2, dim=-1)
|
||||
ape = torch.cat([ape[0], ape[1]], dim=0)
|
||||
@@ -433,11 +440,11 @@ class Compressor(nn.Module):
|
||||
)
|
||||
return kv_score
|
||||
|
||||
def forward(
|
||||
def forward_native(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
attn_backend: AttentionBackend,
|
||||
attn_backend: Optional[AttentionBackend] = None,
|
||||
) -> torch.Tensor:
|
||||
if forward_batch.forward_mode.is_idle():
|
||||
assert x.shape[0] == 0
|
||||
@@ -461,6 +468,26 @@ class Compressor(nn.Module):
|
||||
is_paged=True,
|
||||
)
|
||||
|
||||
def forward_npu(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
attn_backend: Optional[AttentionBackend] = None,
|
||||
) -> torch.Tensor:
|
||||
if forward_batch.forward_mode.is_idle():
|
||||
assert x.shape[0] == 0
|
||||
return x.new_empty(0, self.head_dim)
|
||||
|
||||
if dsa_use_prefill_cp(forward_batch):
|
||||
x = cp_all_gather_rerange_output(
|
||||
x,
|
||||
get_attention_cp_size(),
|
||||
forward_batch,
|
||||
torch.cuda.current_stream(),
|
||||
)
|
||||
|
||||
return get_attn_backend().forward_compress(self, x, forward_batch)
|
||||
|
||||
|
||||
if _is_hip and not envs.SGLANG_OPT_USE_COMPRESSOR_V2.get():
|
||||
from sglang.srt.layers.attention.dsv4.compress_hip import ( # noqa: F811
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
import math
|
||||
from functools import lru_cache
|
||||
from typing import Optional
|
||||
@@ -6,16 +7,29 @@ import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# tilelang isn't shipped on every platform (e.g. Ascend NPU images) and the
|
||||
# only tilelang artifacts in this file are pass_configs that downstream
|
||||
# tilelang.jit decorators would consume — the kernels actually defined here
|
||||
# are Triton. Keep the import optional so this module loads on NPU.
|
||||
try:
|
||||
import tilelang
|
||||
|
||||
tilelang.set_log_level("WARNING")
|
||||
|
||||
pass_configs = {
|
||||
tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
|
||||
tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
|
||||
}
|
||||
except ImportError:
|
||||
pass
|
||||
logger.info(
|
||||
"tilelang not installed; deepseek_v4_rope pass_configs unset. "
|
||||
"Triton kernels in this module still run; only downstream tilelang.jit "
|
||||
"consumers of pass_configs will need to handle the None."
|
||||
)
|
||||
tilelang = None
|
||||
pass_configs = None
|
||||
|
||||
FP8 = "float8_e4m3"
|
||||
BF16 = "bfloat16"
|
||||
@@ -23,9 +37,21 @@ FP32 = "float32"
|
||||
INT32 = "int32"
|
||||
|
||||
|
||||
def _yarn_get_mscale(scale: float = 1.0, mscale: float = 1.0) -> float:
|
||||
if scale <= 1:
|
||||
return 1.0
|
||||
return 0.1 * mscale * math.log(scale) + 1.0
|
||||
|
||||
|
||||
@lru_cache(2)
|
||||
def precompute_freqs_cis(
|
||||
dim, seqlen, original_seq_len, base, factor, beta_fast, beta_slow
|
||||
dim,
|
||||
seqlen,
|
||||
original_seq_len,
|
||||
base,
|
||||
factor,
|
||||
beta_fast,
|
||||
beta_slow,
|
||||
) -> torch.Tensor:
|
||||
|
||||
def find_correction_dim(num_rotations, dim, base, max_seq_len):
|
||||
@@ -434,3 +460,123 @@ def fused_norm_rope_inplace_triton(
|
||||
HAS_WEIGHT=(weight is not None),
|
||||
USE_POS=(positions is not None),
|
||||
)
|
||||
|
||||
|
||||
# Cache contiguous real/imag halves of each freqs_cis (its .real/.imag are
|
||||
# strided views, stride=2 on the interleaved layout), keyed by id.
|
||||
_NPU_ROPE_CONTIG_CACHE: dict[int, tuple] = {}
|
||||
|
||||
|
||||
def _get_contig_freqs_real_imag(
|
||||
freqs_cis: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Return contiguous (real, imag) halves of ``freqs_cis``, cached by id.
|
||||
|
||||
Used by NPU rope paths to avoid the per-call StridedSlice materialization
|
||||
triggered by aclnnIndex over the strided ``.real`` / ``.imag`` views of
|
||||
the complex ``freqs_cis`` buffer. First call per freqs_cis pays the
|
||||
contiguous() once; later calls reuse the cached tensors.
|
||||
|
||||
All callers within a single MQALayer (outer rope, indexer inner rope,
|
||||
compressor epilog rope) get the same freqs_cis instance, so each layer
|
||||
materializes at most one (real, imag) pair.
|
||||
"""
|
||||
cache_key = id(freqs_cis)
|
||||
cached = _NPU_ROPE_CONTIG_CACHE.get(cache_key)
|
||||
if cached is None:
|
||||
cached = (freqs_cis.real.contiguous(), freqs_cis.imag.contiguous())
|
||||
_NPU_ROPE_CONTIG_CACHE[cache_key] = cached
|
||||
return cached
|
||||
|
||||
|
||||
def get_fused_compressor_rope_cos_sin(
|
||||
freqs_cis: torch.Tensor,
|
||||
positions_cmp: torch.Tensor,
|
||||
dtype: torch.dtype,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Build (cos, sin) tensors shaped ``[T, rope_head_dim]`` for the fused
|
||||
compressor op (``torch.ops.custom.compressor``).
|
||||
|
||||
The op consumes ``rope_cos`` / ``rope_sin`` of shape
|
||||
``[min(T, T//cmp_ratio + B), rope_head_dim]`` in bf16/fp16. We index
|
||||
the cached contig real/imag halves of the complex ``freqs_cis`` and
|
||||
interleave-double the last dim to match the kernel's expected layout
|
||||
(matches dsv4_release ``ComplexExpRotaryEmbedding.cos_cache``, which
|
||||
is built as ``complex_cache.real.repeat_interleave(2, dim=-1)``).
|
||||
|
||||
Safe to call from inside a captured aclgraph: both ``index_select`` and
|
||||
``repeat_interleave`` over a graph-input ``positions_cmp`` of fixed
|
||||
capture-time shape produce static-shape outputs. Identical to what the
|
||||
existing inplace_partial_rotary_mul fallback does at
|
||||
:func:`v4_rope_inplace_npu`, just without the inverse / 4D-view step.
|
||||
"""
|
||||
real_contig, imag_contig = _get_contig_freqs_real_imag(freqs_cis)
|
||||
cos_half = real_contig.index_select(0, positions_cmp)
|
||||
sin_half = imag_contig.index_select(0, positions_cmp)
|
||||
cos = cos_half.repeat_interleave(2, dim=-1).to(dtype)
|
||||
sin = sin_half.repeat_interleave(2, dim=-1).to(dtype)
|
||||
return cos, sin
|
||||
|
||||
|
||||
def v4_rope_inplace_npu(
|
||||
q_rope: torch.Tensor,
|
||||
kv_rope: Optional[torch.Tensor],
|
||||
freqs_cis: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
inverse: bool = False,
|
||||
) -> None:
|
||||
"""In-place interleaved RoPE for V4 — torch fallback used on NPU.
|
||||
|
||||
Mirrors main's CUDA `fused_rope` kernel: consecutive (even, odd) pairs
|
||||
of x form complex pairs, with `freqs_cis` a complex tensor where
|
||||
`freqs_cis.real[t, k]` = cos(theta_{t,k}), `freqs_cis.imag` = sin(...)
|
||||
indexed by frequency pair k in [0, rope_dim/2).
|
||||
|
||||
NOTE on V4-Flash YARN `mscale`: when the model was trained with the
|
||||
YARN magnitude-scale `mscale` ≠ 1.0, the cos/sin values stored in
|
||||
`freqs_cis` MUST already be pre-multiplied by `mscale` at precompute
|
||||
time — see `precompute_freqs_cis`. This function
|
||||
just reads what's stored; it does NOT apply mscale here.
|
||||
|
||||
Prefer the NPU-native `torch.ops.custom.inplace_partial_rotary_mul`:
|
||||
the torch fallback differs by ~1 ULP per element vs the kernel because
|
||||
torch does bf16*bf16 muls with bf16 accumulation while the NPU kernel
|
||||
accumulates in fp32; 43 layers × (Q + K) = 86 rope calls compound that
|
||||
drift enough to flip argmax on marginal prompts.
|
||||
"""
|
||||
# Build cos/sin caches in the kernel's expected (T, 1, 1, rope_dim) layout,
|
||||
# each freq value repeated twice for the interleaved pairing convention.
|
||||
freqs_real_contig, freqs_imag_contig = _get_contig_freqs_real_imag(freqs_cis)
|
||||
cos_half = freqs_real_contig[positions] # (T, rope_dim/2)
|
||||
sin_half = freqs_imag_contig[positions]
|
||||
if inverse:
|
||||
sin_half = -sin_half
|
||||
cos_full = cos_half.repeat_interleave(2, dim=-1).to(q_rope.dtype)
|
||||
sin_full = sin_half.repeat_interleave(2, dim=-1).to(q_rope.dtype)
|
||||
rope_dim = cos_full.shape[-1]
|
||||
# repeat_interleave produces a contiguous tensor, so the .view()
|
||||
# below already returns a contiguous result — no .contiguous() needed.
|
||||
cos4 = cos_full.view(-1, 1, 1, rope_dim)
|
||||
sin4 = sin_full.view(-1, 1, 1, rope_dim)
|
||||
# q_rope: (T, n_heads, rope_dim) → (T, 1, n_heads, rope_dim) view
|
||||
# kv_rope: (T, 1, rope_dim) → (T, 1, 1, rope_dim) view
|
||||
q_view = q_rope.unsqueeze(1)
|
||||
torch.ops.custom.inplace_partial_rotary_mul(
|
||||
q_view,
|
||||
cos4,
|
||||
sin4,
|
||||
rotary_mode="interleave",
|
||||
partial_slice=[0, rope_dim],
|
||||
)
|
||||
if kv_rope is not None:
|
||||
if kv_rope.dim() == 3:
|
||||
kv_view = kv_rope.unsqueeze(1)
|
||||
else:
|
||||
kv_view = kv_rope.view(-1, 1, 1, rope_dim)
|
||||
torch.ops.custom.inplace_partial_rotary_mul(
|
||||
kv_view,
|
||||
cos4,
|
||||
sin4,
|
||||
rotary_mode="interleave",
|
||||
partial_slice=[0, rope_dim],
|
||||
)
|
||||
|
||||
@@ -3,8 +3,6 @@ import logging
|
||||
import math
|
||||
from typing import Tuple
|
||||
|
||||
import tilelang
|
||||
import tilelang.language as T
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.utils import is_arch_support_pdl
|
||||
@@ -14,15 +12,55 @@ from sglang.srt.layers.utils.common import strict_contiguous
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
tilelang.set_log_level("WARNING")
|
||||
# Tilelang isn't packaged on every platform (notably Ascend NPU images) but
|
||||
# this module is imported transitively from deepseek_v4.py — module-load
|
||||
# must succeed even when tilelang is missing. The kernels themselves still
|
||||
# require tilelang at runtime; we replace the package with a stub that lets
|
||||
# `@tilelang.jit` decorations and `tilelang.PassConfigKey.*` references parse
|
||||
# without ImportError, and any actual call into the kernels raises a clear
|
||||
# message at execution time instead of crashing on import.
|
||||
try:
|
||||
import tilelang
|
||||
import tilelang.language as T
|
||||
|
||||
# Set once mhc_pre() has compiled every n_splits bucket at startup.
|
||||
_mhc_pre_warmed = False
|
||||
tilelang.set_log_level("WARNING")
|
||||
|
||||
pass_configs = {
|
||||
tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
|
||||
tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
|
||||
}
|
||||
# Set once mhc_pre() has compiled every n_splits bucket at startup.
|
||||
_mhc_pre_warmed = False
|
||||
|
||||
pass_configs = {
|
||||
tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
|
||||
tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
|
||||
}
|
||||
except ImportError:
|
||||
|
||||
class _TilelangMissing:
|
||||
"""Stub so module-level @tilelang.jit and PassConfigKey accesses parse."""
|
||||
|
||||
def __getattr__(self, name):
|
||||
if name == "jit":
|
||||
|
||||
def _jit(*_args, **_kwargs):
|
||||
def _wrap(fn):
|
||||
def _raise(*a, **k):
|
||||
raise RuntimeError(
|
||||
"tilelang is not installed; this kernel cannot run "
|
||||
"on the current platform"
|
||||
)
|
||||
|
||||
return _raise
|
||||
|
||||
return _wrap
|
||||
|
||||
return _jit
|
||||
return _TilelangMissing()
|
||||
|
||||
def __call__(self, *_args, **_kwargs):
|
||||
return _TilelangMissing()
|
||||
|
||||
tilelang = _TilelangMissing()
|
||||
T = _TilelangMissing()
|
||||
pass_configs = None
|
||||
|
||||
FP8 = "float8_e4m3"
|
||||
BF16 = "bfloat16"
|
||||
@@ -1515,3 +1553,60 @@ def mhc_fused_post_pre(
|
||||
comb_mix_cur.view(*outer_shape, hc_mult, hc_mult),
|
||||
layer_input_cur.view(*outer_shape, hidden_size),
|
||||
)
|
||||
|
||||
|
||||
def npu_hc_pre(
|
||||
x: torch.Tensor,
|
||||
hc_fn: torch.Tensor,
|
||||
hc_scale: torch.Tensor,
|
||||
hc_base: torch.Tensor,
|
||||
hc_mult: int,
|
||||
hc_sinkhorn_iters: int,
|
||||
rms_norm_eps: float,
|
||||
hc_eps: float,
|
||||
forward_batch=None,
|
||||
) -> tuple:
|
||||
"""NPU-accelerated hc_pre via the custom_ops kernel.
|
||||
|
||||
Returns (y, post, comb, norm_fused). norm_fused is always False
|
||||
because npu_hc_pre does not fold input_layernorm — the caller must
|
||||
apply it separately.
|
||||
"""
|
||||
shape, dtype = x.size(), x.dtype
|
||||
|
||||
# IDLE / empty short-circuit, mirroring the dsv4-flash source.
|
||||
# The kernel emits post/comb in fp32 (sinkhorn iterates in fp32),
|
||||
# so the dummies must too — otherwise downstream comb/post-aware
|
||||
# ops see a silent fp32 ↔ bf16 split between idle and non-idle
|
||||
# batches.
|
||||
is_idle = forward_batch is not None and forward_batch.forward_mode.is_idle()
|
||||
if is_idle or x.shape[0] == 0:
|
||||
bs = x.shape[0]
|
||||
y = torch.empty((bs, shape[-1]), dtype=dtype, device=x.device)
|
||||
post = torch.empty((bs, hc_mult), dtype=torch.float32, device=x.device)
|
||||
comb = torch.empty(
|
||||
(bs, hc_mult, hc_mult),
|
||||
dtype=torch.float32,
|
||||
device=x.device,
|
||||
)
|
||||
return y, post, comb, False
|
||||
|
||||
# Note the return order: (y, post, comb) — y is the (T, hidden)
|
||||
# mixed activation, post / comb are the hc_post inputs. The
|
||||
# fused kernel emits y in fp32 (sinkhorn iterates in fp32), so
|
||||
# cast back to the input dtype before the downstream
|
||||
# aclnnRmsNorm (which has no x=fp32 / gamma=bf16 overload).
|
||||
y, post, comb = torch.ops.custom.npu_hc_pre(
|
||||
x,
|
||||
hc_fn,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
hc_mult=hc_mult,
|
||||
hc_sinkhorn_iters=hc_sinkhorn_iters,
|
||||
norm_eps=rms_norm_eps,
|
||||
hc_eps=hc_eps,
|
||||
)
|
||||
# npu_hc_pre uses norm_eps for sinkhorn's internal RMS only; it does
|
||||
# not fold input_layernorm. Return norm_fused=False so the caller
|
||||
# applies the layernorm itself, matching the deepgemm/torch paths.
|
||||
return y.to(dtype), post, comb, False
|
||||
|
||||
@@ -19,10 +19,13 @@ from sglang.srt.layers.moe.topk import (
|
||||
_mask_topk_ids_padded_region,
|
||||
_zero_topk_weights_padded_region,
|
||||
)
|
||||
from sglang.srt.utils import is_hip
|
||||
from sglang.srt.utils import is_hip, is_npu
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_is_hip = is_hip()
|
||||
_is_npu = is_npu()
|
||||
|
||||
|
||||
class HashTopK(nn.Module):
|
||||
def __init__(
|
||||
@@ -182,8 +185,7 @@ class HashTopK(nn.Module):
|
||||
)
|
||||
else:
|
||||
topk_weights, topk_ids = self._forward_torch(router_logits, input_ids)
|
||||
|
||||
if is_hip():
|
||||
if _is_hip or _is_npu:
|
||||
topk_weights = topk_weights.to(torch.float32)
|
||||
|
||||
log2phy_prob = None
|
||||
|
||||
@@ -1165,6 +1165,13 @@ def _mask_topk_ids_padded_region(
|
||||
# TODO: let the kernel support other dtypes
|
||||
if _is_cuda and topk_ids.dtype == torch.int32 and fill_value == -1:
|
||||
mask_topk_ids(topk_ids, num_token_non_padded)
|
||||
elif _is_npu:
|
||||
# On NPU, bool-indexed scatter `topk_ids[bool_mask, :] = -1` lowers
|
||||
# to aclnnNonzeroV2 and can trigger an aicore timeout under long
|
||||
# workloads; `torch.where` avoids that nonzero scan.
|
||||
indices = torch.arange(0, topk_ids.shape[0], device=topk_ids.device)
|
||||
mask = (indices >= num_token_non_padded).unsqueeze(-1)
|
||||
topk_ids = torch.where(mask, torch.full_like(topk_ids, -1), topk_ids)
|
||||
else:
|
||||
indices = torch.arange(0, topk_ids.shape[0], device=topk_ids.device)
|
||||
topk_ids[indices >= num_token_non_padded, :] = fill_value
|
||||
|
||||
+20
@@ -132,3 +132,23 @@ class NPUCompressedTensorsW8A8Int8DynamicMoE(CompressedTensorsMoEScheme):
|
||||
) -> CombineInput:
|
||||
|
||||
return self.kernel.apply(layer, dispatch_output)
|
||||
|
||||
def apply_without_routing_weights(
|
||||
self,
|
||||
layer,
|
||||
hidden_states,
|
||||
hidden_states_scale,
|
||||
group_list_type,
|
||||
group_list,
|
||||
output_dtype,
|
||||
):
|
||||
# NPU MoE bypasses MoeRunner: expose the kernel's existing
|
||||
# apply_without_routing_weights directly through the scheme.
|
||||
return self.kernel.apply_without_routing_weights(
|
||||
layer,
|
||||
hidden_states,
|
||||
hidden_states_scale,
|
||||
group_list_type,
|
||||
group_list,
|
||||
output_dtype,
|
||||
)
|
||||
|
||||
@@ -88,6 +88,17 @@ class ModelSlimConfig(QuantizationConfig):
|
||||
|
||||
def __init__(self, quant_config: Dict[str, Any] = {}):
|
||||
super().__init__()
|
||||
keys = [k for k in quant_config if isinstance(k, str)]
|
||||
is_dsv4 = any(k.startswith("hc_head_") for k in keys)
|
||||
if is_dsv4:
|
||||
from sglang.srt.models.deepseek_v4 import DeepseekV4ForCausalLM
|
||||
|
||||
remap = DeepseekV4ForCausalLM.remap_weight_name_to_dpsk_hf_format
|
||||
quant_config = {
|
||||
(remap(k) if isinstance(k, str) else k): v
|
||||
for k, v in quant_config.items()
|
||||
}
|
||||
|
||||
self.quant_description = quant_config
|
||||
ignore = cast(List[str], quant_config.get("ignore", []))
|
||||
self.ignore = ignore if ignore is not None else []
|
||||
|
||||
@@ -68,6 +68,9 @@ from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationM
|
||||
from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank
|
||||
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,
|
||||
@@ -1743,6 +1746,9 @@ 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``.
|
||||
out_cache_loc_dsv4: Optional[Any] = None
|
||||
|
||||
# For hybrid GDN prefix cache
|
||||
mamba_track_indices: torch.Tensor = None # shape: [b], int64
|
||||
@@ -2624,7 +2630,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
if self.model_config.is_encoder_decoder:
|
||||
self.prepare_encoder_info_decode()
|
||||
|
||||
# Allocate memory
|
||||
# Allocate memory (DSV4-NPU c{4,128}_state alloc lens are computed inside
|
||||
# the allocator, triggered from mem_cache/common.py.)
|
||||
self.out_cache_loc = alloc_for_decode(self, token_per_req=1)
|
||||
|
||||
# Update req-level memory management fields
|
||||
@@ -2887,6 +2894,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
if req.decode_batch_idx % eviction_interval == 1:
|
||||
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
|
||||
|
||||
@@ -6,6 +6,11 @@ from typing import TYPE_CHECKING, Optional
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
|
||||
maybe_evict_dsv4_state_on_swa,
|
||||
maybe_write_dsv4_decode,
|
||||
maybe_write_dsv4_extend,
|
||||
)
|
||||
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
|
||||
@@ -21,14 +26,17 @@ from sglang.srt.mem_cache.triton_ops.common import (
|
||||
write_req_to_token_pool_triton,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs, get_global_server_args
|
||||
from sglang.srt.utils import is_hip, support_triton
|
||||
from sglang.srt.utils import is_hip, is_npu, support_triton
|
||||
from sglang.srt.utils.common import ceil_align, is_pin_memory_available
|
||||
|
||||
_is_npu = is_npu()
|
||||
|
||||
_is_hip = is_hip()
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.model_executor.forward_batch_info import DSV4StateLens
|
||||
|
||||
# Needs 2 + 1 slots for mamba request with prefix cache. 2 for ping pong cache, 1 for running mamba state.
|
||||
MAMBA_STATE_PER_REQ_PREFIX_CACHE = 3
|
||||
@@ -95,6 +103,9 @@ def free_swa_out_of_window_slots(
|
||||
req.req_pool_idx, req.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.swa_evicted_seqlen = new_swa_evicted_seqlen
|
||||
|
||||
|
||||
@@ -309,6 +320,25 @@ def evict_from_tree_cache(tree_cache: BasePrefixCache | None, num_tokens: int):
|
||||
tree_cache.evict(EvictParams(num_tokens=num_tokens))
|
||||
|
||||
|
||||
def _compute_dsv4_state_lens(batch, *, is_decode: bool):
|
||||
"""Per-req c{4,128}_state pool alloc lens (a ``DSV4StateLens``) for this
|
||||
alloc step. The DSV4-NPU allocator owns the computation (it also mutates the
|
||||
per-req cumulative state on each ``Req``); we just trigger it here, right
|
||||
before the paged alloc that consumes the result.
|
||||
|
||||
None on CUDA / non-V4 paths (allocator has no ``compute_dsv4_state_lens_*``)
|
||||
so the ``alloc_paged_token_slots_*`` forwarding stays a no-op.
|
||||
"""
|
||||
allocator = batch.token_to_kv_pool_allocator
|
||||
if not hasattr(allocator, "compute_dsv4_state_lens_extend"):
|
||||
return None
|
||||
if is_decode:
|
||||
return allocator.compute_dsv4_state_lens_decode(batch.reqs)
|
||||
return allocator.compute_dsv4_state_lens_extend(
|
||||
batch.reqs, batch.seq_lens_cpu.tolist()
|
||||
)
|
||||
|
||||
|
||||
def alloc_paged_token_slots_extend(
|
||||
tree_cache: BasePrefixCache,
|
||||
prefix_lens: torch.Tensor,
|
||||
@@ -318,6 +348,9 @@ def alloc_paged_token_slots_extend(
|
||||
last_loc: torch.Tensor,
|
||||
extend_num_tokens: int,
|
||||
backup_state: bool = False,
|
||||
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.
|
||||
allocator = tree_cache.token_to_kv_pool_allocator
|
||||
@@ -328,15 +361,35 @@ def alloc_paged_token_slots_extend(
|
||||
if backup_state:
|
||||
state = allocator.backup_state()
|
||||
|
||||
out_cache_loc = allocator.alloc_extend(
|
||||
is_dsv4 = req_pool_indices is not None and hasattr(allocator, "c4_attn_allocator")
|
||||
extra_alloc_kwargs = {}
|
||||
if is_dsv4:
|
||||
extra_alloc_kwargs["req_pool_indices"] = req_pool_indices
|
||||
# Pass the per-req tables in per call for the c-pool / state last_loc
|
||||
# lookup; the allocator holds no reference to the pool.
|
||||
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,
|
||||
prefix_lens_cpu,
|
||||
seq_lens,
|
||||
seq_lens_cpu,
|
||||
last_loc,
|
||||
extend_num_tokens,
|
||||
**extra_alloc_kwargs,
|
||||
)
|
||||
|
||||
if is_dsv4:
|
||||
bundle = out
|
||||
out_cache_loc = None if bundle is None else bundle.out_full_loc
|
||||
if batch is not None:
|
||||
batch.out_cache_loc_dsv4 = bundle
|
||||
else:
|
||||
out_cache_loc = out
|
||||
|
||||
if out_cache_loc is None:
|
||||
error_msg = (
|
||||
f"Prefill out of memory. Try to lower your batch size.\n"
|
||||
@@ -431,6 +484,9 @@ def alloc_for_extend(
|
||||
seq_lens_cpu=batch.seq_lens_cpu,
|
||||
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,
|
||||
)
|
||||
|
||||
# Write to req_to_token_pool
|
||||
@@ -448,6 +504,16 @@ def alloc_for_extend(
|
||||
batch.req_to_token_pool,
|
||||
)
|
||||
|
||||
# DSV4-NPU hook: write c4/c128/swa per-req tables from the stashed bundle.
|
||||
# No-op on non-DSV4 paths (out_cache_loc_dsv4 stays None there).
|
||||
if _is_npu:
|
||||
maybe_write_dsv4_extend(
|
||||
batch,
|
||||
req_pool_indices_cpu,
|
||||
prefix_lens_cpu,
|
||||
batch.seq_lens_cpu,
|
||||
)
|
||||
|
||||
return out_cache_loc, req_pool_indices_device, req_pool_indices_cpu
|
||||
|
||||
|
||||
@@ -457,6 +523,9 @@ def alloc_paged_token_slots_decode(
|
||||
seq_lens_cpu: torch.Tensor,
|
||||
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."""
|
||||
allocator = tree_cache.token_to_kv_pool_allocator
|
||||
@@ -464,7 +533,28 @@ def alloc_paged_token_slots_decode(
|
||||
num_tokens = len(seq_lens) * allocator.page_size
|
||||
evict_from_tree_cache(tree_cache, num_tokens)
|
||||
|
||||
out_cache_loc = allocator.alloc_decode(seq_lens, seq_lens_cpu, last_loc)
|
||||
# DSV4-NPU allocator also needs req_pool_indices + per-req state lens 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")
|
||||
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; the allocator holds
|
||||
# no reference to the pool.
|
||||
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)
|
||||
|
||||
if is_dsv4:
|
||||
bundle = out
|
||||
out_cache_loc = None if bundle is None else bundle.out_full_loc
|
||||
if batch is not None:
|
||||
batch.out_cache_loc_dsv4 = bundle
|
||||
else:
|
||||
out_cache_loc = out
|
||||
|
||||
if out_cache_loc is None:
|
||||
error_msg = (
|
||||
@@ -508,6 +598,9 @@ def alloc_for_decode(batch: ScheduleBatch, token_per_req: int) -> torch.Tensor:
|
||||
seq_lens_cpu=batch.seq_lens_cpu + token_per_req,
|
||||
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,
|
||||
)
|
||||
|
||||
# Write to req_to_token_pool
|
||||
@@ -520,6 +613,15 @@ def alloc_for_decode(batch: ScheduleBatch, token_per_req: int) -> torch.Tensor:
|
||||
(batch.req_pool_indices, locs), out_cache_loc.to(torch.int32)
|
||||
)
|
||||
|
||||
# DSV4-NPU hook: post-decode write of c4/c128/swa per-req tables from the
|
||||
# stashed bundle. No-op on non-DSV4 paths (out_cache_loc_dsv4 stays None).
|
||||
if _is_npu:
|
||||
maybe_write_dsv4_decode(
|
||||
batch,
|
||||
batch.seq_lens_cpu + token_per_req,
|
||||
token_per_req,
|
||||
)
|
||||
|
||||
return out_cache_loc
|
||||
|
||||
|
||||
@@ -576,6 +678,8 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
|
||||
req.mamba_pool_idx is not None
|
||||
), "mamba state is freed while the tree cache does not manage mamba states"
|
||||
tree_cache.req_to_token_pool.free_mamba_cache(req)
|
||||
# The DSV4-NPU ReqToTokenPool subclass's free() additionally releases the
|
||||
# c4/c128 state pages; other ReqToTokenPool subclasses are a no-op here.
|
||||
tree_cache.req_to_token_pool.free(req)
|
||||
|
||||
|
||||
|
||||
@@ -2,15 +2,21 @@ from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
from contextlib import nullcontext
|
||||
from math import gcd
|
||||
|
||||
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
|
||||
from sglang.srt.utils import is_hip, is_npu
|
||||
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:
|
||||
return a // gcd(a, b) * b
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
@@ -109,16 +115,39 @@ class CompressStatePool:
|
||||
last_dim = 3 * head_dim
|
||||
else:
|
||||
self._size = size + self.ring_size + 1
|
||||
self._size = (self._size + ratio - 1) // ratio * ratio
|
||||
# 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
|
||||
)
|
||||
self._size = (self._size + pad_to - 1) // pad_to * pad_to
|
||||
self._logical_size = self._size
|
||||
last_dim = 2 * (1 + overlap) * head_dim
|
||||
|
||||
self.last_dim = last_dim
|
||||
self._alloc_kv_score_buffer(
|
||||
dtype=dtype, device=device, enable_memory_saver=enable_memory_saver
|
||||
)
|
||||
if not online:
|
||||
self.kv_score_buffer[-1].clear()
|
||||
|
||||
def _alloc_kv_score_buffer(
|
||||
self, *, dtype: torch.dtype, device: str, enable_memory_saver: bool
|
||||
) -> None:
|
||||
"""Allocate the flat ``(self._size, self.last_dim)`` kv+score buffer
|
||||
under the memory-saver / custom-mem-pool context and wrap it in
|
||||
: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.
|
||||
"""
|
||||
if _is_hip:
|
||||
self.kv_score_buffer = KVAndScore(
|
||||
torch.empty((self._size, last_dim), dtype=dtype, device=device)
|
||||
torch.empty((self._size, self.last_dim), dtype=dtype, device=device)
|
||||
)
|
||||
if not online:
|
||||
self.kv_score_buffer[-1].clear()
|
||||
else:
|
||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
enable=enable_memory_saver
|
||||
@@ -126,7 +155,6 @@ class CompressStatePool:
|
||||
self.enable_custom_mem_pool, self.custom_mem_pool, _ = (
|
||||
maybe_init_custom_mem_pool(device=device)
|
||||
)
|
||||
|
||||
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||||
with (
|
||||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||||
@@ -135,13 +163,31 @@ class CompressStatePool:
|
||||
):
|
||||
self.kv_score_buffer = KVAndScore(
|
||||
torch.empty(
|
||||
(self._size, last_dim),
|
||||
(self._size, self.last_dim),
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
)
|
||||
if not online:
|
||||
self.kv_score_buffer[-1].clear()
|
||||
|
||||
@property
|
||||
def state_cache_3d(self) -> torch.Tensor:
|
||||
"""``[block_num, page_size, last_dim]`` view of the flat kv+score
|
||||
buffer. ``last_dim = 2*(1+overlap)*head_dim`` — exactly the
|
||||
``2*coff*D`` layout the fused compressor op wants for its
|
||||
``state_cache`` argument (kv at ``[:, :, :coff*D]``, score at
|
||||
``[:, :, coff*D:]``). Only valid for the non-online buffer; the
|
||||
online layout has ``last_dim = 3*head_dim`` which the fused path
|
||||
doesn't use.
|
||||
"""
|
||||
assert not self.online, (
|
||||
"state_cache_3d is for the fused compressor path; "
|
||||
"online (3*head_dim) buffer is indexer-only."
|
||||
)
|
||||
assert self.page_size > 1, (
|
||||
"state_cache_3d requires page_size>1; pool was constructed "
|
||||
"with the default page_size=1 (flat 2D layout)."
|
||||
)
|
||||
return self.kv_score_buffer.kv_score.view(-1, self.page_size, self.last_dim)
|
||||
|
||||
def translate_from_swa_loc_to_state_loc(
|
||||
self, swa_loc: torch.Tensor
|
||||
|
||||
@@ -568,48 +568,46 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
self.unified_swa_pages = self.unified_kv_pool.swa_pages
|
||||
else:
|
||||
self.unified_kv_pool = None
|
||||
self.swa_kv_pool = DeepSeekV4SingleKVPool(
|
||||
swa_size,
|
||||
swa_page_size,
|
||||
dtype,
|
||||
qk_nope_head_dim,
|
||||
qk_rope_head_dim,
|
||||
layer_num,
|
||||
device,
|
||||
enable_memory_saver,
|
||||
self.swa_kv_pool = self._make_kv_pool(
|
||||
size=swa_size,
|
||||
page_size=swa_page_size,
|
||||
dtype=dtype,
|
||||
layer_num=layer_num,
|
||||
device=device,
|
||||
enable_memory_saver=enable_memory_saver,
|
||||
global_page_size=swa_page_size,
|
||||
)
|
||||
|
||||
c4_kv_pool_type = DeepSeekV4SingleKVPool
|
||||
if enable_hisparse:
|
||||
c4_kv_pool_type = HiSparseC4DevicePool
|
||||
self.c4_kv_pool = c4_kv_pool_type(
|
||||
c4_size,
|
||||
c4_page_size,
|
||||
dtype,
|
||||
qk_nope_head_dim,
|
||||
qk_rope_head_dim,
|
||||
c4_layer_num,
|
||||
device,
|
||||
enable_memory_saver,
|
||||
)
|
||||
c4_kv_pool_type = DeepSeekV4SingleKVPool
|
||||
if enable_hisparse:
|
||||
c4_kv_pool_type = HiSparseC4DevicePool
|
||||
self.c4_kv_pool = self._make_kv_pool(
|
||||
size=c4_size,
|
||||
page_size=c4_page_size,
|
||||
dtype=dtype,
|
||||
layer_num=c4_layer_num,
|
||||
device=device,
|
||||
enable_memory_saver=enable_memory_saver,
|
||||
global_page_size=page_size,
|
||||
cls=c4_kv_pool_type,
|
||||
)
|
||||
|
||||
self.c128_kv_pool = DeepSeekV4SingleKVPool(
|
||||
c128_size,
|
||||
c128_page_size,
|
||||
dtype,
|
||||
qk_nope_head_dim,
|
||||
qk_rope_head_dim,
|
||||
c128_layer_num,
|
||||
device,
|
||||
enable_memory_saver,
|
||||
)
|
||||
self.c128_kv_pool = self._make_kv_pool(
|
||||
size=c128_size,
|
||||
page_size=c128_page_size,
|
||||
dtype=dtype,
|
||||
layer_num=c128_layer_num,
|
||||
device=device,
|
||||
enable_memory_saver=enable_memory_saver,
|
||||
global_page_size=page_size,
|
||||
)
|
||||
|
||||
indexer_size = (
|
||||
self.c4_logical_size
|
||||
if (not _is_hip or envs.SGLANG_OPT_USE_COMPRESSOR_V2.get())
|
||||
else c4_size
|
||||
)
|
||||
self.c4_indexer_kv_pool = DeepSeekV4IndexerPool(
|
||||
self.c4_indexer_kv_pool = self._make_indexer_pool(
|
||||
indexer_size,
|
||||
c4_page_size,
|
||||
dtype,
|
||||
@@ -741,6 +739,99 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
|
||||
return data_ptrs, data_lens, item_lens
|
||||
|
||||
def _make_kv_pool(
|
||||
self,
|
||||
*,
|
||||
size: int,
|
||||
page_size: int,
|
||||
dtype: torch.dtype,
|
||||
layer_num: int,
|
||||
device: str,
|
||||
enable_memory_saver: bool,
|
||||
global_page_size: int,
|
||||
cls: type = DeepSeekV4SingleKVPool,
|
||||
) -> DeepSeekV4SingleKVPool:
|
||||
"""Build a full / SWA / c4 / c128 single-KV pool. ``global_page_size``
|
||||
is the model-wide page_size (== ``page_size`` for the SWA pool, larger
|
||||
for the per-ratio c4/c128 pools); the default CUDA pool ignores it.
|
||||
Overridden by :class:`DSV4NPUTokenToKVPool` to swap in the NPU bf16
|
||||
PA_ND variant, which needs ``global_page_size`` for its kernel view."""
|
||||
del global_page_size # CUDA pools key only off their own page_size
|
||||
return cls(
|
||||
size,
|
||||
page_size,
|
||||
dtype,
|
||||
self.qk_nope_head_dim,
|
||||
self.qk_rope_head_dim,
|
||||
layer_num,
|
||||
device,
|
||||
enable_memory_saver,
|
||||
)
|
||||
|
||||
def _make_indexer_pool(
|
||||
self,
|
||||
size: int,
|
||||
page_size: int,
|
||||
dtype: torch.dtype,
|
||||
index_head_dim: int,
|
||||
layer_num: int,
|
||||
device: str,
|
||||
enable_memory_saver: bool,
|
||||
) -> DeepSeekV4IndexerPool:
|
||||
"""Build the c4 lightning-indexer K pool (packed CUDA layout).
|
||||
Overridden by :class:`DSV4NPUTokenToKVPool` to swap in the
|
||||
dedicated-buffer NPU variant (int8 K + fp16 scale)."""
|
||||
return DeepSeekV4IndexerPool(
|
||||
size,
|
||||
page_size,
|
||||
dtype,
|
||||
index_head_dim,
|
||||
layer_num,
|
||||
device,
|
||||
enable_memory_saver,
|
||||
)
|
||||
|
||||
def _state_pool_size(self, ratio: int) -> int:
|
||||
return self.c4_state_pool_size if ratio == 4 else self.c128_state_pool_size
|
||||
|
||||
def _make_attn_state_pool(
|
||||
self, ratio: int, enable_memory_saver: bool
|
||||
) -> CompressStatePool:
|
||||
"""Build the per-layer attention compress-state pool for ``ratio``
|
||||
(4 or 128). Overridden by :class:`DSV4NPUTokenToKVPool` to swap the
|
||||
ring-buffered pool for the NPU paged one."""
|
||||
return CompressStatePool(
|
||||
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,
|
||||
online=(ratio == 128 and ONLINE_C128),
|
||||
swa_page_size=self.swa_page_size,
|
||||
online_mtp_max_draft_tokens=(
|
||||
self.online_mtp_max_draft_tokens if ratio == 128 else 0
|
||||
),
|
||||
)
|
||||
|
||||
def _make_indexer_state_pool(
|
||||
self, ratio: int, enable_memory_saver: bool
|
||||
) -> CompressStatePool:
|
||||
"""Build the per-layer indexer compress-state pool (c4 only)."""
|
||||
return CompressStatePool(
|
||||
size=self._state_pool_size(ratio),
|
||||
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,
|
||||
swa_page_size=self.swa_page_size,
|
||||
)
|
||||
|
||||
def _init_paged_compress_states(self, enable_memory_saver: bool):
|
||||
c4_state_pool_size = self.c4_state_pool_size
|
||||
c128_state_pool_size = self.c128_state_pool_size
|
||||
@@ -754,37 +845,14 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
ratio = self.compression_ratios[idx]
|
||||
if ratio == 0:
|
||||
continue
|
||||
overlap = ratio == 4
|
||||
size = c4_state_pool_size if ratio == 4 else c128_state_pool_size
|
||||
ring_size = self.get_ring_size(ratio)
|
||||
|
||||
self.compress_state_pools[idx] = CompressStatePool(
|
||||
size=size,
|
||||
ring_size=ring_size,
|
||||
overlap=overlap,
|
||||
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,
|
||||
online=(ratio == 128 and ONLINE_C128),
|
||||
swa_page_size=self.swa_page_size,
|
||||
online_mtp_max_draft_tokens=(
|
||||
self.online_mtp_max_draft_tokens if ratio == 128 else 0
|
||||
),
|
||||
self.compress_state_pools[idx] = self._make_attn_state_pool(
|
||||
ratio, enable_memory_saver
|
||||
)
|
||||
|
||||
if ratio == 4:
|
||||
self.indexer_compress_state_pools[idx] = CompressStatePool(
|
||||
size=size,
|
||||
ring_size=ring_size,
|
||||
overlap=overlap,
|
||||
head_dim=self.indexer_head_dim,
|
||||
device=self.device,
|
||||
dtype=self.c4_state_dtype,
|
||||
enable_memory_saver=enable_memory_saver,
|
||||
ratio=ratio,
|
||||
swa_page_size=self.swa_page_size,
|
||||
self.indexer_compress_state_pools[idx] = self._make_indexer_state_pool(
|
||||
ratio, enable_memory_saver
|
||||
)
|
||||
|
||||
def _init_compressed_layer_mapping(self):
|
||||
|
||||
@@ -217,6 +217,65 @@ def compute_local_num_token_non_padded(
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
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).
|
||||
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).
|
||||
|
||||
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
|
||||
class NgramEmbeddingInfo:
|
||||
"""Ngram embedding state for LongCat models."""
|
||||
@@ -286,6 +345,9 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
# The original sequence length without being chunked. Qwen-1M related.
|
||||
orig_seq_lens: Optional[torch.Tensor] = None
|
||||
|
||||
# DSV4-NPU only: per-pool slot bundle from DSV4NPUTokenToKVPoolAllocator,
|
||||
# consumed by the Ascend backend for PA_ND block tables. None elsewhere.
|
||||
out_cache_loc_dsv4: Optional[DSV4OutCacheLoc] = None
|
||||
# The indices to track mamba state with
|
||||
mamba_track_indices: Optional[torch.Tensor] = None # shape: [b], int64
|
||||
# The mask to track mamba state if needed
|
||||
@@ -615,6 +677,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
# Inputs aliased by reference from ScheduleBatch
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
orig_seq_lens=batch.orig_seq_lens,
|
||||
out_cache_loc_dsv4=batch.out_cache_loc_dsv4,
|
||||
mamba_track_indices=batch.mamba_track_indices,
|
||||
mamba_track_mask=batch.mamba_track_mask,
|
||||
mamba_track_seqlens=batch.mamba_track_seqlens,
|
||||
|
||||
@@ -391,7 +391,17 @@ class ModelRunnerKVCacheMixin:
|
||||
start_layer=self.start_layer,
|
||||
)
|
||||
else:
|
||||
self.req_to_token_pool = ReqToTokenPool(
|
||||
# DSV4 on NPU needs an extended ReqToTokenPool holding per-req
|
||||
# swa/c4/c128/c{4,128}_state tables; others stay on the stock one.
|
||||
req_to_token_pool_cls = ReqToTokenPool
|
||||
if _is_npu and is_deepseek_v4(self.model_config.hf_config):
|
||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_req_to_token_pool import (
|
||||
DSV4NPUReqToTokenPool,
|
||||
)
|
||||
|
||||
req_to_token_pool_cls = DSV4NPUReqToTokenPool
|
||||
|
||||
self.req_to_token_pool = req_to_token_pool_cls(
|
||||
size=max_num_reqs,
|
||||
max_context_len=self.model_config.context_len
|
||||
+ extra_max_context_len,
|
||||
@@ -412,7 +422,8 @@ class ModelRunnerKVCacheMixin:
|
||||
|
||||
if is_dsv4_model:
|
||||
swa_page_size = self.page_size
|
||||
assert swa_page_size == 256, "In paged swa mode, page_size must be 256."
|
||||
if not _is_npu:
|
||||
assert swa_page_size == 256, "In paged swa mode, page_size must be 256."
|
||||
|
||||
if self.is_draft_worker:
|
||||
from sglang.srt.models.deepseek_v4_nextn import (
|
||||
@@ -424,7 +435,40 @@ class ModelRunnerKVCacheMixin:
|
||||
] * self.num_effective_layers
|
||||
else:
|
||||
compression_ratios = self.model_config.compress_ratios
|
||||
self.token_to_kv_pool = DeepSeekV4TokenToKVPool(
|
||||
|
||||
# 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.
|
||||
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=self.page_size,
|
||||
max_num_reqs=self.max_running_requests,
|
||||
)
|
||||
c128_state_pool_size = npu_state_pool_size(
|
||||
ratio=128,
|
||||
page_size=self.page_size,
|
||||
max_num_reqs=self.max_running_requests,
|
||||
)
|
||||
else:
|
||||
pool_cls = DeepSeekV4TokenToKVPool
|
||||
c4_state_pool_size = self.c4_state_pool_size
|
||||
c128_state_pool_size = self.c128_state_pool_size
|
||||
|
||||
self.token_to_kv_pool = pool_cls(
|
||||
max_num_reqs=self.max_running_requests,
|
||||
# SWA ring is indexed by req_pool_idx; PD decode inflates req_to_token
|
||||
# past max_running_requests (pre-alloc), so size to the real capacity.
|
||||
@@ -432,8 +476,8 @@ class ModelRunnerKVCacheMixin:
|
||||
swa_size=self.swa_max_total_num_tokens,
|
||||
c4_size=self.c4_max_total_num_tokens,
|
||||
c128_size=self.c128_max_total_num_tokens,
|
||||
c4_state_pool_size=self.c4_state_pool_size,
|
||||
c128_state_pool_size=self.c128_state_pool_size,
|
||||
c4_state_pool_size=c4_state_pool_size,
|
||||
c128_state_pool_size=c128_state_pool_size,
|
||||
page_size=self.page_size,
|
||||
swa_page_size=swa_page_size,
|
||||
sliding_window=self.model_config.window_size,
|
||||
@@ -763,10 +807,21 @@ class ModelRunnerKVCacheMixin:
|
||||
)
|
||||
elif _is_npu and (
|
||||
self.server_args.attention_backend == "ascend"
|
||||
or is_dsv4_model
|
||||
or self.hybrid_gdn_config is not None
|
||||
):
|
||||
if self.is_hybrid_swa:
|
||||
self.token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
||||
# DSV4 on NPU: SWA allocator subclass that also drives the
|
||||
# c4/c128 allocators, producing a DSV4OutCacheLoc per alloc.
|
||||
if is_dsv4_model:
|
||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_allocator import (
|
||||
DSV4NPUTokenToKVPoolAllocator,
|
||||
)
|
||||
|
||||
swa_allocator_cls = DSV4NPUTokenToKVPoolAllocator
|
||||
else:
|
||||
swa_allocator_cls = SWATokenToKVPoolAllocator
|
||||
self.token_to_kv_pool_allocator = swa_allocator_cls(
|
||||
self.full_max_total_num_tokens,
|
||||
self.swa_max_total_num_tokens,
|
||||
page_size=self.page_size,
|
||||
@@ -843,6 +898,13 @@ class ModelRunnerKVCacheMixin:
|
||||
)
|
||||
)
|
||||
|
||||
# DSV4-NPU: wire allocator back-ref into req_to_token_pool so its
|
||||
# free(req) can release c4/c128 pool pages alongside the slot.
|
||||
if hasattr(self.req_to_token_pool, "register_dsv4_allocator"):
|
||||
self.req_to_token_pool.register_dsv4_allocator(
|
||||
self.token_to_kv_pool_allocator
|
||||
)
|
||||
|
||||
else:
|
||||
assert self.is_draft_worker
|
||||
if self.is_hybrid_swa:
|
||||
|
||||
@@ -156,6 +156,7 @@ def build_replay_fb_view(
|
||||
seq_lens_cpu=buffers.seq_lens_cpu[:bs],
|
||||
encoder_lens=buffers.encoder_lens[:bs] if is_encoder_decoder else None,
|
||||
out_cache_loc=getattr(forward_batch, "out_cache_loc", None),
|
||||
out_cache_loc_dsv4=getattr(forward_batch, "out_cache_loc_dsv4", None),
|
||||
spec_info=forward_batch.spec_info,
|
||||
)
|
||||
|
||||
|
||||
@@ -385,9 +385,17 @@ class DeepseekV2MLP(nn.Module):
|
||||
|
||||
# Fallback: fused silu+clamp kernel (still faster than unfused)
|
||||
elif self.swiglu_limit is not None:
|
||||
M, N = gate_up.shape
|
||||
x = gate_up.new_empty((M, N // 2))
|
||||
silu_and_mul_clamp(gate_up, x, float(self.swiglu_limit))
|
||||
if _is_npu:
|
||||
_g, _u = gate_up.chunk(2, dim=-1)
|
||||
_lim = float(self.swiglu_limit)
|
||||
gate_up = torch.cat(
|
||||
[_g.clamp(max=_lim), _u.clamp(min=-_lim, max=_lim)], dim=-1
|
||||
)
|
||||
x = self.act_fn(gate_up)
|
||||
else:
|
||||
M, N = gate_up.shape
|
||||
x = gate_up.new_empty((M, N // 2))
|
||||
silu_and_mul_clamp(gate_up, x, float(self.swiglu_limit))
|
||||
else:
|
||||
x = self.act_fn(gate_up)
|
||||
x, _ = self.down_proj(
|
||||
@@ -493,6 +501,8 @@ class MoEGate(nn.Module):
|
||||
|
||||
elif _use_aiter:
|
||||
logits = aiter_dsv3_router_gemm(hidden_states, self.weight)
|
||||
elif _is_npu:
|
||||
logits = F.linear(hidden_states, self.weight, None)
|
||||
else:
|
||||
if self.is_deepseek_v4:
|
||||
from sglang.jit_kernel.dsv4 import linear_bf16_fp32
|
||||
|
||||
@@ -29,6 +29,7 @@ from sglang.srt.compilation.compilation_config import register_split_op
|
||||
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
get_tensor_model_parallel_world_size,
|
||||
get_tp_group,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
@@ -47,10 +48,15 @@ from sglang.srt.layers.communicator_dsa_cp import (
|
||||
dsa_cp_gather_hidden_states,
|
||||
dsa_cp_reduce_scatter_hidden_states,
|
||||
)
|
||||
from sglang.srt.layers.deepseek_v4_rope import (
|
||||
v4_rope_inplace_npu,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
_DpGatheredBufferWrapper,
|
||||
attn_tp_all_gather,
|
||||
attn_tp_all_reduce,
|
||||
dp_gather_partial,
|
||||
dp_gather_replicate,
|
||||
dp_scatter,
|
||||
get_dp_global_num_tokens,
|
||||
get_global_dp_buffer,
|
||||
@@ -61,7 +67,7 @@ from sglang.srt.layers.dp_attention import (
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
from sglang.srt.layers.mhc import mhc_fused_post_pre
|
||||
from sglang.srt.layers.mhc import mhc_fused_post_pre, npu_hc_pre
|
||||
from sglang.srt.layers.moe import get_moe_a2a_backend, should_use_dp_reduce_scatterv
|
||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
||||
from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8
|
||||
@@ -129,6 +135,11 @@ from sglang.srt.utils import (
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
from sglang.srt.utils.hf_transformers_utils import get_rope_config
|
||||
|
||||
# NPU-only: bind torch_npu here so _compute_q_b / _forward_prepare can call
|
||||
# torch_npu.npu_rms_norm directly (imports elsewhere aren't visible in this module).
|
||||
if _is_npu:
|
||||
import torch_npu
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_FP8_WO_A_GEMM = envs.SGLANG_OPT_FP8_WO_A_GEMM.get()
|
||||
@@ -293,7 +304,12 @@ class MQALayer(nn.Module):
|
||||
if compress_ratio_override is not None
|
||||
else config.compress_ratios[layer_id]
|
||||
)
|
||||
assert compress_ratio in [0, 4, 128]
|
||||
|
||||
assert compress_ratio in (
|
||||
0,
|
||||
4,
|
||||
128,
|
||||
), f"V4 compress_ratio: expected one of (0, 4, 128), got {compress_ratio}"
|
||||
self.compress_ratio: Literal[0, 4, 128] = compress_ratio
|
||||
|
||||
assert self.head_dim == config.head_dim
|
||||
@@ -317,11 +333,9 @@ class MQALayer(nn.Module):
|
||||
|
||||
from sglang.srt.layers.deepseek_v4_rope import precompute_freqs_cis
|
||||
|
||||
assert self.compress_ratio in {0, 4, 128}
|
||||
if self.compress_ratio:
|
||||
original_seq_len = rope_scaling["original_max_position_embeddings"]
|
||||
else:
|
||||
original_seq_len = 0
|
||||
# YARN correction applies to ALL layers (dense and compressed share the same
|
||||
# YARN-corrected inv_freq); only the rope base differs (rope_theta vs compress_rope_theta).
|
||||
original_seq_len = rope_scaling["original_max_position_embeddings"]
|
||||
|
||||
freqs_cis = precompute_freqs_cis(
|
||||
dim=self.qk_rope_head_dim,
|
||||
@@ -354,7 +368,7 @@ class MQALayer(nn.Module):
|
||||
|
||||
self.compressor = None
|
||||
self.indexer = None
|
||||
if self.compress_ratio:
|
||||
if self.compress_ratio in (4, 128):
|
||||
self.compressor = Compressor(
|
||||
config,
|
||||
layer_id=self.layer_id,
|
||||
@@ -436,7 +450,8 @@ class MQALayer(nn.Module):
|
||||
self.hidden_size,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
reduce_results=attn_tp_size > 1,
|
||||
reduce_results=attn_tp_size == get_tensor_model_parallel_world_size()
|
||||
and attn_tp_size > 1,
|
||||
prefix=add_prefix("wo_b", prefix),
|
||||
tp_rank=attn_tp_rank,
|
||||
tp_size=attn_tp_size,
|
||||
@@ -811,6 +826,33 @@ class MQALayer(nn.Module):
|
||||
forward_batch,
|
||||
torch.cuda.current_stream(),
|
||||
)
|
||||
elif _is_npu:
|
||||
q_lora = self.q_norm(q_lora)
|
||||
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]
|
||||
|
||||
if qkv_a is not None:
|
||||
kv = qkv_a[..., self.q_lora_rank :]
|
||||
else:
|
||||
kv, _ = self.wkv(x)
|
||||
kv = self.kv_norm(kv)
|
||||
|
||||
v4_rope_inplace_npu(
|
||||
q[..., -self.qk_rope_head_dim :],
|
||||
kv[..., -self.qk_rope_head_dim :].unsqueeze(1),
|
||||
self.freqs_cis,
|
||||
positions,
|
||||
)
|
||||
attn_backend.store_cache(
|
||||
layer_id=self.layer_id,
|
||||
swa_k=kv,
|
||||
forward_batch=forward_batch,
|
||||
)
|
||||
kv = None
|
||||
if q_out is not None:
|
||||
q_out.copy_(q)
|
||||
else:
|
||||
q_lora = self.q_norm(q_lora)
|
||||
q = self._compute_q_b(q_lora, positions, q_out)
|
||||
@@ -879,9 +921,6 @@ class MQALayer(nn.Module):
|
||||
x_quant=None,
|
||||
) -> torch.Tensor:
|
||||
if not get_attn_tp_context().input_scattered and x.shape[0] == 0:
|
||||
assert (
|
||||
not self.wo_b.reduce_results
|
||||
), "short-circuiting allreduce will lead to hangs"
|
||||
return x
|
||||
|
||||
attn_backend = get_attn_backend()
|
||||
@@ -1000,13 +1039,22 @@ class MQALayer(nn.Module):
|
||||
save_kv_cache=save_kv_cache,
|
||||
)
|
||||
o = o[:, tp_slice, :]
|
||||
fused_rope_inplace(
|
||||
o[..., -self.qk_rope_head_dim :],
|
||||
None,
|
||||
self.freqs_cis,
|
||||
positions=positions,
|
||||
inverse=True,
|
||||
)
|
||||
if _is_npu:
|
||||
v4_rope_inplace_npu(
|
||||
o[..., -self.qk_rope_head_dim :],
|
||||
None,
|
||||
self.freqs_cis,
|
||||
positions,
|
||||
inverse=True,
|
||||
)
|
||||
else:
|
||||
fused_rope_inplace(
|
||||
o[..., -self.qk_rope_head_dim :],
|
||||
None,
|
||||
self.freqs_cis,
|
||||
positions=positions,
|
||||
inverse=True,
|
||||
)
|
||||
|
||||
o = o.view(o.shape[0], self.n_local_groups, -1)
|
||||
|
||||
@@ -1034,6 +1082,8 @@ class MQALayer(nn.Module):
|
||||
o = torch.einsum("tgd,grd->tgr", o, wo_a)
|
||||
|
||||
o, _ = self.wo_b(o.flatten(1))
|
||||
if self.tp_size > 1 and self.tp_size < get_tensor_model_parallel_world_size():
|
||||
o = attn_tp_all_reduce(o)
|
||||
|
||||
return o
|
||||
|
||||
@@ -1233,6 +1283,7 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
hc_scale: torch.Tensor,
|
||||
hc_base: torch.Tensor,
|
||||
norm: Optional[nn.Module] = None,
|
||||
forward_batch: Optional[ForwardBatch] = None,
|
||||
):
|
||||
"""If *norm* is given and the TileLang path is active, the returned
|
||||
hidden_states are already post-norm (the norm is fused into the kernel)."""
|
||||
@@ -1248,6 +1299,19 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
|
||||
shape, dtype = x.size(), x.dtype
|
||||
|
||||
if _is_npu:
|
||||
return npu_hc_pre(
|
||||
x,
|
||||
hc_fn,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
hc_mult=self.hc_mult,
|
||||
hc_sinkhorn_iters=self.hc_sinkhorn_iters,
|
||||
rms_norm_eps=self.rms_norm_eps,
|
||||
hc_eps=self.hc_eps,
|
||||
forward_batch=forward_batch,
|
||||
)
|
||||
|
||||
if x.shape[0] == 0:
|
||||
y = torch.empty((0, shape[-1]), dtype=dtype, device=x.device)
|
||||
post = torch.empty((0, self.hc_mult), dtype=torch.float32, device=x.device)
|
||||
@@ -1339,6 +1403,9 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
(0, self.hc_mult, x.shape[-1]), dtype=x.dtype, device=x.device
|
||||
)
|
||||
|
||||
if _is_npu:
|
||||
return torch.ops.custom.npu_hc_post(x, residual, post, comb)
|
||||
|
||||
if envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get():
|
||||
from sglang.srt.layers.mhc import mhc_post
|
||||
|
||||
@@ -1412,6 +1479,7 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
self.hc_attn_scale,
|
||||
self.hc_attn_base,
|
||||
norm=self.input_layernorm,
|
||||
forward_batch=forward_batch,
|
||||
)
|
||||
if not norm_fused:
|
||||
if _use_aiter and _is_gfx95_supported:
|
||||
@@ -1482,6 +1550,7 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
self.hc_ffn_scale,
|
||||
self.hc_ffn_base,
|
||||
norm=self.post_attention_layernorm,
|
||||
forward_batch=forward_batch,
|
||||
)
|
||||
if not norm_fused:
|
||||
hidden_states = self.post_attention_layernorm(hidden_states)
|
||||
@@ -1695,7 +1764,9 @@ class DeepseekV4Model(nn.Module):
|
||||
dtype=input_ids.dtype,
|
||||
device=input_ids.device,
|
||||
)
|
||||
dp_gather_partial(input_ids_global, input_ids[:, None], forward_batch)
|
||||
# Token ids are replicated within an attention-TP group. Use replicate
|
||||
# gather here to avoid summing duplicated ids when attention_tp_size > 1.
|
||||
dp_gather_replicate(input_ids_global, input_ids[:, None], forward_batch)
|
||||
input_ids_global = input_ids_global.squeeze(-1)
|
||||
else:
|
||||
input_ids_global = input_ids
|
||||
@@ -1886,6 +1957,7 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
if self.capture_aux_hidden_states:
|
||||
hidden_states, aux_hidden_states = hidden_states
|
||||
hidden_states, pre_hc_head = hidden_states
|
||||
|
||||
return self.logits_processor(
|
||||
input_ids,
|
||||
hidden_states,
|
||||
@@ -1930,7 +2002,10 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
for layer_id in range(self.model.start_layer, self.model.end_layer):
|
||||
layer = self.model.layers[layer_id]
|
||||
self_attn = layer.self_attn
|
||||
if self_attn.compress_ratio != 0 and not self_attn.compressor.ape_converted:
|
||||
if (
|
||||
self_attn.compress_ratio in (4, 128)
|
||||
and not self_attn.compressor.ape_converted
|
||||
):
|
||||
self_attn.compressor.apply_ape_hotfix()
|
||||
if (
|
||||
self_attn.compress_ratio == 4
|
||||
@@ -1941,7 +2016,9 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
|
||||
@staticmethod
|
||||
def remap_weight_name_to_dpsk_hf_format(
|
||||
name: str, is_nextn: bool = False, num_hidden_layers: Optional[int] = None
|
||||
name: str,
|
||||
is_nextn: bool = False,
|
||||
num_hidden_layers: Optional[int] = None,
|
||||
) -> str:
|
||||
if name == "embed.weight":
|
||||
return "model.embed_tokens.weight"
|
||||
@@ -2339,8 +2416,10 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
del self.lm_head.weight
|
||||
self.model.embed_tokens.weight = embed
|
||||
self.lm_head.weight = head
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
# Hot weight reload (RL workflows). Use the device-agnostic module
|
||||
# accessor so this works on both CUDA/HIP and NPU.
|
||||
torch.get_device_module().empty_cache()
|
||||
torch.get_device_module().synchronize()
|
||||
|
||||
@classmethod
|
||||
def get_model_config_for_expert_location(cls, config):
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import logging
|
||||
|
||||
from sglang.srt.server_args import ServerArgs, get_global_server_args
|
||||
from sglang.srt.utils.common import is_blackwell, is_hip, is_musa
|
||||
from sglang.srt.utils.common import is_blackwell, is_hip, is_musa, is_npu
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -236,7 +236,11 @@ class DraftBackendFactory:
|
||||
)
|
||||
|
||||
def _create_dsv4_decode_backend(self):
|
||||
if is_hip():
|
||||
# On NPU the "dsv4" backend resolves to the Ascend V4 subclass; its
|
||||
# draft path reuses the Ascend multi-step draft backend.
|
||||
if is_npu():
|
||||
return self._create_ascend_decode_backend()
|
||||
elif is_hip():
|
||||
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
|
||||
DeepseekV4MultiStepBackend,
|
||||
)
|
||||
@@ -333,7 +337,11 @@ class DraftBackendFactory:
|
||||
return None
|
||||
|
||||
def _create_dsv4_prefill_backend(self):
|
||||
if is_hip():
|
||||
# On NPU the "dsv4" backend resolves to the Ascend V4 subclass; its
|
||||
# draft-extend path reuses the Ascend prefill draft backend.
|
||||
if is_npu():
|
||||
return self._create_ascend_prefill_backend()
|
||||
elif is_hip():
|
||||
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
|
||||
DeepseekV4HipRadixBackend,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user