feat(kvcache): checkpoint decode state at the tool-call opener token

This commit is contained in:
Xiaoze Fan
2026-08-13 18:32:02 +00:00
parent ee54f638fd
commit ea5348b4e3
10 changed files with 160 additions and 8 deletions
+7
View File
@@ -53,6 +53,13 @@ class Req:
mamba_restore_src: int | None = None # on a prefix hit: tree snapshot slot to COW into the live slot (first chunk only)
swa_evicted_seqlen: int = 0 # SWA radix: positions < this had their swa KV freed (slid out of window) during decode
decode_batch_idx: int = 0 # SWA radix: # of decode forwards done; the proactive free_swa skips the first (overlap guard)
# Set once, at the first sampled tool-call opener token (scheduler detection): the state
# length just after that token (its index + 1). A client-side rewrite of the echoed tool
# call diverges strictly after this point, so it is the deepest reuse boundary that
# survives such a rewrite. GDN: the state is frozen into a ping-pong slot when cached_len
# reaches it (snapshot_toolcall_anchor) and donated at finish. SWA: caps the proactive
# out-of-window eviction so the window ending here stays resumable.
toolcall_anchor_len: int | None = None
# Abort arrived while this request's forward was in flight (overlap scheduling). The abort
# handler must not free resources under an in-flight forward; it sets this flag and
# _process_last_data frees the request when the batch drains (after copy_done.synchronize).
+4 -1
View File
@@ -381,7 +381,10 @@ def _swa_per_req_swa_floor(config) -> int:
window = next(g.sliding_window for g in config.model_config.kv_cache_group_specs() if g.is_swa)
ps = config.page_size
locked = ((window + _SWA_RETAIN_GAP + ps - 1) // ps) * ps
return locked + window + _SWA_EVICTION_INTERVAL + 2 * ps
floor = locked + window + _SWA_EVICTION_INTERVAL + 2 * ps
if getattr(config, "special_token_ckpt", False):
floor += window + _SWA_RETAIN_GAP + _SWA_EVICTION_INTERVAL
return floor
def _swa_pool_floor(config) -> int:
+68 -2
View File
@@ -141,6 +141,33 @@ class CacheManager:
self.linear_state_pool.free(er.mamba_slots)
self._free(er.kv_indices)
def snapshot_toolcall_anchor(self, reqs: List[Req]) -> None:
"""Freeze each decoding request's GDN state at its tool-call anchor, into the ping-pong
slot that is idle during decode (the kernel-side ×CHUNK track only runs on prefill
extends). Must run on the engine stream before the current step's kernels: cached_len
equals the anchor exactly when every enqueued step up to the anchor-consuming one has
been issued and the next (current) one has not, so the copy lands between them in
stream order. Reuses ``mamba_last_track_seqlen`` as the pending-donate mark -- the
prefill track's own pending freeze was consumed by the prefill-commit ``cache_req``
before any decode drain could set an anchor."""
if not self.is_hybrid:
return
pool = self.linear_state_pool
for r in reqs:
a = r.toolcall_anchor_len
if (
a is None
or r.mamba_ping_pong is None
or r.mamba_last_track_seqlen is not None
or r.cached_len != a
or align_down(a, self.page_size) != a
):
continue
dst = r.mamba_ping_pong[r.mamba_next_track_idx]
pool.copy_from(r.linear_slot_idx, dst)
r.mamba_last_track_seqlen = a
r.mamba_next_track_idx = 1 - r.mamba_next_track_idx
def maybe_free_swa_out_of_window(self, reqs: List[Req], *, forward_iter: int) -> None:
"""Proactively free each decoding request's now-out-of-window SWA slots, bounding its swa
footprint to ~one window so a smaller-than-full swa pool (swa_full_tokens_ratio<1) stays
@@ -158,6 +185,22 @@ class CacheManager:
continue # overlap guard: extend forward may still be running
floor = req.cache_handle.cached_len # reused prefix -> its swa is tree-owned, not ours
threshold = (req.device_len - 1) - window - self.page_size
if req.toolcall_anchor_len is not None:
# Keep the window ending at the anchor resumable: a client-side rewrite of the
# echoed tool call forks after the anchor, and a resume there needs
# [anchor - window, anchor) live. The finish-insert then adopts (rather than
# tombstones) these never-evicted slots; they stay unlocked, so real pool
# pressure can still reclaim them (same soft retention as the prompt-end pin).
cap = req.toolcall_anchor_len - window - _SWA_RETAIN_GAP
if threshold - cap > window + _SWA_RETAIN_GAP:
# The decode ran on far past the anchor (a tool call is normally within
# tens of tokens of the end). Holding the cap would grow this request's
# live swa without bound ("SWA pool exhausted" is unhandled) -- drop the
# anchor and let normal eviction resume. This bound is what the
# anchor-retention term in _swa_per_req_swa_floor sizes the pool for.
req.toolcall_anchor_len = None
else:
threshold = min(threshold, cap)
new_evicted = align_down(threshold, self.page_size)
start = max(req.swa_evicted_seqlen, floor)
if new_evicted > start:
@@ -308,6 +351,29 @@ class CacheManager:
return
if finished:
# A pending freeze (the tool-call anchor, or a prefill ×64 track the request
# finished too early to chunk-commit) is a strictly shorter prefix than the live
# donate below: insert it first and advance the dedup-free floor to its boundary
# -- [prefix_len, L) is now tree-owned by the donated node, so only [old, prefix_len)
# is this request's dup to free. The frozen slot is consumed either way (taken by
# the tree or freed here) and both ping-pong refs are dropped before
# _free_req_slots so nothing double-frees.
free_upto = old_handle.cached_len
L = req.mamba_last_track_seqlen
if (
L is not None
and 0 < L <= req.cached_len
and align_down(L, self.page_size) == L
and req.mamba_ping_pong is not None
):
frozen_idx = 1 - req.mamba_next_track_idx
frozen = req.mamba_ping_pong[frozen_idx]
prefix_len, mamba_exist = self.prefix_cache.insert(
req.input_ids[:L], page_indices[:L], frozen)
pool.free([s for s in req.mamba_ping_pong if mamba_exist or s != frozen])
req.mamba_ping_pong = None
self._free(page_indices[free_upto : max(free_upto, prefix_len)])
free_upto = max(free_upto, L)
# Donate the live slot (final full-sequence state). The live state is at cached_len;
# only attach it when cached_len is itself the page-aligned node boundary (always for
# page_size==1). For page_size>1 a non-aligned cached_len would attach an over-advanced
@@ -319,11 +385,11 @@ class CacheManager:
prefix_len, mamba_exist = self.prefix_cache.insert(
req.input_ids[:insert_len], page_indices[:insert_len], req.linear_slot_idx)
self.unlock(old_handle)
self._free(page_indices[old_handle.cached_len : prefix_len])
self._free(page_indices[free_upto : max(free_upto, prefix_len)])
keep_live = not mamba_exist # tree now owns linear_slot_idx
else:
self.unlock(old_handle)
self._free(page_indices[old_handle.cached_len :])
self._free(page_indices[free_upto :])
self._free_req_slots(req, keep_live=keep_live)
return
+1
View File
@@ -17,6 +17,7 @@ class SchedulerConfig(EngineConfig):
cache_type: str = "radix"
offline_mode: bool = False
decode_log_interval: int = 40
special_token_ckpt: bool = False
# networking config
_unique_suffix: str = field(default_factory=_get_pid_suffix)
+24 -1
View File
@@ -18,7 +18,12 @@ from freetoken.message import (
PromptAdmittedMsg,
UserMsg,
)
from freetoken.utils import init_logger, load_eos_token_ids, load_tokenizer
from freetoken.utils import (
init_logger,
load_eos_token_ids,
load_tokenizer,
load_toolcall_anchor_id,
)
from .cache import CacheManager
from .config import SchedulerConfig
@@ -106,6 +111,16 @@ class Scheduler(SchedulerIOMixin):
self._pending_rebuild: CacheRebuildBackendMsg | None = None
self.tokenizer = load_tokenizer(config.model_path)
self.eos_token_ids = load_eos_token_ids(config.model_path, self.tokenizer)
self.toolcall_anchor_id = None
if config.special_token_ckpt and (
self.cache_manager.is_hybrid or self.cache_manager.is_swa
):
from freetoken.server.function_call_parser import toolcall_opener_for
self.toolcall_anchor_id = load_toolcall_anchor_id(
self.tokenizer,
toolcall_opener_for(getattr(config, "tool_call_parser", "")),
)
self.token_pool = self.table_manager.token_pool
# Floor the prefill chunk by the cache manager's cap (DSV4: ~half the window pool) so a
# sliding-window cache chunks long prompts and frees out-of-window pages between chunks
@@ -336,6 +351,12 @@ class Scheduler(SchedulerIOMixin):
if finished
else None
)
if (
next_token == self.toolcall_anchor_id
and req.toolcall_anchor_len is None
and not finished
):
req.toolcall_anchor_len = req.input_ids.numel()
reply.append(
DetokenizeMsg(
uid=req.uid,
@@ -842,6 +863,8 @@ class Scheduler(SchedulerIOMixin):
def _forward(self, forward_input: ForwardInput) -> ForwardOutput:
batch, sample_args, input_mapping, output_mapping = forward_input
batch.input_ids = self.token_pool[input_mapping]
if self.toolcall_anchor_id is not None and not batch.is_prefill:
self.cache_manager.snapshot_toolcall_anchor(batch.reqs)
forward_output = self.engine.forward_batch(batch, sample_args)
self.token_pool[output_mapping] = forward_output.next_tokens_gpu
self.decode_manager.filter_reqs(forward_input.batch.reqs)
+15
View File
@@ -543,6 +543,21 @@ def parse_args(
),
)
parser.add_argument(
"--enable-special-token-ckpt",
action="store_true",
dest="special_token_ckpt",
default=ServerArgs.special_token_ckpt,
help=(
"Checkpoint decode state at special tokens (currently the tool-call opener). "
"When a GDN-hybrid or SWA model samples its tool-call opener token, the "
"scheduler preserves a reuse point just after it (GDN: a state snapshot "
"donated to the prefix cache; SWA: the trailing window is kept resumable), so "
"a client that rewrites the echoed tool call only invalidates the call body, "
"not the turn."
),
)
parser.add_argument(
"--moe-prefill-hit-d2d",
action="store_true",
@@ -239,6 +239,14 @@ class BaseFormatDetector(ABC):
# when this holds; otherwise they send the full arguments once at close.
args_fragments_prefix_stable = True
# The marker that UNIQUELY opens a tool-call block in this wire format, or None when
# no such marker exists (gpt-oss: <|channel|> opens every channel header, tool call or
# not; DSV32/DSV4: the DSML opener is a multi-piece composite). Consumed by the
# scheduler's special-token checkpoint, which only uses it when the tokenizer encodes
# it as a single token. Often equals bot_token, but declared separately because
# bot_token is a parse trigger, not a uniqueness claim.
toolcall_opener: str | None = None
def __init__(self):
# Streaming state management
# Buffer for accumulating incomplete patterns that arrive across multiple streaming chunks
@@ -899,7 +907,7 @@ class Qwen25Detector(BaseFormatDetector):
Reference: https://huggingface.co/Qwen/Qwen2.5-0.5B-Instruct?chat_template=default
"""
toolcall_opener = "<tool_call>"
def __init__(self):
"""
Initializes the detector with necessary state variables.
@@ -1091,7 +1099,7 @@ class Llama32Detector(BaseFormatDetector):
<python_tag>{"name":"xxx", "arguments":{...}}
```
"""
toolcall_opener = "<|python_tag|>"
def __init__(self):
super().__init__()
self.bot_token = "<|python_tag|>"
@@ -1181,7 +1189,7 @@ class Glm47Detector(BaseFormatDetector):
Reference: https://github.com/vllm-project/vllm/blob/main/vllm/tool_parsers/glm4_moe_tool_parser.py
"""
toolcall_opener = "<tool_call>"
def __init__(self):
super().__init__()
self.bot_token = "<tool_call>"
@@ -1859,6 +1867,7 @@ class DeepSeekV32Detector(BaseFormatDetector):
class Qwen3CoderDetector(InvokeParamStreamMixin, BaseFormatDetector):
toolcall_opener = "<tool_call>"
_ps_trim = "\n"
_ps_trim_single = True
_ps_missing_type = "string"
@@ -2024,7 +2033,7 @@ class Qwen3CoderDetector(InvokeParamStreamMixin, BaseFormatDetector):
class Gemma4Detector(BaseFormatDetector):
"""FreeToken serving adapter for Gemma4's compact tool-call format."""
toolcall_opener = "<|tool_call>"
def __init__(self):
super().__init__()
self.bot_token = "<|tool_call>"
@@ -2313,6 +2322,7 @@ class Gemma4Detector(BaseFormatDetector):
class MiniMaxDetector(InvokeParamStreamMixin, BaseFormatDetector):
toolcall_opener = "<minimax:tool_call>"
_ps_trim = "\n"
_ps_trim_single = False
_ps_missing_type = "loose"
@@ -2740,6 +2750,13 @@ class FunctionCallParser:
SUPPORTED_TOOL_CALL_PARSERS = list(FunctionCallParser.ToolCallParserEnum.keys())
def toolcall_opener_for(tool_call_parser: str) -> str | None:
"""The configured parser's unique tool-call opening marker, or None when the format has
no single unique opener (see ``BaseFormatDetector.toolcall_opener``)."""
detector = FunctionCallParser.ToolCallParserEnum.get(tool_call_parser)
return detector.toolcall_opener if detector is not None else None
def _coerce_tools(tools: List[Any] | None) -> List[Tool]:
return [_coerce_tool(tool) for tool in tools or []]
+2
View File
@@ -11,6 +11,7 @@ from .hf import (
load_eos_token_ids,
load_generation_sampling,
load_tokenizer,
load_toolcall_anchor_id,
)
from .logger import init_logger
from .misc import UNSET, Unset, align_ceil, align_down, call_if_main, div_ceil, div_even, mem_GB
@@ -31,6 +32,7 @@ __all__ = [
"load_eos_token_ids",
"load_generation_sampling",
"load_tokenizer",
"load_toolcall_anchor_id",
"init_logger",
"is_arch_supported",
"is_sm90_family",
+17
View File
@@ -72,6 +72,23 @@ def load_eos_token_ids(
return frozenset(ids)
def load_toolcall_anchor_id(
tokenizer: PreTrainedTokenizerBase, opener: str | None
) -> int | None:
"""The single token id of ``opener`` -- the wire format's unique tool-call opening
marker, declared by the model's detector (``BaseFormatDetector.toolcall_opener``).
None when there is no opener or the tokenizer spells it with more than one token:
the scheduler's special-token checkpoint matches sampled ids one at a time, so only
a one-token opener can anchor."""
if not opener:
return None
try:
ids = tokenizer.encode(opener, add_special_tokens=False)
except Exception:
return None
return int(ids[0]) if len(ids) == 1 else None
def load_generation_sampling(model_path: str) -> dict[str, Any]:
"""Recommended sampling defaults from ``generation_config.json`` (sglang's
``sampling_defaults='model'``). Returns ``{temperature, top_k, top_p}`` for the keys
@@ -59,6 +59,7 @@ def _setup():
prefill_manager=pm,
finished_reqs=set(),
eos_token_ids=set(),
toolcall_anchor_id=None,
config=SimpleNamespace(page_size=1),
status_reporter=SimpleNamespace(report_batch=lambda *_, **__: None),
send_result=sent.extend,