[Kimi] Support DCP + DSpark (ported from kimi-k3 branch) (#32828)
This commit is contained in:
@@ -80,14 +80,21 @@ _TOKENSPEED_MAX_Q_LEN = 8
|
||||
|
||||
|
||||
def _get_tokenspeed_workspace(
|
||||
device: torch.device, num_heads: int, kv_lora_rank: int
|
||||
device: torch.device,
|
||||
num_heads: int,
|
||||
kv_lora_rank: int,
|
||||
max_q_len: int = _TOKENSPEED_MAX_Q_LEN,
|
||||
) -> torch.Tensor:
|
||||
from sglang.srt.runtime_context import get_resources
|
||||
|
||||
# DCP target verification gathers Q to the full head count before launching
|
||||
# TokenSpeed; size for that launch shape, not the rank-local head count.
|
||||
num_heads *= get_parallel().attn_dcp_size
|
||||
max_q_len = max(max_q_len, _TOKENSPEED_MAX_Q_LEN)
|
||||
needed = (
|
||||
tokenspeed_mla.get_num_sm(device)
|
||||
* num_heads
|
||||
* _TOKENSPEED_MAX_Q_LEN
|
||||
* max_q_len
|
||||
* (kv_lora_rank + 1)
|
||||
* 4
|
||||
)
|
||||
@@ -133,7 +140,12 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
||||
self._tokenspeed_workspace: Optional[torch.Tensor] = None
|
||||
if is_tokenspeed_mla_available():
|
||||
self._tokenspeed_workspace = _get_tokenspeed_workspace(
|
||||
self.device, self.num_q_heads, self.kv_lora_rank
|
||||
self.device,
|
||||
self.num_q_heads,
|
||||
self.kv_lora_rank,
|
||||
max_q_len=(
|
||||
model_runner.server_args.max_speculative_num_draft_tokens or 1
|
||||
),
|
||||
)
|
||||
|
||||
# Pre-JIT the prefill kernel variants. Each cute.compile takes 1-2
|
||||
|
||||
@@ -16,12 +16,7 @@ from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
|
||||
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
|
||||
from sglang.srt.runtime_context import (
|
||||
get_schedule,
|
||||
get_server_args,
|
||||
get_serving,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_serving, get_spec
|
||||
from sglang.srt.utils.common import ceil_align
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -183,8 +178,8 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
|
||||
def _release_overallocated_kv_indices(
|
||||
req: Req, start_p: int, end_p: int, tree_cache: BasePrefixCache
|
||||
) -> None:
|
||||
global_server_args = get_server_args()
|
||||
page_size = get_schedule().page_size
|
||||
allocator = tree_cache.token_to_kv_pool_allocator
|
||||
page_size = allocator.page_size
|
||||
spec_algo = get_spec().speculative_algorithm
|
||||
|
||||
# strip_thinking_cache intentionally reports output tokens as overallocated
|
||||
@@ -201,11 +196,9 @@ def _release_overallocated_kv_indices(
|
||||
indices_to_free = tree_cache.req_to_token_pool.req_to_token[req.req_pool_idx][
|
||||
start_p:end_p
|
||||
]
|
||||
# start_p is ceil-aligned above: never shares a page with
|
||||
# cache_finished_req's tail frees in the same group.
|
||||
tree_cache.token_to_kv_pool_allocator.free_segment(
|
||||
indices_to_free, start_pos=start_p
|
||||
)
|
||||
# start_p is aligned to the allocator's physical page size above, so it
|
||||
# never shares a page with cache_finished_req's tail free in this group.
|
||||
allocator.free_segment(indices_to_free, start_pos=start_p)
|
||||
|
||||
|
||||
def available_and_evictable_str(tree_cache: BasePrefixCache) -> str:
|
||||
|
||||
@@ -275,6 +275,16 @@ class KVCacheConfigurator:
|
||||
full_max_total_num_tokens = config.full_max_total_num_tokens
|
||||
swa_max_total_num_tokens = config.swa_max_total_num_tokens
|
||||
|
||||
# Draft pools are replicated, not DCP-sharded, yet consume the shared
|
||||
# allocator's virtual locs in [0, max_total * dcp_size) untranslated.
|
||||
dcp_size = self.server_args.dcp_size
|
||||
if self.is_draft_worker and dcp_size > 1:
|
||||
max_total_num_tokens *= dcp_size
|
||||
if full_max_total_num_tokens is not None:
|
||||
full_max_total_num_tokens *= dcp_size
|
||||
if swa_max_total_num_tokens is not None:
|
||||
swa_max_total_num_tokens *= dcp_size
|
||||
|
||||
# DSV4 compressed-attention pool sizes. Draft worker reuses target's
|
||||
# full/swa sizes but does NOT own c4/c128/state pools (those live on
|
||||
# the target rank only); zero them out regardless of what config holds.
|
||||
|
||||
@@ -177,7 +177,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
||||
self._cell_size = scale_kv_cell_size_per_token_for_dflash(
|
||||
target_cell_size_per_token=self._cell_size,
|
||||
target_num_layers=int(num_layers),
|
||||
draft_num_layers=int(draft_num_layers),
|
||||
draft_num_layers=int(draft_num_layers) * kvc.server_args.dcp_size,
|
||||
)
|
||||
|
||||
def _compute_cell_size(self, kvc: KVCacheConfigurator, num_layers: int) -> int:
|
||||
|
||||
@@ -5,6 +5,7 @@ from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.configs.hybrid_arch import mambaish_config
|
||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
@@ -57,6 +58,7 @@ from sglang.srt.speculative.spec_utils import (
|
||||
GrammarTree,
|
||||
build_grammar_vocab_mask,
|
||||
draft_tp_context,
|
||||
prepare_mamba_track_for_verify,
|
||||
)
|
||||
from sglang.srt.utils import get_available_gpu_memory, is_cuda
|
||||
|
||||
@@ -245,6 +247,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
)
|
||||
|
||||
self._forced_budget_frac: Optional[float] = None
|
||||
self._need_mamba_verify_commit = False
|
||||
|
||||
self._observers = DsparkStepObservers(
|
||||
planner=self._verify_planner,
|
||||
@@ -296,6 +299,12 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
def init_attention_backends(self):
|
||||
with self._draft_context():
|
||||
self._draft_worker.init_attention_backends()
|
||||
self._need_mamba_verify_commit = mambaish_config(
|
||||
self.model_runner.model_config
|
||||
) is not None and hasattr(
|
||||
self.model_runner.attn_backend,
|
||||
"update_mamba_state_after_mtp_verify",
|
||||
)
|
||||
|
||||
def init_cuda_graphs(self):
|
||||
capture_decode_cuda_graph = not get_exec().graph.disable_cuda_graph
|
||||
@@ -587,6 +596,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
and self._simulate_acc_len <= 0
|
||||
and not batch.has_grammar
|
||||
)
|
||||
prepare_mamba_track_for_verify(batch)
|
||||
with self._observers.segment(InfoSegment.TARGET_VERIFY):
|
||||
if run_compact:
|
||||
target_verify, hidden_strided = self._verify_executor.run_compact(
|
||||
@@ -644,6 +654,13 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
else:
|
||||
on_publish(accept.new_seq_lens)
|
||||
|
||||
self._commit_target_mamba_states_after_verify(
|
||||
batch=batch,
|
||||
seq_lens_pre_verify=prefix_lens,
|
||||
seq_lens_post_verify=accept.new_seq_lens,
|
||||
commit_lens=accept.commit_lens,
|
||||
)
|
||||
|
||||
folded_commit = folded_accept and epilogue.folds_commit
|
||||
if not folded_commit:
|
||||
self._verify_executor.commit_hidden(
|
||||
@@ -699,5 +716,52 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
new_seq_lens=accept.new_seq_lens,
|
||||
)
|
||||
|
||||
def _commit_target_mamba_states_after_verify(
|
||||
self,
|
||||
*,
|
||||
batch: ScheduleBatch,
|
||||
seq_lens_pre_verify: torch.Tensor,
|
||||
seq_lens_post_verify: torch.Tensor,
|
||||
commit_lens: torch.Tensor,
|
||||
) -> None:
|
||||
"""Commit the last accepted verify step's KDA/mamba state (chain
|
||||
layout: step index = commit_lens - 1) into the persistent caches."""
|
||||
if not self._need_mamba_verify_commit:
|
||||
return
|
||||
# Chain layout only: step index = commit_lens - 1. A tree (topk > 1)
|
||||
# layout would need the accept-index mapping the shared spec_utils
|
||||
# commit helper does.
|
||||
assert self.server_args.speculative_eagle_topk in (None, 1)
|
||||
attn_backend = self.target_worker.model_runner.attn_backend
|
||||
|
||||
last_correct_step_indices = commit_lens.to(torch.int64) - 1
|
||||
mamba_steps_to_track = None
|
||||
|
||||
if batch.mamba_track_indices is not None:
|
||||
mamba_track_interval = self.server_args.mamba_track_interval
|
||||
to_track_mask = (
|
||||
seq_lens_pre_verify // mamba_track_interval
|
||||
!= seq_lens_post_verify // mamba_track_interval
|
||||
)
|
||||
tracking_point = (
|
||||
seq_lens_post_verify // mamba_track_interval * mamba_track_interval
|
||||
)
|
||||
to_track_ith = torch.clamp(tracking_point - seq_lens_pre_verify - 1, min=0)
|
||||
can_track_mask = to_track_mask & (
|
||||
to_track_ith < commit_lens.to(to_track_ith.dtype)
|
||||
)
|
||||
mamba_steps_to_track = torch.where(
|
||||
can_track_mask,
|
||||
to_track_ith.to(torch.int64),
|
||||
torch.full_like(to_track_ith, -1, dtype=torch.int64),
|
||||
)
|
||||
|
||||
attn_backend.update_mamba_state_after_mtp_verify(
|
||||
last_correct_step_indices=last_correct_step_indices,
|
||||
mamba_track_indices=batch.mamba_track_indices,
|
||||
mamba_steps_to_track=mamba_steps_to_track,
|
||||
model=self.target_worker.model_runner.model,
|
||||
)
|
||||
|
||||
def get_confidence_budget_prepare(self):
|
||||
return self._verify_planner.confidence_budget_prepare()
|
||||
|
||||
Reference in New Issue
Block a user