Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 35 additions & 10 deletions python/sglang/srt/arg_groups/attention_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,22 @@
logger = logging.getLogger(__name__)


def _replayssm_is_kda(model_config: Any) -> bool:
"""Arg-time predicate for the KDA gate the pool later exposes as
`MambaPool.replayssm_is_kda`.

Reading `mamba2_cache_params.is_kda` -- what the pool itself reads -- is not
an option here: building the cache params needs the parallel widths, and
this handler runs before the process groups exist. `KimiLinearCacheParams`
is the only params class whose `is_kda` is True and it has exactly two
producers, `KimiLinearConfig` and `BailingHybridConfig` with `use_kda`,
which is precisely what `kimi_linear_config` matches.
"""
from sglang.srt.configs.hybrid_arch import kimi_linear_config

return kimi_linear_config(model_config) is not None


def handle_attention_backend_compatibility(server_args: Any):

cfg = resolving_view(server_args)
Expand Down Expand Up @@ -333,12 +349,17 @@ def handle_linear_attn_backend(server_args: Any):
# the COW copy-into-slot path resets the ring cursor) -- so the
# --disable-radix-cache requirement is dropped.
#
# Slice 2b only wires the no_buffer mamba scheduler strategy (the
# default). The extra_buffer strategy donates the track snapshot via
# `donate_mamba_ping_pong_slot` with a separate ping-pong slot swap that
# does NOT route through MambaPool.copy_from, so the ReplaySSM ring
# cursor of the donated/kept slot would not be reset there. Handling
# that donation path is a follow-up; for now require no_buffer.
# Both mamba scheduler strategies are wired for GDN. no_buffer resets the
# ReplaySSM ring cursor in MambaPool.copy_from; under extra_buffer every slot
# enters a ping-pong buffer with a zero cursor, and the decode kernel
# force-flushes the ring into temporal[slot] at the radix track boundary, so
# a donated track snapshot is a complete checkpoint.
#
# KDA still requires no_buffer: HybridLinearAttnBackend only builds
# `replayssm_force_flush` when not is_kda, so a KDA track snapshot would be
# taken with ring entries still unfolded -- the ping-pong swap copies the SSM
# state but not the pending ring. Lift this once KDA force-flushes at radix
# tracking boundaries.
if cfg.enable_linear_replayssm:
if decode not in {"triton", "helion"}:
raise ValueError(
Expand All @@ -347,11 +368,15 @@ def handle_linear_attn_backend(server_args: Any):
f"--linear-attn-decode-backend={decode!r}."
)

if mamba_extra_buffer_of(resolved_view(server_args)):
Comment thread
yuan-luo marked this conversation as resolved.
if mamba_extra_buffer_of(resolved_view(server_args)) and _replayssm_is_kda(
model_config_of(server_args)
):
raise ValueError(
"--enable-linear-replayssm requires --mamba-radix-cache-strategy "
"no_buffer (the default); the extra_buffer ping-pong "
"donation path is not yet supported (follow-up). Got "
"--enable-linear-replayssm on a KDA model requires "
"--mamba-radix-cache-strategy no_buffer (the default); the "
"extra_buffer ping-pong donation path does not force-flush the "
"KDA ReplaySSM ring at radix tracking boundaries, so the donated "
"snapshot would miss pending ring entries. Got "
f"--mamba-radix-cache-strategy={cfg.mamba_radix_cache_strategy!r}."
)
if cfg.disaggregation_mode != "null":
Expand Down
21 changes: 21 additions & 0 deletions python/sglang/srt/mem_cache/memory_pool.py
Original file line number Diff line number Diff line change
Expand Up @@ -1466,6 +1466,12 @@ def _alloc_ping_pong_buffer(self, req: Req):
"Not enough space for mamba ping pong idx, "
"try to increase --mamba-full-memory-ratio."
)
# ReplaySSM: a recycled slot carries its previous owner's ring cursor,
# and `replayssm_write_pos` is documented as "reset on slot (re)alloc".
# `alloc` does this for the live decode slot; do the same here so a
# ping-pong slot also enters the buffer with an empty ring.
if self.mamba_pool.replayssm_write_pos is not None:
self.mamba_pool.replayssm_write_pos[slots] = 0
buf = torch.full(
(self.mamba_ping_pong_track_buffer_size,),
-1,
Expand All @@ -1489,6 +1495,16 @@ def set_mamba_ping_pong_slot(self, req: Req, idx: int, value):
set_mamba_track_indices_from_reqs reads correct slot indices.
"""
req.kv.mamba_ping_pong_track_buffer[idx] = value
# ReplaySSM: this is the single point where an already-allocated slot
# enters a ping-pong buffer (the donate swap and both lazy on-demand
# paths), so resetting here upholds the invariant for all of them. A
# device-side id must not be compared on the host -- that would force a
# cudaStreamSynchronize -- and only the "clear this entry" callers pass
# a host-side -1, so branch on the type rather than the value.
write_pos_buf = self.mamba_pool.replayssm_write_pos
if write_pos_buf is not None:
if isinstance(value, torch.Tensor) or value >= 0:
write_pos_buf[value] = 0
self.req_index_to_mamba_ping_pong_track_buffer_mapping[req.kv.req_pool_idx] = (
req.kv.mamba_ping_pong_track_buffer
)
Expand All @@ -1513,6 +1529,11 @@ def donate_mamba_ping_pong_slot(
f"next_track_idx={req.kv.mamba_next_track_idx}, "
f"rid={req.rid}"
)
# The ReplaySSM ring cursor of the donated slot is already 0 (it was
# reset when the slot entered the buffer and only live decode slots
# advance the cursor), and set_mamba_ping_pong_slot resets new_slot as
# it goes in, so the donated checkpoint and the replacement tracking
# window are both consistent without an explicit reset here.
self.set_mamba_ping_pong_slot(req, donate_idx, new_slot[0])
return mamba_value_donated

Expand Down
Loading