[mem_cache] Move mamba state and retraction_backup into ReqKvInfo (#37164)
This commit is contained in:
@@ -285,7 +285,7 @@ class DecodeKVCacheOffloadManager:
|
|||||||
self.token_to_kv_pool_allocator.free(overalloc_indices)
|
self.token_to_kv_pool_allocator.free(overalloc_indices)
|
||||||
|
|
||||||
self.req_to_token_pool.free(req)
|
self.req_to_token_pool.free(req)
|
||||||
req.kv.mark_released()
|
req.kv.mark_kv_released()
|
||||||
self.tree_cache.protected_size_ -= len(req.prefix_indices)
|
self.tree_cache.protected_size_ -= len(req.prefix_indices)
|
||||||
self.offloaded_state.pop(req, None)
|
self.offloaded_state.pop(req, None)
|
||||||
|
|
||||||
|
|||||||
@@ -1036,7 +1036,7 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
else:
|
else:
|
||||||
logger.warning(error_message)
|
logger.warning(error_message)
|
||||||
req.time_stats.trace_ctx.abort(abort_info={"reason": error_message})
|
req.time_stats.trace_ctx.abort(abort_info={"reason": error_message})
|
||||||
if req.kv.is_held or req.mamba_pool_idx is not None:
|
if req.kv.holds_kv or req.kv.holds_mamba:
|
||||||
release_kv_cache(req, self.tree_cache)
|
release_kv_cache(req, self.tree_cache)
|
||||||
maybe_release_metadata_buffer(req, self.req_to_metadata_buffer_idx_allocator)
|
maybe_release_metadata_buffer(req, self.req_to_metadata_buffer_idx_allocator)
|
||||||
req.pending_bootstrap = False
|
req.pending_bootstrap = False
|
||||||
|
|||||||
@@ -222,7 +222,7 @@ class MlxAuxiliaryStateReqToTokenPool(ReqToTokenPool):
|
|||||||
|
|
||||||
* Radix cache enabled: the ``MlxAuxiliaryStateComponent`` of the unified
|
* Radix cache enabled: the ``MlxAuxiliaryStateComponent`` of the unified
|
||||||
radix cache owns release — on finish it either frees the slot or
|
radix cache owns release — on finish it either frees the slot or
|
||||||
transfers it to the tree, nulling ``req.mamba_pool_idx`` before the
|
transfers it to the tree, nulling ``req.kv.mamba_pool_idx`` before the
|
||||||
request row is freed. The pool must NOT free auxiliary slots itself.
|
request row is freed. The pool must NOT free auxiliary slots itself.
|
||||||
* Radix cache disabled (``ChunkCache``): no tree component exists, and
|
* Radix cache disabled (``ChunkCache``): no tree component exists, and
|
||||||
``release_kv_cache``'s ``free_mamba_cache`` fallback is gated on
|
``release_kv_cache``'s ``free_mamba_cache`` fallback is gated on
|
||||||
@@ -270,13 +270,13 @@ class MlxAuxiliaryStateReqToTokenPool(ReqToTokenPool):
|
|||||||
|
|
||||||
auxiliary_state_indices = []
|
auxiliary_state_indices = []
|
||||||
for req in reqs:
|
for req in reqs:
|
||||||
if getattr(req, "mamba_pool_idx", None) is not None:
|
if req.kv.holds_mamba:
|
||||||
mid = req.mamba_pool_idx
|
mid = req.kv.mamba_pool_idx
|
||||||
else:
|
else:
|
||||||
allocated = self.auxiliary_state_pool.alloc(1)
|
allocated = self.auxiliary_state_pool.alloc(1)
|
||||||
assert allocated is not None, "Not enough MLX auxiliary state slots"
|
assert allocated is not None, "Not enough MLX auxiliary state slots"
|
||||||
mid = allocated[0]
|
mid = allocated[0]
|
||||||
req.mamba_pool_idx = mid
|
req.kv.mamba_pool_idx = mid
|
||||||
auxiliary_state_indices.append(mid.to(dtype=torch.int32))
|
auxiliary_state_indices.append(mid.to(dtype=torch.int32))
|
||||||
self.req_index_to_auxiliary_state_index_mapping[select_index] = torch.stack(
|
self.req_index_to_auxiliary_state_index_mapping[select_index] = torch.stack(
|
||||||
auxiliary_state_indices
|
auxiliary_state_indices
|
||||||
@@ -293,16 +293,16 @@ class MlxAuxiliaryStateReqToTokenPool(ReqToTokenPool):
|
|||||||
return 0
|
return 0
|
||||||
|
|
||||||
def free_mamba_cache(self, req, mamba_ping_pong_track_buffer_to_keep=None):
|
def free_mamba_cache(self, req, mamba_ping_pong_track_buffer_to_keep=None):
|
||||||
if getattr(req, "mamba_pool_idx", None) is not None:
|
if req.kv.holds_mamba:
|
||||||
self.auxiliary_state_pool.free(req.mamba_pool_idx.unsqueeze(0))
|
self.auxiliary_state_pool.free(req.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
track_buffer = getattr(req, "mamba_ping_pong_track_buffer", None)
|
track_buffer = req.kv.mamba_ping_pong_track_buffer
|
||||||
if track_buffer is not None:
|
if track_buffer is not None:
|
||||||
if mamba_ping_pong_track_buffer_to_keep is None:
|
if mamba_ping_pong_track_buffer_to_keep is None:
|
||||||
self.auxiliary_state_pool.free(track_buffer)
|
self.auxiliary_state_pool.free(track_buffer)
|
||||||
req.mamba_ping_pong_track_buffer = None
|
req.kv.mamba_ping_pong_track_buffer = None
|
||||||
req.mamba_next_track_idx = None
|
req.kv.mamba_next_track_idx = None
|
||||||
req.mamba_last_track_idx = None
|
req.kv.mamba_last_track_idx = None
|
||||||
|
|
||||||
def free_auxiliary_state_cache(self, req, track_buffer_to_keep=None):
|
def free_auxiliary_state_cache(self, req, track_buffer_to_keep=None):
|
||||||
self.free_mamba_cache(
|
self.free_mamba_cache(
|
||||||
@@ -314,7 +314,7 @@ class MlxAuxiliaryStateReqToTokenPool(ReqToTokenPool):
|
|||||||
if self._owns_auxiliary_state_release:
|
if self._owns_auxiliary_state_release:
|
||||||
# No-radix configuration: nothing else will ever release the
|
# No-radix configuration: nothing else will ever release the
|
||||||
# auxiliary slot, so return it with the request row. Keyed on
|
# auxiliary slot, so return it with the request row. Keyed on
|
||||||
# req.mamba_pool_idx (None-safe, nulled by free_mamba_cache), NOT
|
# req.kv.mamba_pool_idx (None-safe, nulled by free_mamba_cache), NOT
|
||||||
# on req_index_to_auxiliary_state_index_mapping, which may point
|
# on req_index_to_auxiliary_state_index_mapping, which may point
|
||||||
# at a slot the radix tree owns.
|
# at a slot the radix tree owns.
|
||||||
self.free_mamba_cache(req)
|
self.free_mamba_cache(req)
|
||||||
@@ -347,13 +347,13 @@ class MlxAuxiliaryStateComponent(MambaComponent):
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _tracked_value(req) -> tuple[object | None, bool]:
|
def _tracked_value(req) -> tuple[object | None, bool]:
|
||||||
track_buffer = getattr(req, "mamba_ping_pong_track_buffer", None)
|
track_buffer = req.kv.mamba_ping_pong_track_buffer
|
||||||
track_len = getattr(req, "mamba_last_track_seqlen", None)
|
track_len = req.kv.mamba_last_track_seqlen
|
||||||
if track_buffer is not None and track_len is not None:
|
if track_buffer is not None and track_len is not None:
|
||||||
return track_buffer[0].unsqueeze(-1).clone(), True
|
return track_buffer[0].unsqueeze(-1).clone(), True
|
||||||
if getattr(req, "mamba_pool_idx", None) is None:
|
if not req.kv.holds_mamba:
|
||||||
return None, False
|
return None, False
|
||||||
return req.mamba_pool_idx.unsqueeze(-1).clone(), False
|
return req.kv.mamba_pool_idx.unsqueeze(-1).clone(), False
|
||||||
|
|
||||||
def prepare_for_caching_req(
|
def prepare_for_caching_req(
|
||||||
self,
|
self,
|
||||||
@@ -362,7 +362,7 @@ class MlxAuxiliaryStateComponent(MambaComponent):
|
|||||||
token_ids_len: int,
|
token_ids_len: int,
|
||||||
is_finished: bool,
|
is_finished: bool,
|
||||||
) -> int | None:
|
) -> int | None:
|
||||||
cache_len = getattr(req, "mamba_last_track_seqlen", None)
|
cache_len = req.kv.mamba_last_track_seqlen
|
||||||
auxiliary_value, uses_track_slot = self._tracked_value(req)
|
auxiliary_value, uses_track_slot = self._tracked_value(req)
|
||||||
setattr(insert_params, "mlx_auxiliary_state_uses_track_slot", uses_track_slot)
|
setattr(insert_params, "mlx_auxiliary_state_uses_track_slot", uses_track_slot)
|
||||||
|
|
||||||
@@ -408,13 +408,13 @@ class MlxAuxiliaryStateComponent(MambaComponent):
|
|||||||
if bool(
|
if bool(
|
||||||
getattr(insert_params, "mlx_auxiliary_state_uses_track_slot", False)
|
getattr(insert_params, "mlx_auxiliary_state_uses_track_slot", False)
|
||||||
):
|
):
|
||||||
track_buffer = getattr(req, "mamba_ping_pong_track_buffer", None)
|
track_buffer = req.kv.mamba_ping_pong_track_buffer
|
||||||
if track_buffer is not None:
|
if track_buffer is not None:
|
||||||
self.cache.req_to_token_pool.auxiliary_state_pool.free(track_buffer)
|
self.cache.req_to_token_pool.auxiliary_state_pool.free(track_buffer)
|
||||||
req.mamba_ping_pong_track_buffer = None
|
req.kv.mamba_ping_pong_track_buffer = None
|
||||||
req.mamba_next_track_idx = None
|
req.kv.mamba_next_track_idx = None
|
||||||
req.mamba_last_track_idx = None
|
req.kv.mamba_last_track_idx = None
|
||||||
req.mamba_last_track_seqlen = None
|
req.kv.mamba_last_track_seqlen = None
|
||||||
return
|
return
|
||||||
|
|
||||||
auxiliary_value_exists = (
|
auxiliary_value_exists = (
|
||||||
@@ -433,11 +433,11 @@ class MlxAuxiliaryStateComponent(MambaComponent):
|
|||||||
self.cache.req_to_token_pool.free_auxiliary_state_cache(req)
|
self.cache.req_to_token_pool.free_auxiliary_state_cache(req)
|
||||||
else:
|
else:
|
||||||
# The radix tree now owns the live auxiliary-state slot.
|
# The radix tree now owns the live auxiliary-state slot.
|
||||||
track_buffer = getattr(req, "mamba_ping_pong_track_buffer", None)
|
track_buffer = req.kv.mamba_ping_pong_track_buffer
|
||||||
if track_buffer is not None:
|
if track_buffer is not None:
|
||||||
self.cache.req_to_token_pool.auxiliary_state_pool.free(track_buffer)
|
self.cache.req_to_token_pool.auxiliary_state_pool.free(track_buffer)
|
||||||
req.mamba_ping_pong_track_buffer = None
|
req.kv.mamba_ping_pong_track_buffer = None
|
||||||
req.mamba_next_track_idx = None
|
req.kv.mamba_next_track_idx = None
|
||||||
req.mamba_last_track_idx = None
|
req.kv.mamba_last_track_idx = None
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
req.mamba_last_track_seqlen = None
|
req.kv.mamba_last_track_seqlen = None
|
||||||
|
|||||||
@@ -379,7 +379,7 @@ class MlxModelRunner:
|
|||||||
|
|
||||||
chunk_size = mamba_cache_chunk_size()
|
chunk_size = mamba_cache_chunk_size()
|
||||||
track_len = prefix_len + (new_token_count // chunk_size) * chunk_size
|
track_len = prefix_len + (new_token_count // chunk_size) * chunk_size
|
||||||
branching_len = getattr(req, "mamba_branching_seqlen", None)
|
branching_len = req.mamba_branching_seqlen
|
||||||
if (
|
if (
|
||||||
branching_len is not None
|
branching_len is not None
|
||||||
and prefix_len < branching_len <= prefix_len + new_token_count
|
and prefix_len < branching_len <= prefix_len + new_token_count
|
||||||
@@ -407,7 +407,7 @@ class MlxModelRunner:
|
|||||||
if pool is None or not hasattr(pool, "store_cache"):
|
if pool is None or not hasattr(pool, "store_cache"):
|
||||||
return
|
return
|
||||||
|
|
||||||
track_buffer = getattr(req, "mamba_ping_pong_track_buffer", None)
|
track_buffer = req.kv.mamba_ping_pong_track_buffer
|
||||||
if track_buffer is None:
|
if track_buffer is None:
|
||||||
track_buffer = pool.alloc(1)
|
track_buffer = pool.alloc(1)
|
||||||
if track_buffer is None:
|
if track_buffer is None:
|
||||||
@@ -416,16 +416,16 @@ class MlxModelRunner:
|
|||||||
"falling back to leaf-only auxiliary-state radix caching."
|
"falling back to leaf-only auxiliary-state radix caching."
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
req.mamba_ping_pong_track_buffer = track_buffer
|
req.kv.mamba_ping_pong_track_buffer = track_buffer
|
||||||
req.mamba_next_track_idx = 0
|
req.kv.mamba_next_track_idx = 0
|
||||||
req.mamba_last_track_idx = 0
|
req.kv.mamba_last_track_idx = 0
|
||||||
|
|
||||||
pool.store_cache(
|
pool.store_cache(
|
||||||
track_buffer[0],
|
track_buffer[0],
|
||||||
cache,
|
cache,
|
||||||
self._cache_layout.auxiliary_layer_indices,
|
self._cache_layout.auxiliary_layer_indices,
|
||||||
)
|
)
|
||||||
req.mamba_last_track_seqlen = track_len
|
req.kv.mamba_last_track_seqlen = track_len
|
||||||
|
|
||||||
def _cache_with_pool_backed_attention(
|
def _cache_with_pool_backed_attention(
|
||||||
self, prefix_slot_ids: list[int], prefix_len: int
|
self, prefix_slot_ids: list[int], prefix_len: int
|
||||||
@@ -905,7 +905,7 @@ class MlxModelRunner:
|
|||||||
"""
|
"""
|
||||||
prefix_len = len(prefix_slot_ids)
|
prefix_len = len(prefix_slot_ids)
|
||||||
if req is not None:
|
if req is not None:
|
||||||
req.mamba_last_track_seqlen = None
|
req.kv.mamba_last_track_seqlen = None
|
||||||
if self._enable_sampling:
|
if self._enable_sampling:
|
||||||
self._req_sampling[req_id] = (
|
self._req_sampling[req_id] = (
|
||||||
MlxSamplingParams.from_req(
|
MlxSamplingParams.from_req(
|
||||||
|
|||||||
@@ -172,7 +172,7 @@ class MlxTpModelWorker(TpModelWorker):
|
|||||||
self._mlx_runner.store_auxiliary_state_for_request(req.rid)
|
self._mlx_runner.store_auxiliary_state_for_request(req.rid)
|
||||||
# Prefer the just-snapshotted live auxiliary state for the final
|
# Prefer the just-snapshotted live auxiliary state for the final
|
||||||
# insert. Any older tracked slot is released during component cleanup.
|
# insert. Any older tracked slot is released during component cleanup.
|
||||||
req.mamba_last_track_seqlen = None
|
req.kv.mamba_last_track_seqlen = None
|
||||||
|
|
||||||
def _route_extend_request(self, rid: str, decoding_rids: set[str]) -> str:
|
def _route_extend_request(self, rid: str, decoding_rids: set[str]) -> str:
|
||||||
"""Classify a request within an extend / mixed batch.
|
"""Classify a request within an extend / mixed batch.
|
||||||
|
|||||||
@@ -816,11 +816,12 @@ class ReqLogprob:
|
|||||||
@dataclasses.dataclass(slots=True, kw_only=True)
|
@dataclasses.dataclass(slots=True, kw_only=True)
|
||||||
class ReqKvInfo:
|
class ReqKvInfo:
|
||||||
# Device KV a request holds outside the prefix cache. Always present on the Req;
|
# Device KV a request holds outside the prefix cache. Always present on the Req;
|
||||||
# whether any KV is held is `is_held` (a row is registered).
|
# whether any KV is held is `holds_kv` (a row is registered).
|
||||||
|
# Match observations and scheduling state stay on the Req itself.
|
||||||
req_pool_idx: Optional[int] = None # req_to_token row, the register for the slots
|
req_pool_idx: Optional[int] = None # req_to_token row, the register for the slots
|
||||||
|
|
||||||
# The request's own KV is [cache_protected_len, kv_allocated_len).
|
# The request's own KV is [cache_protected_len, kv_allocated_len).
|
||||||
cache_protected_len: int = 0 # tree cache owns [0, here) (matched or inserted)
|
cache_protected_len: int = 0 # Tree cache owns [0, here) (matched or inserted)
|
||||||
kv_committed_len: int = 0 # KV content committed up to here, <= kv_allocated_len
|
kv_committed_len: int = 0 # KV content committed up to here, <= kv_allocated_len
|
||||||
kv_allocated_len: int = 0
|
kv_allocated_len: int = 0
|
||||||
|
|
||||||
@@ -828,6 +829,21 @@ class ReqKvInfo:
|
|||||||
swa_evict_floor: int = 0 # [0, here) never window-evicted (prefill-aware SWA)
|
swa_evict_floor: int = 0 # [0, here) never window-evicted (prefill-aware SWA)
|
||||||
swa_evicted_seqlen: int = 0 # SWA eviction cursor
|
swa_evicted_seqlen: int = 0 # SWA eviction cursor
|
||||||
|
|
||||||
|
# Host-side KV backup the request holds across a retraction (unified cache).
|
||||||
|
retraction_backup: Optional[RetractionBackup] = None
|
||||||
|
|
||||||
|
# Mamba state: an independent resource; whether it is held is `holds_mamba`.
|
||||||
|
mamba_pool_idx: Optional[torch.Tensor] = None # shape (1)
|
||||||
|
mamba_ping_pong_track_buffer: Optional[torch.Tensor] = None # shape (2)
|
||||||
|
mamba_next_track_idx: Optional[int] = None # 0 or 1
|
||||||
|
mamba_last_track_idx: Optional[int] = None # 0 or 1
|
||||||
|
# Seq len of the last cached mamba state
|
||||||
|
mamba_last_track_seqlen: Optional[int] = None
|
||||||
|
# Deferred COW: source mamba pool index from radix cache node (copy on forward stream)
|
||||||
|
mamba_cow_src_index: Optional[torch.Tensor] = None
|
||||||
|
# Deferred clear: newly allocated mamba slot needs zeroing on forward stream
|
||||||
|
mamba_needs_clear: bool = False
|
||||||
|
|
||||||
def swa_dead_lo(self, page_size: int) -> int:
|
def swa_dead_lo(self, page_size: int) -> int:
|
||||||
# Lowest SWA position this request may free itself: above the tree-owned
|
# Lowest SWA position this request may free itself: above the tree-owned
|
||||||
# prefix and above the eviction shield, page-aligned upward.
|
# prefix and above the eviction shield, page-aligned upward.
|
||||||
@@ -837,14 +853,18 @@ class ReqKvInfo:
|
|||||||
return lo
|
return lo
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_held(self) -> bool:
|
def holds_kv(self) -> bool:
|
||||||
return self.req_pool_idx is not None
|
return self.req_pool_idx is not None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_released(self) -> bool:
|
def holds_mamba(self) -> bool:
|
||||||
|
return self.mamba_pool_idx is not None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_kv_released(self) -> bool:
|
||||||
return self.kv_allocated_len == 0 and self.swa_evicted_seqlen == 0
|
return self.kv_allocated_len == 0 and self.swa_evicted_seqlen == 0
|
||||||
|
|
||||||
def mark_released(self) -> None:
|
def mark_kv_released(self) -> None:
|
||||||
self.kv_allocated_len = 0
|
self.kv_allocated_len = 0
|
||||||
self.swa_evicted_seqlen = 0
|
self.swa_evicted_seqlen = 0
|
||||||
|
|
||||||
@@ -929,7 +949,6 @@ class Req(ReqDllmMixin):
|
|||||||
|
|
||||||
# For req-level memory management
|
# For req-level memory management
|
||||||
self.kv = ReqKvInfo()
|
self.kv = ReqKvInfo()
|
||||||
self.retraction_backup: Optional[RetractionBackup] = None
|
|
||||||
|
|
||||||
# for cross-encoder model
|
# for cross-encoder model
|
||||||
self.token_type_ids = token_type_ids
|
self.token_type_ids = token_type_ids
|
||||||
@@ -974,21 +993,6 @@ class Req(ReqDllmMixin):
|
|||||||
self.lora_id = lora_id
|
self.lora_id = lora_id
|
||||||
self.routing_key = routing_key
|
self.routing_key = routing_key
|
||||||
|
|
||||||
# Memory pool info
|
|
||||||
self.mamba_pool_idx: Optional[torch.Tensor] = None # shape (1)
|
|
||||||
self.mamba_ping_pong_track_buffer: Optional[torch.Tensor] = None # shape (2)
|
|
||||||
self.mamba_next_track_idx: Optional[int] = None # 0 or 1
|
|
||||||
self.mamba_last_track_idx: Optional[int] = None # 0 or 1
|
|
||||||
self.mamba_last_track_seqlen: Optional[int] = (
|
|
||||||
None # seq len of the last cached mamba state
|
|
||||||
)
|
|
||||||
# the branching point seqlen to track mamba state. If set, given by prefix match,
|
|
||||||
# it will be the tracked seqlen in the ping pong buffer for the right prefill pass.
|
|
||||||
self.mamba_branching_seqlen: Optional[int] = None
|
|
||||||
# Deferred COW: source mamba pool index from radix cache node (copy on forward stream)
|
|
||||||
self.mamba_cow_src_index: Optional[torch.Tensor] = None
|
|
||||||
# Deferred clear: newly allocated mamba slot needs zeroing on forward stream
|
|
||||||
self.mamba_needs_clear: bool = False
|
|
||||||
# Lazy extra buffer: skip radix cache insert when prealloc failed at
|
# Lazy extra buffer: skip radix cache insert when prealloc failed at
|
||||||
# boundary — the forward overwrites the only slot, corrupting the state.
|
# boundary — the forward overwrites the only slot, corrupting the state.
|
||||||
self.mamba_lazy_is_insert: bool = True
|
self.mamba_lazy_is_insert: bool = True
|
||||||
@@ -1041,6 +1045,10 @@ class Req(ReqDllmMixin):
|
|||||||
self.host_hit_length = 0
|
self.host_hit_length = 0
|
||||||
self.swa_host_hit_length = 0
|
self.swa_host_hit_length = 0
|
||||||
self.mamba_host_hit_length = 0
|
self.mamba_host_hit_length = 0
|
||||||
|
# The branching point seqlen to track mamba state. If set, given by prefix
|
||||||
|
# match, it will be the tracked seqlen in the ping pong buffer for the
|
||||||
|
# right prefill pass.
|
||||||
|
self.mamba_branching_seqlen: Optional[int] = None
|
||||||
# Total cached prefix length (on-device prefix_indices + host_hit_length),
|
# Total cached prefix length (on-device prefix_indices + host_hit_length),
|
||||||
# capped at the max allowed prefix. Set during prefix matching at schedule
|
# capped at the max allowed prefix. Set during prefix matching at schedule
|
||||||
# time and used to estimate uncached tokens / sort by longest prefix for
|
# time and used to estimate uncached tokens / sort by longest prefix for
|
||||||
@@ -1744,16 +1752,16 @@ class Req(ReqDllmMixin):
|
|||||||
self.temp_input_token_ids_logprobs_val = None
|
self.temp_input_token_ids_logprobs_val = None
|
||||||
self.temp_input_token_ids_logprobs_idx = None
|
self.temp_input_token_ids_logprobs_idx = None
|
||||||
self.inflight_middle_chunks = 0
|
self.inflight_middle_chunks = 0
|
||||||
self.mamba_pool_idx = None
|
self.kv.mamba_pool_idx = None
|
||||||
self.mamba_ping_pong_track_buffer = None
|
self.kv.mamba_ping_pong_track_buffer = None
|
||||||
self.mamba_next_track_idx = None
|
self.kv.mamba_next_track_idx = None
|
||||||
self.mamba_last_track_idx = None
|
self.kv.mamba_last_track_idx = None
|
||||||
self.mamba_last_track_seqlen = None
|
self.kv.mamba_last_track_seqlen = None
|
||||||
self.mamba_branching_seqlen = None
|
self.mamba_branching_seqlen = None
|
||||||
self.mamba_cow_src_index = None
|
self.kv.mamba_cow_src_index = None
|
||||||
self.mamba_needs_clear = False
|
self.kv.mamba_needs_clear = False
|
||||||
self.already_computed = 0
|
self.already_computed = 0
|
||||||
assert not self.kv.is_held, "expect it is already released"
|
assert not self.kv.holds_kv, "expect it is already released"
|
||||||
self.kv.kv_committed_len = 0
|
self.kv.kv_committed_len = 0
|
||||||
self.extend_batch_idx = 0
|
self.extend_batch_idx = 0
|
||||||
self.decode_batch_idx = 0
|
self.decode_batch_idx = 0
|
||||||
@@ -1784,34 +1792,34 @@ class Req(ReqDllmMixin):
|
|||||||
mamba_pool = self._mamba_pool_needing_backup(
|
mamba_pool = self._mamba_pool_needing_backup(
|
||||||
req_to_token_pool, token_to_kv_pool_allocator
|
req_to_token_pool, token_to_kv_pool_allocator
|
||||||
)
|
)
|
||||||
self.retraction_backup = RetractionBackup(
|
self.kv.retraction_backup = RetractionBackup(
|
||||||
cpu_tensors=token_to_kv_pool_allocator.get_cpu_copy(
|
cpu_tensors=token_to_kv_pool_allocator.get_cpu_copy(
|
||||||
token_indices, mamba_indices=self.mamba_pool_idx
|
token_indices, mamba_indices=self.kv.mamba_pool_idx
|
||||||
),
|
),
|
||||||
mamba_cpu=(
|
mamba_cpu=(
|
||||||
mamba_pool.get_cpu_copy(self.mamba_pool_idx.unsqueeze(0))
|
mamba_pool.get_cpu_copy(self.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
if mamba_pool is not None and self.mamba_pool_idx is not None
|
if mamba_pool is not None and self.kv.holds_mamba
|
||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def load_kv_cache(self, req_to_token_pool, token_to_kv_pool_allocator):
|
def load_kv_cache(self, req_to_token_pool, token_to_kv_pool_allocator):
|
||||||
assert self.retraction_backup is not None
|
assert self.kv.retraction_backup is not None
|
||||||
token_indices = req_to_token_pool.req_to_token[
|
token_indices = req_to_token_pool.req_to_token[
|
||||||
self.kv.req_pool_idx, : self.seqlen - 1
|
self.kv.req_pool_idx, : self.seqlen - 1
|
||||||
]
|
]
|
||||||
# Loads both the kv cache and mamba state if exists
|
# Loads both the kv cache and mamba state if exists
|
||||||
mamba_cpu = self.retraction_backup.mamba_cpu
|
mamba_cpu = self.kv.retraction_backup.mamba_cpu
|
||||||
if mamba_cpu is not None and self.mamba_pool_idx is not None:
|
if mamba_cpu is not None and self.kv.holds_mamba:
|
||||||
req_to_token_pool.mamba_pool.load_cpu_copy(
|
req_to_token_pool.mamba_pool.load_cpu_copy(
|
||||||
mamba_cpu, self.mamba_pool_idx.unsqueeze(0)
|
mamba_cpu, self.kv.mamba_pool_idx.unsqueeze(0)
|
||||||
)
|
)
|
||||||
token_to_kv_pool_allocator.load_cpu_copy(
|
token_to_kv_pool_allocator.load_cpu_copy(
|
||||||
self.retraction_backup.cpu_tensors,
|
self.kv.retraction_backup.cpu_tensors,
|
||||||
token_indices,
|
token_indices,
|
||||||
mamba_indices=self.mamba_pool_idx,
|
mamba_indices=self.kv.mamba_pool_idx,
|
||||||
)
|
)
|
||||||
self.retraction_backup = None
|
self.kv.retraction_backup = None
|
||||||
|
|
||||||
def build_rebootstrap_payload(self) -> dict:
|
def build_rebootstrap_payload(self) -> dict:
|
||||||
"""Build the prefill ``/generate`` payload that asks the original prefill
|
"""Build the prefill ``/generate`` payload that asks the original prefill
|
||||||
@@ -1964,7 +1972,11 @@ def set_mamba_track_indices_from_reqs(
|
|||||||
# gone through _alloc_ping_pong_buffer yet (e.g., spec v2 verify path).
|
# gone through _alloc_ping_pong_buffer yet (e.g., spec v2 verify path).
|
||||||
# Default to 0 (first ping-pong slot) to avoid TypeError.
|
# Default to 0 (first ping-pong slot) to avoid TypeError.
|
||||||
track_positions = [
|
track_positions = [
|
||||||
req.mamba_next_track_idx if req.mamba_next_track_idx is not None else 0
|
(
|
||||||
|
req.kv.mamba_next_track_idx
|
||||||
|
if req.kv.mamba_next_track_idx is not None
|
||||||
|
else 0
|
||||||
|
)
|
||||||
for req in batch.reqs
|
for req in batch.reqs
|
||||||
]
|
]
|
||||||
batch.mamba_track_buffer_indices = list(track_positions)
|
batch.mamba_track_buffer_indices = list(track_positions)
|
||||||
@@ -2171,8 +2183,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# For hybrid GDN prefix cache
|
# For hybrid GDN prefix cache
|
||||||
mamba_track_indices: torch.Tensor = None # shape: [b], int64
|
mamba_track_indices: torch.Tensor = None # shape: [b], int64
|
||||||
# Per-batch snapshot of the logical ping-pong positions selected for this
|
# Per-batch snapshot of the logical ping-pong positions selected for this
|
||||||
# forward (normally req.mamba_next_track_idx; spec may override it). Result
|
# forward (normally req.kv.mamba_next_track_idx; spec may override it). Result
|
||||||
# processing uses it to update req.mamba_last_track_idx, since both req-level
|
# processing uses it to update req.kv.mamba_last_track_idx, since both req-level
|
||||||
# indices may advance under overlap.
|
# indices may advance under overlap.
|
||||||
mamba_track_buffer_indices: Optional[List[int]] = None # shape: [b], 0 or 1
|
mamba_track_buffer_indices: Optional[List[int]] = None # shape: [b], 0 or 1
|
||||||
mamba_track_mask: torch.Tensor = None # shape: [b], bool
|
mamba_track_mask: torch.Tensor = None # shape: [b], bool
|
||||||
@@ -2694,7 +2706,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Collect mamba init info for deferred ops on forward stream
|
# Collect mamba init info for deferred ops on forward stream
|
||||||
if any(req.mamba_pool_idx is not None for req in reqs):
|
if any(req.kv.holds_mamba for req in reqs):
|
||||||
self._collect_deferred_mamba_cow_and_clear(reqs)
|
self._collect_deferred_mamba_cow_and_clear(reqs)
|
||||||
|
|
||||||
if self.model_config.is_encoder_decoder:
|
if self.model_config.is_encoder_decoder:
|
||||||
@@ -2736,7 +2748,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
return i + 1
|
return i + 1
|
||||||
|
|
||||||
mask = req.extend_range.length >= checkpoint_grid
|
mask = req.extend_range.length >= checkpoint_grid
|
||||||
track_index = req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item()
|
track_index = req.kv.mamba_ping_pong_track_buffer[
|
||||||
|
req.kv.mamba_next_track_idx
|
||||||
|
].item()
|
||||||
mamba_track_seqlen = -1
|
mamba_track_seqlen = -1
|
||||||
if mask:
|
if mask:
|
||||||
# mamba_track_seqlen is used to calculate the indices to track in
|
# mamba_track_seqlen is used to calculate the indices to track in
|
||||||
@@ -2769,11 +2783,11 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# In lazy mode, skip the swap — the second ping-pong slot is not
|
# In lazy mode, skip the swap — the second ping-pong slot is not
|
||||||
# allocated yet; it will be allocated on demand at the track boundary
|
# allocated yet; it will be allocated on demand at the track boundary
|
||||||
# in mamba_lazy_prealloc_at_boundary during prepare_for_decode.
|
# in mamba_lazy_prealloc_at_boundary during prepare_for_decode.
|
||||||
req.mamba_last_track_idx = req.mamba_next_track_idx
|
req.kv.mamba_last_track_idx = req.kv.mamba_next_track_idx
|
||||||
if not mamba_extra_buffer_lazy_enabled():
|
if not mamba_extra_buffer_lazy_enabled():
|
||||||
req.mamba_next_track_idx = (
|
req.kv.mamba_next_track_idx = (
|
||||||
self.req_to_token_pool.get_mamba_ping_pong_other_idx(
|
self.req_to_token_pool.get_mamba_ping_pong_other_idx(
|
||||||
req.mamba_next_track_idx
|
req.kv.mamba_next_track_idx
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
if req.mamba_branching_seqlen is not None:
|
if req.mamba_branching_seqlen is not None:
|
||||||
@@ -2792,7 +2806,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# See _force_track_h() for more details.
|
# See _force_track_h() for more details.
|
||||||
mamba_track_seqlen = _force_track_h(req.mamba_branching_seqlen)
|
mamba_track_seqlen = _force_track_h(req.mamba_branching_seqlen)
|
||||||
mamba_track_seqlen_aligned = req.mamba_branching_seqlen
|
mamba_track_seqlen_aligned = req.mamba_branching_seqlen
|
||||||
req.mamba_last_track_seqlen = mamba_track_seqlen_aligned
|
req.kv.mamba_last_track_seqlen = mamba_track_seqlen_aligned
|
||||||
|
|
||||||
return _MambaRadixCacheV2TrackEntry(
|
return _MambaRadixCacheV2TrackEntry(
|
||||||
track_mask=mask,
|
track_mask=mask,
|
||||||
@@ -2806,14 +2820,14 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
cow_dst_tensors = []
|
cow_dst_tensors = []
|
||||||
clear_tensors = []
|
clear_tensors = []
|
||||||
for req in reqs:
|
for req in reqs:
|
||||||
if req.mamba_cow_src_index is not None:
|
if req.kv.mamba_cow_src_index is not None:
|
||||||
cow_src_tensors.append(req.mamba_cow_src_index)
|
cow_src_tensors.append(req.kv.mamba_cow_src_index)
|
||||||
cow_dst_tensors.append(req.mamba_pool_idx.unsqueeze(0))
|
cow_dst_tensors.append(req.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req.mamba_cow_src_index = None
|
req.kv.mamba_cow_src_index = None
|
||||||
req.mamba_needs_clear = False
|
req.kv.mamba_needs_clear = False
|
||||||
elif req.mamba_needs_clear:
|
elif req.kv.mamba_needs_clear:
|
||||||
clear_tensors.append(req.mamba_pool_idx.unsqueeze(0))
|
clear_tensors.append(req.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req.mamba_needs_clear = False
|
req.kv.mamba_needs_clear = False
|
||||||
self.mamba_cow_src_indices = (
|
self.mamba_cow_src_indices = (
|
||||||
torch.cat(cow_src_tensors) if cow_src_tensors else None
|
torch.cat(cow_src_tensors) if cow_src_tensors else None
|
||||||
)
|
)
|
||||||
@@ -3100,12 +3114,12 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
"""
|
"""
|
||||||
pool = self.req_to_token_pool
|
pool = self.req_to_token_pool
|
||||||
for i, req in enumerate(self.reqs):
|
for i, req in enumerate(self.reqs):
|
||||||
buf = req.mamba_ping_pong_track_buffer
|
buf = req.kv.mamba_ping_pong_track_buffer
|
||||||
assert buf is not None
|
assert buf is not None
|
||||||
# Skip reqs not at a track boundary
|
# Skip reqs not at a track boundary
|
||||||
if self.seq_lens_cpu[i].item() % mamba_track_interval != 0:
|
if self.seq_lens_cpu[i].item() % mamba_track_interval != 0:
|
||||||
continue
|
continue
|
||||||
other_idx = 1 - req.mamba_next_track_idx
|
other_idx = 1 - req.kv.mamba_next_track_idx
|
||||||
if buf[other_idx].item() != -1:
|
if buf[other_idx].item() != -1:
|
||||||
# With overlap the previous forward's post-processing
|
# With overlap the previous forward's post-processing
|
||||||
# (which frees this slot) hasn't run yet. Skip.
|
# (which frees this slot) hasn't run yet. Skip.
|
||||||
@@ -3118,7 +3132,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
new_slot = pool.mamba_allocator.alloc(1)
|
new_slot = pool.mamba_allocator.alloc(1)
|
||||||
if new_slot is not None:
|
if new_slot is not None:
|
||||||
pool.set_mamba_ping_pong_slot(req, other_idx, new_slot[0])
|
pool.set_mamba_ping_pong_slot(req, other_idx, new_slot[0])
|
||||||
req.mamba_next_track_idx = other_idx
|
req.kv.mamba_next_track_idx = other_idx
|
||||||
|
|
||||||
def mamba_lazy_spec_prepare(self, mamba_track_interval: int, max_draft_tokens: int):
|
def mamba_lazy_spec_prepare(self, mamba_track_interval: int, max_draft_tokens: int):
|
||||||
"""Lazy-mode spec counterpart of mamba_lazy_prealloc_at_boundary.
|
"""Lazy-mode spec counterpart of mamba_lazy_prealloc_at_boundary.
|
||||||
@@ -3134,16 +3148,16 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
pool = self.req_to_token_pool
|
pool = self.req_to_token_pool
|
||||||
track_positions: List[int] = []
|
track_positions: List[int] = []
|
||||||
for req in self.reqs:
|
for req in self.reqs:
|
||||||
buf = req.mamba_ping_pong_track_buffer
|
buf = req.kv.mamba_ping_pong_track_buffer
|
||||||
assert buf is not None
|
assert buf is not None
|
||||||
if not mamba_lazy_spec_in_window(
|
if not mamba_lazy_spec_in_window(
|
||||||
req, mamba_track_interval, max_draft_tokens
|
req, mamba_track_interval, max_draft_tokens
|
||||||
):
|
):
|
||||||
# No crossing reachable: the scatter mask stays -1, the
|
# No crossing reachable: the scatter mask stays -1, the
|
||||||
# position is never written.
|
# position is never written.
|
||||||
track_positions.append(req.mamba_next_track_idx)
|
track_positions.append(req.kv.mamba_next_track_idx)
|
||||||
continue
|
continue
|
||||||
other_idx = 1 - req.mamba_next_track_idx
|
other_idx = 1 - req.kv.mamba_next_track_idx
|
||||||
has_pending = buf[other_idx].item() != -1
|
has_pending = buf[other_idx].item() != -1
|
||||||
if not has_pending:
|
if not has_pending:
|
||||||
if envs.SGLANG_TEST_MAMBA_LAZY_ALLOC_FAIL.get():
|
if envs.SGLANG_TEST_MAMBA_LAZY_ALLOC_FAIL.get():
|
||||||
@@ -3157,7 +3171,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
has_pending = True
|
has_pending = True
|
||||||
# On failure the verify scatters in place into the keep slot.
|
# On failure the verify scatters in place into the keep slot.
|
||||||
track_positions.append(
|
track_positions.append(
|
||||||
other_idx if has_pending else req.mamba_next_track_idx
|
other_idx if has_pending else req.kv.mamba_next_track_idx
|
||||||
)
|
)
|
||||||
self.mamba_lazy_spec_track_positions_cpu = track_positions
|
self.mamba_lazy_spec_track_positions_cpu = track_positions
|
||||||
|
|
||||||
@@ -3497,7 +3511,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# seqlen progress is monotonic per KV handle.
|
# seqlen progress is monotonic per KV handle.
|
||||||
if (
|
if (
|
||||||
req.decode_batch_idx >= 1
|
req.decode_batch_idx >= 1
|
||||||
and req.kv.is_held
|
and req.kv.holds_kv
|
||||||
and req.seqlen - 1 - sliding_window_size
|
and req.seqlen - 1 - sliding_window_size
|
||||||
>= req.kv.swa_evicted_seqlen + eviction_interval
|
>= req.kv.swa_evicted_seqlen + eviction_interval
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -833,7 +833,7 @@ class PrefillAdder:
|
|||||||
backstopped by the fail-loud RuntimeError in `alloc_req_slots`. FIXME: if
|
backstopped by the fail-loud RuntimeError in `alloc_req_slots`. FIXME: if
|
||||||
over-admission crashes under pressure, make this more conservative (e.g.
|
over-admission crashes under pressure, make this more conservative (e.g.
|
||||||
multiply by `MAMBA_STATE_PER_REQ_PREFIX_CACHE`)."""
|
multiply by `MAMBA_STATE_PER_REQ_PREFIX_CACHE`)."""
|
||||||
if self._mamba_slot_cost and req.mamba_pool_idx is None:
|
if self._mamba_slot_cost and not req.kv.holds_mamba:
|
||||||
return self._mamba_slot_cost
|
return self._mamba_slot_cost
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
|
|||||||
@@ -3587,15 +3587,13 @@ class Scheduler(
|
|||||||
if not added:
|
if not added:
|
||||||
# init_next_round_input() may stage deferred Mamba COW/clear
|
# init_next_round_input() may stage deferred Mamba COW/clear
|
||||||
# metadata before add_one_req() rejects the request.
|
# metadata before add_one_req() rejects the request.
|
||||||
req.mamba_cow_src_index = None
|
req.kv.mamba_cow_src_index = None
|
||||||
req.mamba_needs_clear = False
|
req.kv.mamba_needs_clear = False
|
||||||
if req.mamba_pool_idx is not None and not getattr(
|
if req.kv.holds_mamba and not getattr(req, "session", None):
|
||||||
req, "session", None
|
|
||||||
):
|
|
||||||
self.tree_cache.req_to_token_pool.mamba_allocator.free(
|
self.tree_cache.req_to_token_pool.mamba_allocator.free(
|
||||||
req.mamba_pool_idx.unsqueeze(-1)
|
req.kv.mamba_pool_idx.unsqueeze(-1)
|
||||||
)
|
)
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
break
|
break
|
||||||
|
|
||||||
if mamba_allocator is not None:
|
if mamba_allocator is not None:
|
||||||
@@ -4852,7 +4850,7 @@ class Scheduler(
|
|||||||
|
|
||||||
# For mamba radix cache
|
# For mamba radix cache
|
||||||
if (
|
if (
|
||||||
req.mamba_pool_idx is not None
|
req.kv.holds_mamba
|
||||||
and self.disaggregation_mode != DisaggregationMode.DECODE
|
and self.disaggregation_mode != DisaggregationMode.DECODE
|
||||||
):
|
):
|
||||||
release_kv_cache(req, self.tree_cache, is_insert=False)
|
release_kv_cache(req, self.tree_cache, is_insert=False)
|
||||||
@@ -4867,10 +4865,7 @@ class Scheduler(
|
|||||||
self.ipc_channels.send_to_tokenizer.send_output(
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
_make_abort_req(req), req
|
_make_abort_req(req), req
|
||||||
)
|
)
|
||||||
if (
|
if req.kv.holds_kv or req.kv.holds_mamba:
|
||||||
req.kv.req_pool_idx is not None
|
|
||||||
or getattr(req, "mamba_pool_idx", None) is not None
|
|
||||||
):
|
|
||||||
release_kv_cache(req, self.tree_cache, is_insert=False)
|
release_kv_cache(req, self.tree_cache, is_insert=False)
|
||||||
logger.debug(f"Abort dLLM queued request. {req.rid=}")
|
logger.debug(f"Abort dLLM queued request. {req.rid=}")
|
||||||
|
|
||||||
|
|||||||
@@ -1125,14 +1125,14 @@ class SchedulerBatchResultProcessor:
|
|||||||
known_mamba_boundary = bool(batch.mamba_track_mask_next_cpu[i])
|
known_mamba_boundary = bool(batch.mamba_track_mask_next_cpu[i])
|
||||||
|
|
||||||
if completed_mamba_boundary and not lazy:
|
if completed_mamba_boundary and not lazy:
|
||||||
req.mamba_last_track_idx = batch.mamba_track_buffer_indices[i]
|
req.kv.mamba_last_track_idx = batch.mamba_track_buffer_indices[i]
|
||||||
req.mamba_last_track_seqlen = req.kv.kv_committed_len - lookahead
|
req.kv.mamba_last_track_seqlen = req.kv.kv_committed_len - lookahead
|
||||||
elif (
|
elif (
|
||||||
req.finished()
|
req.finished()
|
||||||
and lazy
|
and lazy
|
||||||
and lookahead == 1
|
and lookahead == 1
|
||||||
and known_mamba_boundary
|
and known_mamba_boundary
|
||||||
and req.mamba_next_track_idx == req.mamba_last_track_idx
|
and req.kv.mamba_next_track_idx == req.kv.mamba_last_track_idx
|
||||||
):
|
):
|
||||||
req.mamba_lazy_is_insert = False
|
req.mamba_lazy_is_insert = False
|
||||||
|
|
||||||
@@ -1216,7 +1216,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
Lazy: keep the same index (prealloc handles the swap) and run
|
Lazy: keep the same index (prealloc handles the swap) and run
|
||||||
post-decode cleanup to free the temporary second slot.
|
post-decode cleanup to free the temporary second slot.
|
||||||
"""
|
"""
|
||||||
if req.mamba_ping_pong_track_buffer is None:
|
if req.kv.mamba_ping_pong_track_buffer is None:
|
||||||
return
|
return
|
||||||
|
|
||||||
lazy = mamba_extra_buffer_lazy_enabled()
|
lazy = mamba_extra_buffer_lazy_enabled()
|
||||||
@@ -1238,17 +1238,17 @@ class SchedulerBatchResultProcessor:
|
|||||||
if not at_boundary:
|
if not at_boundary:
|
||||||
return
|
return
|
||||||
|
|
||||||
track_idx = req.mamba_next_track_idx
|
track_idx = req.kv.mamba_next_track_idx
|
||||||
if not known_boundary and batch.mamba_track_buffer_indices is not None:
|
if not known_boundary and batch.mamba_track_buffer_indices is not None:
|
||||||
track_idx = batch.mamba_track_buffer_indices[i]
|
track_idx = batch.mamba_track_buffer_indices[i]
|
||||||
if not known_boundary:
|
if not known_boundary:
|
||||||
req.mamba_last_track_seqlen = track_seqlen
|
req.kv.mamba_last_track_seqlen = track_seqlen
|
||||||
if lazy:
|
if lazy:
|
||||||
self.mamba_lazy_post_decode_at_boundary(req, batch, track_idx)
|
self.mamba_lazy_post_decode_at_boundary(req, batch, track_idx)
|
||||||
else:
|
else:
|
||||||
if not known_boundary:
|
if not known_boundary:
|
||||||
req.mamba_last_track_idx = track_idx
|
req.kv.mamba_last_track_idx = track_idx
|
||||||
req.mamba_next_track_idx = (
|
req.kv.mamba_next_track_idx = (
|
||||||
batch.req_to_token_pool.get_mamba_ping_pong_other_idx(track_idx)
|
batch.req_to_token_pool.get_mamba_ping_pong_other_idx(track_idx)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1274,12 +1274,12 @@ class SchedulerBatchResultProcessor:
|
|||||||
if req.finished():
|
if req.finished():
|
||||||
# Skip the donation if a scatter wrote or may still write the keep slot.
|
# Skip the donation if a scatter wrote or may still write the keep slot.
|
||||||
keep_written_by_this_step = (
|
keep_written_by_this_step = (
|
||||||
crossed and planned_pos == req.mamba_next_track_idx
|
crossed and planned_pos == req.kv.mamba_next_track_idx
|
||||||
)
|
)
|
||||||
other_idx = 1 - req.mamba_next_track_idx
|
other_idx = 1 - req.kv.mamba_next_track_idx
|
||||||
# Recompute the in-flight verify's plan (kv_committed_len is
|
# Recompute the in-flight verify's plan (kv_committed_len is
|
||||||
# frozen since its prepare, so the recompute is exact).
|
# frozen since its prepare, so the recompute is exact).
|
||||||
keep_may_be_written_in_flight = req.mamba_ping_pong_track_buffer[
|
keep_may_be_written_in_flight = req.kv.mamba_ping_pong_track_buffer[
|
||||||
other_idx
|
other_idx
|
||||||
].item() == -1 and mamba_lazy_spec_in_window(
|
].item() == -1 and mamba_lazy_spec_in_window(
|
||||||
req,
|
req,
|
||||||
@@ -1296,18 +1296,18 @@ class SchedulerBatchResultProcessor:
|
|||||||
|
|
||||||
if not crossed or planned_pos is None:
|
if not crossed or planned_pos is None:
|
||||||
return
|
return
|
||||||
if planned_pos != req.mamba_next_track_idx:
|
if planned_pos != req.kv.mamba_next_track_idx:
|
||||||
# Promote pending -> keep: free the old checkpoint, repoint.
|
# Promote pending -> keep: free the old checkpoint, repoint.
|
||||||
pool = batch.req_to_token_pool
|
pool = batch.req_to_token_pool
|
||||||
keep_idx = req.mamba_next_track_idx
|
keep_idx = req.kv.mamba_next_track_idx
|
||||||
keep_val = req.mamba_ping_pong_track_buffer[keep_idx]
|
keep_val = req.kv.mamba_ping_pong_track_buffer[keep_idx]
|
||||||
pool.mamba_allocator.free(keep_val.unsqueeze(0))
|
pool.mamba_allocator.free(keep_val.unsqueeze(0))
|
||||||
pool.set_mamba_ping_pong_slot(req, keep_idx, -1)
|
pool.set_mamba_ping_pong_slot(req, keep_idx, -1)
|
||||||
req.mamba_next_track_idx = planned_pos
|
req.kv.mamba_next_track_idx = planned_pos
|
||||||
# else: in-place fallback, or promoted by an earlier confirmation —
|
# else: in-place fallback, or promoted by an earlier confirmation —
|
||||||
# keep holds the track_seqlen state either way.
|
# keep holds the track_seqlen state either way.
|
||||||
req.mamba_last_track_idx = planned_pos
|
req.kv.mamba_last_track_idx = planned_pos
|
||||||
req.mamba_last_track_seqlen = track_seqlen
|
req.kv.mamba_last_track_seqlen = track_seqlen
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _mamba_assert_committed_len_lookahead(req: Req) -> None:
|
def _mamba_assert_committed_len_lookahead(req: Req) -> None:
|
||||||
@@ -1355,13 +1355,13 @@ class SchedulerBatchResultProcessor:
|
|||||||
self, req: Req, batch: ScheduleBatch, track_idx: int
|
self, req: Req, batch: ScheduleBatch, track_idx: int
|
||||||
):
|
):
|
||||||
"""Commit a completed lazy-mode boundary and free its old slot."""
|
"""Commit a completed lazy-mode boundary and free its old slot."""
|
||||||
req.mamba_last_track_idx = track_idx
|
req.kv.mamba_last_track_idx = track_idx
|
||||||
req.mamba_next_track_idx = track_idx
|
req.kv.mamba_next_track_idx = track_idx
|
||||||
other_idx = 1 - track_idx
|
other_idx = 1 - track_idx
|
||||||
other_val = req.mamba_ping_pong_track_buffer[other_idx].item()
|
other_val = req.kv.mamba_ping_pong_track_buffer[other_idx].item()
|
||||||
if other_val != -1:
|
if other_val != -1:
|
||||||
pool = batch.req_to_token_pool
|
pool = batch.req_to_token_pool
|
||||||
pool.mamba_allocator.free(
|
pool.mamba_allocator.free(
|
||||||
req.mamba_ping_pong_track_buffer[other_idx].unsqueeze(0)
|
req.kv.mamba_ping_pong_track_buffer[other_idx].unsqueeze(0)
|
||||||
)
|
)
|
||||||
pool.set_mamba_ping_pong_slot(req, other_idx, -1)
|
pool.set_mamba_ping_pong_slot(req, other_idx, -1)
|
||||||
|
|||||||
@@ -252,7 +252,7 @@ class SchedulerInvariantChecker:
|
|||||||
swa_uncached = 0
|
swa_uncached = 0
|
||||||
for batch in batches:
|
for batch in batches:
|
||||||
for req in batch.reqs:
|
for req in batch.reqs:
|
||||||
if not req.kv.is_held:
|
if not req.kv.holds_kv:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
allocated_len = req.kv.kv_allocated_len
|
allocated_len = req.kv.kv_allocated_len
|
||||||
@@ -324,7 +324,7 @@ class SchedulerInvariantChecker:
|
|||||||
batch = self.get_last_batch()
|
batch = self.get_last_batch()
|
||||||
if batch is not None:
|
if batch is not None:
|
||||||
for req in batch.reqs:
|
for req in batch.reqs:
|
||||||
if not req.kv.is_held:
|
if not req.kv.holds_kv:
|
||||||
continue
|
continue
|
||||||
_add_owner(
|
_add_owner(
|
||||||
req,
|
req,
|
||||||
@@ -336,7 +336,7 @@ class SchedulerInvariantChecker:
|
|||||||
sess = getattr(self.tree_cache, "slots", None)
|
sess = getattr(self.tree_cache, "slots", None)
|
||||||
if sess:
|
if sess:
|
||||||
for sid, slot in sess.items():
|
for sid, slot in sess.items():
|
||||||
if slot.kv.is_held:
|
if slot.kv.holds_kv:
|
||||||
_add_owner(
|
_add_owner(
|
||||||
slot,
|
slot,
|
||||||
f"slot {sid[:8]}",
|
f"slot {sid[:8]}",
|
||||||
|
|||||||
@@ -723,15 +723,15 @@ class SchedulerPPMixin:
|
|||||||
latencies.append(latency_ms)
|
latencies.append(latency_ms)
|
||||||
|
|
||||||
# Release KV and Mamba cache
|
# Release KV and Mamba cache
|
||||||
if req.kv.is_held:
|
if req.kv.holds_kv:
|
||||||
kv_indices = self.req_to_token_pool.req_to_token[
|
kv_indices = self.req_to_token_pool.req_to_token[
|
||||||
req.kv.req_pool_idx, : req.extend_range.end
|
req.kv.req_pool_idx, : req.extend_range.end
|
||||||
]
|
]
|
||||||
self.token_to_kv_pool_allocator.free(kv_indices)
|
self.token_to_kv_pool_allocator.free(kv_indices)
|
||||||
if req.mamba_pool_idx is not None:
|
if req.kv.holds_mamba:
|
||||||
self.req_to_token_pool.free_mamba_cache(req)
|
self.req_to_token_pool.free_mamba_cache(req)
|
||||||
self.req_to_token_pool.free(req)
|
self.req_to_token_pool.free(req)
|
||||||
req.kv.mark_released()
|
req.kv.mark_kv_released()
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"[PP Dynamic Chunk] [PP0] Profiled {len(seq_lens)} samples: "
|
f"[PP Dynamic Chunk] [PP0] Profiled {len(seq_lens)} samples: "
|
||||||
|
|||||||
@@ -62,7 +62,7 @@ def free_swa_out_of_window_slots(
|
|||||||
is_chunk_cache: bool = False,
|
is_chunk_cache: bool = False,
|
||||||
retain_floor: int | None = None,
|
retain_floor: int | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
if not req.kv.is_held:
|
if not req.kv.holds_kv:
|
||||||
return
|
return
|
||||||
|
|
||||||
# For swa radix cache, we need to evict the tokens that are not in the tree cache and also not in the sliding window
|
# For swa radix cache, we need to evict the tokens that are not in the tree cache and also not in the sliding window
|
||||||
@@ -159,8 +159,8 @@ def retraction_backup(
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
unified_cache = cast("UnifiedRadixCache", tree_cache)
|
unified_cache = cast("UnifiedRadixCache", tree_cache)
|
||||||
req.retraction_backup = unified_cache.retraction_backup(req)
|
req.kv.retraction_backup = unified_cache.retraction_backup(req)
|
||||||
return req.retraction_backup is not None
|
return req.kv.retraction_backup is not None
|
||||||
|
|
||||||
|
|
||||||
def retraction_restore(
|
def retraction_restore(
|
||||||
@@ -179,38 +179,38 @@ def retraction_restore(
|
|||||||
return
|
return
|
||||||
|
|
||||||
unified_cache = cast("UnifiedRadixCache", tree_cache)
|
unified_cache = cast("UnifiedRadixCache", tree_cache)
|
||||||
assert req.retraction_backup is not None
|
assert req.kv.retraction_backup is not None
|
||||||
unified_cache.retraction_restore(req, req.retraction_backup)
|
unified_cache.retraction_restore(req, req.kv.retraction_backup)
|
||||||
req.retraction_backup = None
|
req.kv.retraction_backup = None
|
||||||
|
|
||||||
|
|
||||||
def retraction_discard(req: Req, tree_cache: BasePrefixCache, backend: str) -> None:
|
def retraction_discard(req: Req, tree_cache: BasePrefixCache, backend: str) -> None:
|
||||||
if backend == "cpu_tensor":
|
if backend == "cpu_tensor":
|
||||||
req.retraction_backup = None
|
req.kv.retraction_backup = None
|
||||||
return
|
return
|
||||||
if backend != "host_pool":
|
if backend != "host_pool":
|
||||||
raise ValueError(f"Unknown retraction backup backend: {backend}")
|
raise ValueError(f"Unknown retraction backup backend: {backend}")
|
||||||
if req.retraction_backup is None:
|
if req.kv.retraction_backup is None:
|
||||||
return
|
return
|
||||||
|
|
||||||
unified_cache = cast("UnifiedRadixCache", tree_cache)
|
unified_cache = cast("UnifiedRadixCache", tree_cache)
|
||||||
unified_cache.retraction_discard(req.retraction_backup)
|
unified_cache.retraction_discard(req.kv.retraction_backup)
|
||||||
req.retraction_backup = None
|
req.kv.retraction_backup = None
|
||||||
|
|
||||||
|
|
||||||
def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = True):
|
def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = True):
|
||||||
assert (not req.kv.is_held) == req.kv.is_released
|
assert (not req.kv.holds_kv) == req.kv.is_kv_released
|
||||||
# MambaRadixCache may alloc mamba state before alloc KV cache
|
# MambaRadixCache may alloc mamba state before alloc KV cache
|
||||||
if not req.kv.is_held:
|
if not req.kv.holds_kv:
|
||||||
assert (
|
assert (
|
||||||
tree_cache.supports_mamba()
|
tree_cache.supports_mamba()
|
||||||
), "Only MambaRadixCache allow freeing before alloc"
|
), "Only MambaRadixCache allow freeing before alloc"
|
||||||
# TODO (csy, hanming): clean up this early allocation logic
|
# TODO (csy, hanming): clean up this early allocation logic
|
||||||
if req.mamba_pool_idx is not None:
|
if req.kv.holds_mamba:
|
||||||
tree_cache.req_to_token_pool.mamba_allocator.free(
|
tree_cache.req_to_token_pool.mamba_allocator.free(
|
||||||
req.mamba_pool_idx.unsqueeze(-1)
|
req.kv.mamba_pool_idx.unsqueeze(-1)
|
||||||
)
|
)
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
return
|
return
|
||||||
|
|
||||||
effective_kv_committed_len = req.effective_kv_committed_len()
|
effective_kv_committed_len = req.effective_kv_committed_len()
|
||||||
@@ -222,8 +222,8 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
|
|||||||
|
|
||||||
# StreamingSession.cache_finished_req handles speculative tail trim
|
# StreamingSession.cache_finished_req handles speculative tail trim
|
||||||
# internally, then sets req_pool_idx = None.
|
# internally, then sets req_pool_idx = None.
|
||||||
assert (not req.kv.is_held) == req.kv.is_released
|
assert (not req.kv.holds_kv) == req.kv.is_kv_released
|
||||||
if not req.kv.is_held:
|
if not req.kv.holds_kv:
|
||||||
return
|
return
|
||||||
|
|
||||||
start_p, end_p = effective_kv_committed_len, req.kv.kv_allocated_len
|
start_p, end_p = effective_kv_committed_len, req.kv.kv_allocated_len
|
||||||
@@ -234,13 +234,13 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
|
|||||||
not tree_cache.supports_mamba()
|
not tree_cache.supports_mamba()
|
||||||
):
|
):
|
||||||
assert (
|
assert (
|
||||||
req.mamba_pool_idx is not None
|
req.kv.holds_mamba
|
||||||
), "mamba state is freed while the tree cache does not manage mamba states"
|
), "mamba state is freed while the tree cache does not manage mamba states"
|
||||||
tree_cache.req_to_token_pool.free_mamba_cache(req)
|
tree_cache.req_to_token_pool.free_mamba_cache(req)
|
||||||
# The DSV4-NPU ReqToTokenPool subclass's free() additionally releases the
|
# The DSV4-NPU ReqToTokenPool subclass's free() additionally releases the
|
||||||
# c4/c128 state pages; other ReqToTokenPool subclasses are a no-op here.
|
# c4/c128 state pages; other ReqToTokenPool subclasses are a no-op here.
|
||||||
tree_cache.req_to_token_pool.free(req)
|
tree_cache.req_to_token_pool.free(req)
|
||||||
req.kv.mark_released()
|
req.kv.mark_kv_released()
|
||||||
|
|
||||||
|
|
||||||
def _release_overallocated_kv_indices(
|
def _release_overallocated_kv_indices(
|
||||||
|
|||||||
@@ -561,7 +561,7 @@ class MambaRadixCache(BasePrefixCache):
|
|||||||
|
|
||||||
if is_insert:
|
if is_insert:
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.enable_mamba_extra_buffer:
|
||||||
cache_len = req.mamba_last_track_seqlen
|
cache_len = req.kv.mamba_last_track_seqlen
|
||||||
else:
|
else:
|
||||||
cache_len = len(token_ids)
|
cache_len = len(token_ids)
|
||||||
# ReplaySSM (no_buffer): `temporal[slot]` lags the live state by
|
# ReplaySSM (no_buffer): `temporal[slot]` lags the live state by
|
||||||
@@ -571,8 +571,8 @@ class MambaRadixCache(BasePrefixCache):
|
|||||||
# with its key length. page_size is asserted == 1, so no realign.
|
# with its key length. page_size is asserted == 1, so no realign.
|
||||||
write_pos_buf = self.req_to_token_pool.mamba_pool.replayssm_write_pos
|
write_pos_buf = self.req_to_token_pool.mamba_pool.replayssm_write_pos
|
||||||
if write_pos_buf is not None:
|
if write_pos_buf is not None:
|
||||||
cache_len -= int(write_pos_buf[req.mamba_pool_idx].item())
|
cache_len -= int(write_pos_buf[req.kv.mamba_pool_idx].item())
|
||||||
write_pos_buf[req.mamba_pool_idx] = 0
|
write_pos_buf[req.kv.mamba_pool_idx] = 0
|
||||||
if cache_len is None:
|
if cache_len is None:
|
||||||
cache_len = 0
|
cache_len = 0
|
||||||
if cache_len != len(token_ids):
|
if cache_len != len(token_ids):
|
||||||
@@ -602,16 +602,16 @@ class MambaRadixCache(BasePrefixCache):
|
|||||||
mamba_ping_pong_track_buffer_to_keep = (
|
mamba_ping_pong_track_buffer_to_keep = (
|
||||||
self.req_to_token_pool.get_mamba_ping_pong_keep_idx(req)
|
self.req_to_token_pool.get_mamba_ping_pong_keep_idx(req)
|
||||||
)
|
)
|
||||||
src_active = req.mamba_ping_pong_track_buffer[
|
src_active = req.kv.mamba_ping_pong_track_buffer[
|
||||||
mamba_ping_pong_track_buffer_to_keep
|
mamba_ping_pong_track_buffer_to_keep
|
||||||
].unsqueeze(-1)
|
].unsqueeze(-1)
|
||||||
if _MAMBA_DEBUG_ASSERTS:
|
if _MAMBA_DEBUG_ASSERTS:
|
||||||
# .item() forces a cudaStreamSynchronize; only pay it when debugging.
|
# .item() forces a cudaStreamSynchronize; only pay it when debugging.
|
||||||
assert src_active.item() != -1, (
|
assert src_active.item() != -1, (
|
||||||
f"Cached mamba slot is -1: keep_idx={mamba_ping_pong_track_buffer_to_keep}, "
|
f"Cached mamba slot is -1: keep_idx={mamba_ping_pong_track_buffer_to_keep}, "
|
||||||
f"buf={req.mamba_ping_pong_track_buffer.tolist()}, "
|
f"buf={req.kv.mamba_ping_pong_track_buffer.tolist()}, "
|
||||||
f"next_track_idx={req.mamba_next_track_idx}, "
|
f"next_track_idx={req.kv.mamba_next_track_idx}, "
|
||||||
f"last_track_seqlen={req.mamba_last_track_seqlen}, "
|
f"last_track_seqlen={req.kv.mamba_last_track_seqlen}, "
|
||||||
f"rid={req.rid}"
|
f"rid={req.rid}"
|
||||||
)
|
)
|
||||||
if self.int8_ckpt_pool is not None:
|
if self.int8_ckpt_pool is not None:
|
||||||
@@ -623,10 +623,10 @@ class MambaRadixCache(BasePrefixCache):
|
|||||||
else:
|
else:
|
||||||
if self.int8_ckpt_pool is not None:
|
if self.int8_ckpt_pool is not None:
|
||||||
mamba_value = self._commit_int8_checkpoint(
|
mamba_value = self._commit_int8_checkpoint(
|
||||||
req.mamba_pool_idx.unsqueeze(-1)
|
req.kv.mamba_pool_idx.unsqueeze(-1)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
mamba_value = req.mamba_pool_idx.unsqueeze(-1).clone()
|
mamba_value = req.kv.mamba_pool_idx.unsqueeze(-1).clone()
|
||||||
mamba_ping_pong_track_buffer_to_keep = None
|
mamba_ping_pong_track_buffer_to_keep = None
|
||||||
|
|
||||||
result = self.insert(
|
result = self.insert(
|
||||||
@@ -685,7 +685,7 @@ class MambaRadixCache(BasePrefixCache):
|
|||||||
|
|
||||||
token_ids = req.get_fill_ids()
|
token_ids = req.get_fill_ids()
|
||||||
cache_len = (
|
cache_len = (
|
||||||
req.mamba_last_track_seqlen
|
req.kv.mamba_last_track_seqlen
|
||||||
if self.enable_mamba_extra_buffer
|
if self.enable_mamba_extra_buffer
|
||||||
else len(token_ids)
|
else len(token_ids)
|
||||||
)
|
)
|
||||||
@@ -726,7 +726,7 @@ class MambaRadixCache(BasePrefixCache):
|
|||||||
self.req_to_token_pool.mamba_allocator.free(src_active)
|
self.req_to_token_pool.mamba_allocator.free(src_active)
|
||||||
else:
|
else:
|
||||||
mamba_value_donated = self._commit_int8_checkpoint(
|
mamba_value_donated = self._commit_int8_checkpoint(
|
||||||
req.mamba_pool_idx.view(-1)
|
req.kv.mamba_pool_idx.view(-1)
|
||||||
)
|
)
|
||||||
elif self.enable_mamba_extra_buffer:
|
elif self.enable_mamba_extra_buffer:
|
||||||
new_slot = self._alloc_mamba_slot()
|
new_slot = self._alloc_mamba_slot()
|
||||||
@@ -739,7 +739,7 @@ class MambaRadixCache(BasePrefixCache):
|
|||||||
# virtual->physical (identity for the non-unified memory pool) before the copy.
|
# virtual->physical (identity for the non-unified memory pool) before the copy.
|
||||||
translate = self.req_to_token_pool.translate_mamba_indices
|
translate = self.req_to_token_pool.translate_mamba_indices
|
||||||
self.req_to_token_pool.mamba_pool.copy_from(
|
self.req_to_token_pool.mamba_pool.copy_from(
|
||||||
translate(req.mamba_pool_idx.unsqueeze(0)),
|
translate(req.kv.mamba_pool_idx.unsqueeze(0)),
|
||||||
translate(mamba_value_donated),
|
translate(mamba_value_donated),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -799,7 +799,7 @@ class MambaRadixCache(BasePrefixCache):
|
|||||||
[new_indices, kv_indices_orig[len(new_indices) :]]
|
[new_indices, kv_indices_orig[len(new_indices) :]]
|
||||||
)
|
)
|
||||||
req.kv.cache_protected_len = len(new_indices)
|
req.kv.cache_protected_len = len(new_indices)
|
||||||
req.mamba_last_track_seqlen = None
|
req.kv.mamba_last_track_seqlen = None
|
||||||
req.last_node = new_last_node
|
req.last_node = new_last_node
|
||||||
|
|
||||||
def pretty_print(self) -> None:
|
def pretty_print(self) -> None:
|
||||||
@@ -1179,7 +1179,7 @@ class MambaRadixCache(BasePrefixCache):
|
|||||||
|
|
||||||
# Defer COW to forward stream: record source index, allocate destination
|
# Defer COW to forward stream: record source index, allocate destination
|
||||||
if cow_mamba and last_node.mamba_value is not None:
|
if cow_mamba and last_node.mamba_value is not None:
|
||||||
if req.mamba_pool_idx is None:
|
if not req.kv.holds_mamba:
|
||||||
dst_index = self.req_to_token_pool.mamba_allocator.alloc(1)
|
dst_index = self.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
if dst_index is None:
|
if dst_index is None:
|
||||||
self.inc_lock_ref(last_node)
|
self.inc_lock_ref(last_node)
|
||||||
@@ -1187,9 +1187,9 @@ class MambaRadixCache(BasePrefixCache):
|
|||||||
dst_index = self.req_to_token_pool.mamba_allocator.alloc(1)
|
dst_index = self.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
self.dec_lock_ref(last_node)
|
self.dec_lock_ref(last_node)
|
||||||
assert dst_index is not None, "Can not alloc mamba cache"
|
assert dst_index is not None, "Can not alloc mamba cache"
|
||||||
req.mamba_pool_idx = dst_index[0]
|
req.kv.mamba_pool_idx = dst_index[0]
|
||||||
req.mamba_cow_src_index = last_node.mamba_value
|
req.kv.mamba_cow_src_index = last_node.mamba_value
|
||||||
req.mamba_needs_clear = False
|
req.kv.mamba_needs_clear = False
|
||||||
|
|
||||||
value = value[:best_value_len]
|
value = value[:best_value_len]
|
||||||
if value:
|
if value:
|
||||||
|
|||||||
@@ -1359,32 +1359,34 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
mamba_indices: list[torch.Tensor] = []
|
mamba_indices: list[torch.Tensor] = []
|
||||||
mamba_ping_pong_track_buffers: list[torch.Tensor] = []
|
mamba_ping_pong_track_buffers: list[torch.Tensor] = []
|
||||||
for req in reqs:
|
for req in reqs:
|
||||||
if req.mamba_pool_idx is not None: # for radix cache / continuing chunked
|
if req.kv.holds_mamba: # for radix cache / continuing chunked
|
||||||
pass
|
pass
|
||||||
else:
|
else:
|
||||||
mid = self.mamba_allocator.alloc(1)
|
mid = self.mamba_allocator.alloc(1)
|
||||||
assert (
|
assert (
|
||||||
mid is not None
|
mid is not None
|
||||||
), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size. {mid=}, {self.mamba_pool.size=}, {self.mamba_allocator.available_size()=}, {len(reqs)=}"
|
), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size. {mid=}, {self.mamba_pool.size=}, {self.mamba_allocator.available_size()=}, {len(reqs)=}"
|
||||||
req.mamba_pool_idx = mid[0]
|
req.kv.mamba_pool_idx = mid[0]
|
||||||
req.mamba_needs_clear = True
|
req.kv.mamba_needs_clear = True
|
||||||
# GDN ReplaySSM: a freshly (re)assigned slot starts an empty
|
# GDN ReplaySSM: a freshly (re)assigned slot starts an empty
|
||||||
# ring. write_pos=0 means "ring empty", so the decode kernel
|
# ring. write_pos=0 means "ring empty", so the decode kernel
|
||||||
# ignores ring contents and reads only the checkpoint state
|
# ignores ring contents and reads only the checkpoint state
|
||||||
# (the post-prefill state that prefill wrote into this slot).
|
# (the post-prefill state that prefill wrote into this slot).
|
||||||
if self.mamba_pool.replayssm_write_pos is not None:
|
if self.mamba_pool.replayssm_write_pos is not None:
|
||||||
self.mamba_pool.replayssm_write_pos[req.mamba_pool_idx] = 0
|
self.mamba_pool.replayssm_write_pos[req.kv.mamba_pool_idx] = 0
|
||||||
# ReplaySSM spec-verify ring: an empty ring also resets the
|
# ReplaySSM spec-verify ring: an empty ring also resets the
|
||||||
# circular origin + flush flag so the first verify step on this
|
# circular origin + flush flag so the first verify step on this
|
||||||
# freshly-prefilled slot reconstructs from the checkpoint alone.
|
# freshly-prefilled slot reconstructs from the checkpoint alone.
|
||||||
if self.mamba_pool.replayssm_cache_base is not None:
|
if self.mamba_pool.replayssm_cache_base is not None:
|
||||||
self.mamba_pool.replayssm_cache_base[req.mamba_pool_idx] = 0
|
self.mamba_pool.replayssm_cache_base[req.kv.mamba_pool_idx] = 0
|
||||||
self.mamba_pool.replayssm_is_flush[req.mamba_pool_idx] = 0
|
self.mamba_pool.replayssm_is_flush[req.kv.mamba_pool_idx] = 0
|
||||||
mamba_indices.append(req.mamba_pool_idx)
|
mamba_indices.append(req.kv.mamba_pool_idx)
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.enable_mamba_extra_buffer:
|
||||||
if req.mamba_ping_pong_track_buffer is None:
|
if req.kv.mamba_ping_pong_track_buffer is None:
|
||||||
self._alloc_ping_pong_buffer(req)
|
self._alloc_ping_pong_buffer(req)
|
||||||
mamba_ping_pong_track_buffers.append(req.mamba_ping_pong_track_buffer)
|
mamba_ping_pong_track_buffers.append(
|
||||||
|
req.kv.mamba_ping_pong_track_buffer
|
||||||
|
)
|
||||||
assert len(select_index) == len(
|
assert len(select_index) == len(
|
||||||
mamba_indices
|
mamba_indices
|
||||||
), "Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size."
|
), "Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size."
|
||||||
@@ -1449,7 +1451,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
|
|
||||||
def get_mamba_ping_pong_keep_idx(self, req: Req) -> int:
|
def get_mamba_ping_pong_keep_idx(self, req: Req) -> int:
|
||||||
"""Return the ping-pong index holding the most recent tracked state."""
|
"""Return the ping-pong index holding the most recent tracked state."""
|
||||||
return req.mamba_last_track_idx
|
return req.kv.mamba_last_track_idx
|
||||||
|
|
||||||
def _alloc_ping_pong_buffer(self, req: Req):
|
def _alloc_ping_pong_buffer(self, req: Req):
|
||||||
"""Allocate the ping-pong track buffer for a new request.
|
"""Allocate the ping-pong track buffer for a new request.
|
||||||
@@ -1474,9 +1476,9 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
device=slots.device,
|
device=slots.device,
|
||||||
)
|
)
|
||||||
buf[:n] = slots
|
buf[:n] = slots
|
||||||
req.mamba_ping_pong_track_buffer = buf
|
req.kv.mamba_ping_pong_track_buffer = buf
|
||||||
req.mamba_next_track_idx = 0
|
req.kv.mamba_next_track_idx = 0
|
||||||
req.mamba_last_track_idx = (
|
req.kv.mamba_last_track_idx = (
|
||||||
0
|
0
|
||||||
if self.enable_mamba_extra_buffer_lazy
|
if self.enable_mamba_extra_buffer_lazy
|
||||||
else self.get_mamba_ping_pong_other_idx(0)
|
else self.get_mamba_ping_pong_other_idx(0)
|
||||||
@@ -1489,9 +1491,9 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
req_index_to_mamba_ping_pong_track_buffer_mapping in sync so that
|
req_index_to_mamba_ping_pong_track_buffer_mapping in sync so that
|
||||||
set_mamba_track_indices_from_reqs reads correct slot indices.
|
set_mamba_track_indices_from_reqs reads correct slot indices.
|
||||||
"""
|
"""
|
||||||
req.mamba_ping_pong_track_buffer[idx] = value
|
req.kv.mamba_ping_pong_track_buffer[idx] = value
|
||||||
self.req_index_to_mamba_ping_pong_track_buffer_mapping[req.kv.req_pool_idx] = (
|
self.req_index_to_mamba_ping_pong_track_buffer_mapping[req.kv.req_pool_idx] = (
|
||||||
req.mamba_ping_pong_track_buffer
|
req.kv.mamba_ping_pong_track_buffer
|
||||||
)
|
)
|
||||||
|
|
||||||
def donate_mamba_ping_pong_slot(
|
def donate_mamba_ping_pong_slot(
|
||||||
@@ -1504,14 +1506,14 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
"""
|
"""
|
||||||
donate_idx = self.get_mamba_ping_pong_keep_idx(req)
|
donate_idx = self.get_mamba_ping_pong_keep_idx(req)
|
||||||
mamba_value_donated = (
|
mamba_value_donated = (
|
||||||
req.mamba_ping_pong_track_buffer[donate_idx].unsqueeze(-1).clone()
|
req.kv.mamba_ping_pong_track_buffer[donate_idx].unsqueeze(-1).clone()
|
||||||
)
|
)
|
||||||
if _MAMBA_DEBUG_ASSERTS:
|
if _MAMBA_DEBUG_ASSERTS:
|
||||||
# .item() forces a cudaStreamSynchronize; only pay it when debugging.
|
# .item() forces a cudaStreamSynchronize; only pay it when debugging.
|
||||||
assert mamba_value_donated.item() != -1, (
|
assert mamba_value_donated.item() != -1, (
|
||||||
f"Donated mamba slot is -1: donate_idx={donate_idx}, "
|
f"Donated mamba slot is -1: donate_idx={donate_idx}, "
|
||||||
f"buf={req.mamba_ping_pong_track_buffer.tolist()}, "
|
f"buf={req.kv.mamba_ping_pong_track_buffer.tolist()}, "
|
||||||
f"next_track_idx={req.mamba_next_track_idx}, "
|
f"next_track_idx={req.kv.mamba_next_track_idx}, "
|
||||||
f"rid={req.rid}"
|
f"rid={req.rid}"
|
||||||
)
|
)
|
||||||
self.set_mamba_ping_pong_slot(req, donate_idx, new_slot[0])
|
self.set_mamba_ping_pong_slot(req, donate_idx, new_slot[0])
|
||||||
@@ -1520,10 +1522,10 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
def free_mamba_cache(
|
def free_mamba_cache(
|
||||||
self, req: Req, mamba_ping_pong_track_buffer_to_keep: Optional[int] = None
|
self, req: Req, mamba_ping_pong_track_buffer_to_keep: Optional[int] = None
|
||||||
):
|
):
|
||||||
mamba_index = req.mamba_pool_idx
|
mamba_index = req.kv.mamba_pool_idx
|
||||||
assert mamba_index is not None, "double free? mamba_index is None"
|
assert mamba_index is not None, "double free? mamba_index is None"
|
||||||
self.mamba_allocator.free(mamba_index.unsqueeze(0))
|
self.mamba_allocator.free(mamba_index.unsqueeze(0))
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
|
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.enable_mamba_extra_buffer:
|
||||||
mamba_ping_pong_track_buffer_to_free = (
|
mamba_ping_pong_track_buffer_to_free = (
|
||||||
@@ -1565,13 +1567,13 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.mamba_allocator.free(mamba_ping_pong_track_buffer_to_free)
|
self.mamba_allocator.free(mamba_ping_pong_track_buffer_to_free)
|
||||||
# Match the req.mamba_pool_idx=None clear above so the next
|
# Match the req.kv.mamba_pool_idx=None clear above so the next
|
||||||
# alloc() doesn't see a stale ping-pong reference on the req
|
# alloc() doesn't see a stale ping-pong reference on the req
|
||||||
# and skip allocation (which would silently reuse a freed
|
# and skip allocation (which would silently reuse a freed
|
||||||
# tensor on the req side while the new pool slot leaks).
|
# tensor on the req side while the new pool slot leaks).
|
||||||
req.mamba_ping_pong_track_buffer = None
|
req.kv.mamba_ping_pong_track_buffer = None
|
||||||
req.mamba_next_track_idx = None
|
req.kv.mamba_next_track_idx = None
|
||||||
req.mamba_last_track_idx = None
|
req.kv.mamba_last_track_idx = None
|
||||||
|
|
||||||
def clear(self):
|
def clear(self):
|
||||||
logger.info("Reset HybridReqToTokenPool")
|
logger.info("Reset HybridReqToTokenPool")
|
||||||
|
|||||||
@@ -197,7 +197,7 @@ class MambaComponent(TreeComponent):
|
|||||||
return result
|
return result
|
||||||
req = params.req
|
req = params.req
|
||||||
assert req is not None
|
assert req is not None
|
||||||
if req.mamba_pool_idx is None:
|
if not req.kv.holds_mamba:
|
||||||
dst_index = self.cache.req_to_token_pool.mamba_allocator.alloc(1)
|
dst_index = self.cache.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
if dst_index is None:
|
if dst_index is None:
|
||||||
# Pin the window via inc/dec_lock_ref so evict's SWA release
|
# Pin the window via inc/dec_lock_ref so evict's SWA release
|
||||||
@@ -210,9 +210,9 @@ class MambaComponent(TreeComponent):
|
|||||||
result.best_match_node, lock_result.to_dec_params()
|
result.best_match_node, lock_result.to_dec_params()
|
||||||
)
|
)
|
||||||
assert dst_index is not None, "Can not alloc mamba cache"
|
assert dst_index is not None, "Can not alloc mamba cache"
|
||||||
req.mamba_pool_idx = dst_index[0]
|
req.kv.mamba_pool_idx = dst_index[0]
|
||||||
req.mamba_cow_src_index = src_index
|
req.kv.mamba_cow_src_index = src_index
|
||||||
req.mamba_needs_clear = False
|
req.kv.mamba_needs_clear = False
|
||||||
return result
|
return result
|
||||||
|
|
||||||
def commit_insert_component_data(
|
def commit_insert_component_data(
|
||||||
@@ -534,7 +534,7 @@ class MambaComponent(TreeComponent):
|
|||||||
is_finished: bool,
|
is_finished: bool,
|
||||||
) -> Optional[int]:
|
) -> Optional[int]:
|
||||||
if self.cache.enable_mamba_extra_buffer:
|
if self.cache.enable_mamba_extra_buffer:
|
||||||
cache_len = req.mamba_last_track_seqlen
|
cache_len = req.kv.mamba_last_track_seqlen
|
||||||
else:
|
else:
|
||||||
cache_len = token_ids_len
|
cache_len = token_ids_len
|
||||||
# ReplaySSM (no_buffer): `temporal[slot]` lags the live state by the
|
# ReplaySSM (no_buffer): `temporal[slot]` lags the live state by the
|
||||||
@@ -548,8 +548,8 @@ class MambaComponent(TreeComponent):
|
|||||||
self.cache.req_to_token_pool.mamba_pool.replayssm_write_pos
|
self.cache.req_to_token_pool.mamba_pool.replayssm_write_pos
|
||||||
)
|
)
|
||||||
if write_pos_buf is not None:
|
if write_pos_buf is not None:
|
||||||
cache_len -= int(write_pos_buf[req.mamba_pool_idx].item())
|
cache_len -= int(write_pos_buf[req.kv.mamba_pool_idx].item())
|
||||||
write_pos_buf[req.mamba_pool_idx] = 0
|
write_pos_buf[req.kv.mamba_pool_idx] = 0
|
||||||
|
|
||||||
if is_finished:
|
if is_finished:
|
||||||
if cache_len is None:
|
if cache_len is None:
|
||||||
@@ -559,10 +559,10 @@ class MambaComponent(TreeComponent):
|
|||||||
req
|
req
|
||||||
)
|
)
|
||||||
active_value = (
|
active_value = (
|
||||||
req.mamba_ping_pong_track_buffer[keep_idx].unsqueeze(-1).clone()
|
req.kv.mamba_ping_pong_track_buffer[keep_idx].unsqueeze(-1).clone()
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
active_value = req.mamba_pool_idx.unsqueeze(-1).clone()
|
active_value = req.kv.mamba_pool_idx.unsqueeze(-1).clone()
|
||||||
if self.int8_ckpt_pool is not None:
|
if self.int8_ckpt_pool is not None:
|
||||||
insert_params.mamba_value = self._commit_int8_checkpoint(active_value)
|
insert_params.mamba_value = self._commit_int8_checkpoint(active_value)
|
||||||
else:
|
else:
|
||||||
@@ -584,7 +584,7 @@ class MambaComponent(TreeComponent):
|
|||||||
self.cache.req_to_token_pool.mamba_allocator.free(src_active)
|
self.cache.req_to_token_pool.mamba_allocator.free(src_active)
|
||||||
else:
|
else:
|
||||||
mamba_value_donated = self._commit_int8_checkpoint(
|
mamba_value_donated = self._commit_int8_checkpoint(
|
||||||
req.mamba_pool_idx.view(-1)
|
req.kv.mamba_pool_idx.view(-1)
|
||||||
)
|
)
|
||||||
elif self.cache.enable_mamba_extra_buffer:
|
elif self.cache.enable_mamba_extra_buffer:
|
||||||
new_slot = self._alloc_mamba_slot()
|
new_slot = self._alloc_mamba_slot()
|
||||||
@@ -599,7 +599,7 @@ class MambaComponent(TreeComponent):
|
|||||||
# virtual->physical (identity for the non-unified memory pool) first.
|
# virtual->physical (identity for the non-unified memory pool) first.
|
||||||
translate = self.cache.req_to_token_pool.translate_mamba_indices
|
translate = self.cache.req_to_token_pool.translate_mamba_indices
|
||||||
self.cache.req_to_token_pool.mamba_pool.copy_from(
|
self.cache.req_to_token_pool.mamba_pool.copy_from(
|
||||||
translate(req.mamba_pool_idx.unsqueeze(0)),
|
translate(req.kv.mamba_pool_idx.unsqueeze(0)),
|
||||||
translate(mamba_value_donated),
|
translate(mamba_value_donated),
|
||||||
)
|
)
|
||||||
insert_params.mamba_value = mamba_value_donated
|
insert_params.mamba_value = mamba_value_donated
|
||||||
@@ -647,7 +647,7 @@ class MambaComponent(TreeComponent):
|
|||||||
insert_result is None or insert_result.mamba_exist
|
insert_result is None or insert_result.mamba_exist
|
||||||
):
|
):
|
||||||
self._free_mamba_value(insert_params.mamba_value)
|
self._free_mamba_value(insert_params.mamba_value)
|
||||||
req.mamba_last_track_seqlen = None
|
req.kv.mamba_last_track_seqlen = None
|
||||||
|
|
||||||
def build_external_linker_transfer(
|
def build_external_linker_transfer(
|
||||||
self,
|
self,
|
||||||
@@ -669,7 +669,7 @@ class MambaComponent(TreeComponent):
|
|||||||
) -> PrepareLoadBackResult:
|
) -> PrepareLoadBackResult:
|
||||||
if (
|
if (
|
||||||
req is None
|
req is None
|
||||||
or req.mamba_pool_idx is not None
|
or req.kv.holds_mamba
|
||||||
or not self.tree_core.component_has_host_value_only(
|
or not self.tree_core.component_has_host_value_only(
|
||||||
node_id, self.component_type
|
node_id, self.component_type
|
||||||
)
|
)
|
||||||
@@ -680,7 +680,7 @@ class MambaComponent(TreeComponent):
|
|||||||
self.cache.evict_for_alloc(EvictParams(num_tokens=0, mamba_num=1))
|
self.cache.evict_for_alloc(EvictParams(num_tokens=0, mamba_num=1))
|
||||||
dst = self.cache.req_to_token_pool.mamba_allocator.alloc(1)
|
dst = self.cache.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
assert dst is not None, "Cannot alloc mamba for load_back"
|
assert dst is not None, "Cannot alloc mamba for load_back"
|
||||||
req.mamba_pool_idx = dst[0]
|
req.kv.mamba_pool_idx = dst[0]
|
||||||
return PrepareLoadBackResult(allocated_mamba_slot=dst)
|
return PrepareLoadBackResult(allocated_mamba_slot=dst)
|
||||||
|
|
||||||
def finalize_load_back(
|
def finalize_load_back(
|
||||||
@@ -689,7 +689,7 @@ class MambaComponent(TreeComponent):
|
|||||||
# A called-off load-back returns the slot prepare allocated and clears req (the H->D copy never ran).
|
# A called-off load-back returns the slot prepare allocated and clears req (the H->D copy never ran).
|
||||||
if not success and prep.allocated_mamba_slot is not None:
|
if not success and prep.allocated_mamba_slot is not None:
|
||||||
self.cache.req_to_token_pool.mamba_allocator.free(prep.allocated_mamba_slot)
|
self.cache.req_to_token_pool.mamba_allocator.free(prep.allocated_mamba_slot)
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
|
|
||||||
def prepare_prefetch(
|
def prepare_prefetch(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -1937,7 +1937,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
) -> tuple[PoolTransfer, dict[ComponentType, list[PoolTransfer]]]:
|
) -> tuple[PoolTransfer, dict[ComponentType, list[PoolTransfer]]]:
|
||||||
"""Build the H->D load-back KV transfer plus per-component aux transfers."""
|
"""Build the H->D load-back KV transfer plus per-component aux transfers."""
|
||||||
# Component hooks take primitives, not Req: extract its fields here.
|
# Component hooks take primitives, not Req: extract its fields here.
|
||||||
mamba_pool_idx = req.mamba_pool_idx if req is not None else None
|
mamba_pool_idx = req.kv.mamba_pool_idx if req is not None else None
|
||||||
node = self.node_by_id(node_id)
|
node = self.node_by_id(node_id)
|
||||||
kv_xfer = self.components_by_type[BASE_COMPONENT_TYPE].build_hicache_transfers(
|
kv_xfer = self.components_by_type[BASE_COMPONENT_TYPE].build_hicache_transfers(
|
||||||
node, CacheTransferPhase.LOAD_BACK
|
node, CacheTransferPhase.LOAD_BACK
|
||||||
|
|||||||
@@ -2870,7 +2870,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
def swa_retain_floor(self, req) -> int | None:
|
def swa_retain_floor(self, req) -> int | None:
|
||||||
if not self.is_mamba_enabled or self._sliding_window_size is None:
|
if not self.is_mamba_enabled or self._sliding_window_size is None:
|
||||||
return None
|
return None
|
||||||
checkpoint = req.mamba_last_track_seqlen
|
checkpoint = req.kv.mamba_last_track_seqlen
|
||||||
if checkpoint is None:
|
if checkpoint is None:
|
||||||
return None
|
return None
|
||||||
return checkpoint - self._sliding_window_size
|
return checkpoint - self._sliding_window_size
|
||||||
|
|||||||
@@ -53,14 +53,6 @@ class SessionSlot:
|
|||||||
# releases only what it took (may share the node with another req).
|
# releases only what it took (may share the node with another req).
|
||||||
skip_lock_node_ids: dict = field(default_factory=dict)
|
skip_lock_node_ids: dict = field(default_factory=dict)
|
||||||
|
|
||||||
# Mamba states
|
|
||||||
mamba_pool_idx: Any = None
|
|
||||||
mamba_ping_pong_track_buffer: Any = None
|
|
||||||
mamba_next_track_idx: Any = None
|
|
||||||
mamba_last_track_idx: Any = None
|
|
||||||
mamba_last_track_seqlen: Any = None
|
|
||||||
mamba_branching_seqlen: Any = None
|
|
||||||
|
|
||||||
def save_from_req(self, req: Req, is_first: bool):
|
def save_from_req(self, req: Req, is_first: bool):
|
||||||
"""Save KV state from a finishing request into this slot."""
|
"""Save KV state from a finishing request into this slot."""
|
||||||
kv = req.detach_kv()
|
kv = req.detach_kv()
|
||||||
@@ -74,37 +66,14 @@ class SessionSlot:
|
|||||||
# Later turns run on the slot's record (see restore_to_req).
|
# Later turns run on the slot's record (see restore_to_req).
|
||||||
assert kv is self.kv
|
assert kv is self.kv
|
||||||
|
|
||||||
self.mamba_pool_idx = req.mamba_pool_idx
|
|
||||||
self.mamba_ping_pong_track_buffer = req.mamba_ping_pong_track_buffer
|
|
||||||
self.mamba_next_track_idx = req.mamba_next_track_idx
|
|
||||||
self.mamba_last_track_idx = req.mamba_last_track_idx
|
|
||||||
self.mamba_last_track_seqlen = req.mamba_last_track_seqlen
|
|
||||||
self.mamba_branching_seqlen = req.mamba_branching_seqlen
|
|
||||||
|
|
||||||
# The mamba state moved to the slot too; clear the req's references so a
|
|
||||||
# later alloc/retract path cannot mistake slot-owned state for its own.
|
|
||||||
req.mamba_pool_idx = None
|
|
||||||
req.mamba_ping_pong_track_buffer = None
|
|
||||||
req.mamba_next_track_idx = None
|
|
||||||
req.mamba_last_track_idx = None
|
|
||||||
req.mamba_last_track_seqlen = None
|
|
||||||
req.mamba_branching_seqlen = None
|
|
||||||
|
|
||||||
def restore_to_req(self, req: Req):
|
def restore_to_req(self, req: Req):
|
||||||
"""Restore KV state from this slot into an incoming request."""
|
"""Restore KV state from this slot into an incoming request."""
|
||||||
req.kv = self.kv
|
req.kv = self.kv
|
||||||
req.swa_uuid_for_lock = self.swa_uuid_for_lock
|
req.swa_uuid_for_lock = self.swa_uuid_for_lock
|
||||||
req.skip_lock_node_ids = self.skip_lock_node_ids
|
req.skip_lock_node_ids = self.skip_lock_node_ids
|
||||||
|
|
||||||
req.mamba_pool_idx = self.mamba_pool_idx
|
# NOTE: the slot keeps sharing the record it just handed out. During
|
||||||
req.mamba_ping_pong_track_buffer = self.mamba_ping_pong_track_buffer
|
# chunked prefill, a request may be rejected by
|
||||||
req.mamba_next_track_idx = self.mamba_next_track_idx
|
|
||||||
req.mamba_last_track_idx = self.mamba_last_track_idx
|
|
||||||
req.mamba_last_track_seqlen = self.mamba_last_track_seqlen
|
|
||||||
req.mamba_branching_seqlen = self.mamba_branching_seqlen
|
|
||||||
|
|
||||||
# NOTE: req_pool_idx and mamba_pool_idx are intentionally NOT cleared
|
|
||||||
# from the slot. During chunked prefill, a request may be rejected by
|
|
||||||
# the scheduler (e.g. budget exhausted) and retried in the next cycle.
|
# the scheduler (e.g. budget exhausted) and retried in the next cycle.
|
||||||
# Each retry calls match_prefix -> restore_to_req again, so the slot
|
# Each retry calls match_prefix -> restore_to_req again, so the slot
|
||||||
# must remain intact for idempotent restoration.
|
# must remain intact for idempotent restoration.
|
||||||
@@ -176,7 +145,7 @@ class StreamingSession(BasePrefixCache):
|
|||||||
return session_id in self.slots
|
return session_id in self.slots
|
||||||
|
|
||||||
def any_holding_kv(self) -> bool:
|
def any_holding_kv(self) -> bool:
|
||||||
return any(s.kv.is_held for s in self.slots.values())
|
return any(s.kv.holds_kv for s in self.slots.values())
|
||||||
|
|
||||||
# -- Try-handle entries for composition (see class docstring) --
|
# -- Try-handle entries for composition (see class docstring) --
|
||||||
|
|
||||||
@@ -204,7 +173,7 @@ class StreamingSession(BasePrefixCache):
|
|||||||
if not _is_streaming(req):
|
if not _is_streaming(req):
|
||||||
return None
|
return None
|
||||||
slot = self.slots.get(req.session.session_id)
|
slot = self.slots.get(req.session.session_id)
|
||||||
if slot is None or not slot.kv.is_held:
|
if slot is None or not slot.kv.holds_kv:
|
||||||
return None
|
return None
|
||||||
if req.to_finish is not None:
|
if req.to_finish is not None:
|
||||||
req.session.abort_req()
|
req.session.abort_req()
|
||||||
@@ -310,25 +279,17 @@ class StreamingSession(BasePrefixCache):
|
|||||||
kv = req.detach_kv()
|
kv = req.detach_kv()
|
||||||
if slot is None:
|
if slot is None:
|
||||||
# First-request mid-processing abort: create ephemeral
|
# First-request mid-processing abort: create ephemeral
|
||||||
# slot from req state so release_session handles cleanup.
|
# slot from req state so release_session handles cleanup;
|
||||||
# Include last_node from the req so
|
# the detached record carries the mamba refs for
|
||||||
# release_session calls dec_lock_ref on the tree lock.
|
# _free_slot_mamba, and last_node lets release_session
|
||||||
# Also carry the mamba refs over so _free_slot_mamba can
|
# dec_lock_ref the tree lock.
|
||||||
# return the (possibly extra_buffer ping-pong) slots to
|
|
||||||
# the mamba pool; otherwise the abort orphans them.
|
|
||||||
slot = SessionSlot(
|
slot = SessionSlot(
|
||||||
kv=kv,
|
kv=kv,
|
||||||
last_node=req.last_node,
|
last_node=req.last_node,
|
||||||
swa_uuid_for_lock=req.swa_uuid_for_lock,
|
swa_uuid_for_lock=req.swa_uuid_for_lock,
|
||||||
skip_lock_node_ids=req.skip_lock_node_ids,
|
skip_lock_node_ids=req.skip_lock_node_ids,
|
||||||
mamba_pool_idx=req.mamba_pool_idx,
|
|
||||||
mamba_ping_pong_track_buffer=req.mamba_ping_pong_track_buffer,
|
|
||||||
)
|
)
|
||||||
self.slots[session_id] = slot
|
self.slots[session_id] = slot
|
||||||
# Slot now owns the mamba state — drop the req's refs so
|
|
||||||
# the abort fall-through doesn't double-free.
|
|
||||||
req.mamba_pool_idx = None
|
|
||||||
req.mamba_ping_pong_track_buffer = None
|
|
||||||
else:
|
else:
|
||||||
assert kv is slot.kv
|
assert kv is slot.kv
|
||||||
self.release_session(session_id)
|
self.release_session(session_id)
|
||||||
@@ -424,7 +385,7 @@ class StreamingSession(BasePrefixCache):
|
|||||||
protected_len = slot.kv.cache_protected_len
|
protected_len = slot.kv.cache_protected_len
|
||||||
lock_node = slot.last_node
|
lock_node = slot.last_node
|
||||||
tokens_freed = (
|
tokens_freed = (
|
||||||
max(0, slot.kv.kv_allocated_len - protected_len) if slot.kv.is_held else 0
|
max(0, slot.kv.kv_allocated_len - protected_len) if slot.kv.holds_kv else 0
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"Session KV released: %s (%d tokens freed)", session_id, tokens_freed
|
"Session KV released: %s (%d tokens freed)", session_id, tokens_freed
|
||||||
@@ -439,7 +400,7 @@ class StreamingSession(BasePrefixCache):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
if slot.kv.is_held:
|
if slot.kv.holds_kv:
|
||||||
start = protected_len
|
start = protected_len
|
||||||
end = slot.kv.kv_allocated_len
|
end = slot.kv.kv_allocated_len
|
||||||
if start < end:
|
if start < end:
|
||||||
@@ -467,7 +428,7 @@ class StreamingSession(BasePrefixCache):
|
|||||||
active_pool_idxs is not None
|
active_pool_idxs is not None
|
||||||
and slot.kv.req_pool_idx in active_pool_idxs
|
and slot.kv.req_pool_idx in active_pool_idxs
|
||||||
)
|
)
|
||||||
if slot.kv.is_held and not in_batch:
|
if slot.kv.holds_kv and not in_batch:
|
||||||
allocated = ceil_align(slot.kv.kv_allocated_len, self.page_size)
|
allocated = ceil_align(slot.kv.kv_allocated_len, self.page_size)
|
||||||
total += allocated - slot.kv.cache_protected_len
|
total += allocated - slot.kv.cache_protected_len
|
||||||
return total
|
return total
|
||||||
@@ -484,7 +445,7 @@ class StreamingSession(BasePrefixCache):
|
|||||||
active_pool_idxs is not None
|
active_pool_idxs is not None
|
||||||
and slot.kv.req_pool_idx in active_pool_idxs
|
and slot.kv.req_pool_idx in active_pool_idxs
|
||||||
)
|
)
|
||||||
if slot.kv.is_held and not in_batch:
|
if slot.kv.holds_kv and not in_batch:
|
||||||
allocated = ceil_align(slot.kv.kv_allocated_len, self.page_size)
|
allocated = ceil_align(slot.kv.kv_allocated_len, self.page_size)
|
||||||
total += allocated - max(
|
total += allocated - max(
|
||||||
slot.kv.cache_protected_len, slot.kv.swa_evicted_seqlen
|
slot.kv.cache_protected_len, slot.kv.swa_evicted_seqlen
|
||||||
@@ -498,7 +459,7 @@ class StreamingSession(BasePrefixCache):
|
|||||||
in_batch = (
|
in_batch = (
|
||||||
active_pool_idxs is not None and s.kv.req_pool_idx in active_pool_idxs
|
active_pool_idxs is not None and s.kv.req_pool_idx in active_pool_idxs
|
||||||
)
|
)
|
||||||
return s.kv.is_held and not in_batch
|
return s.kv.holds_kv and not in_batch
|
||||||
|
|
||||||
return sum(_owned(s) for s in self.slots.values())
|
return sum(_owned(s) for s in self.slots.values())
|
||||||
|
|
||||||
@@ -517,10 +478,10 @@ class StreamingSession(BasePrefixCache):
|
|||||||
)
|
)
|
||||||
if in_batch:
|
if in_batch:
|
||||||
continue
|
continue
|
||||||
if slot.mamba_pool_idx is not None:
|
if slot.kv.holds_mamba:
|
||||||
total += slot.mamba_pool_idx.numel()
|
total += slot.kv.mamba_pool_idx.numel()
|
||||||
if slot.mamba_ping_pong_track_buffer is not None:
|
if slot.kv.mamba_ping_pong_track_buffer is not None:
|
||||||
total += slot.mamba_ping_pong_track_buffer.numel()
|
total += slot.kv.mamba_ping_pong_track_buffer.numel()
|
||||||
return total
|
return total
|
||||||
|
|
||||||
def _free_slot_mamba(self, slot: SessionSlot) -> None:
|
def _free_slot_mamba(self, slot: SessionSlot) -> None:
|
||||||
@@ -528,12 +489,12 @@ class StreamingSession(BasePrefixCache):
|
|||||||
mamba_allocator = getattr(self.req_to_token_pool, "mamba_allocator", None)
|
mamba_allocator = getattr(self.req_to_token_pool, "mamba_allocator", None)
|
||||||
if mamba_allocator is None:
|
if mamba_allocator is None:
|
||||||
return
|
return
|
||||||
if slot.mamba_pool_idx is not None:
|
if slot.kv.holds_mamba:
|
||||||
mamba_allocator.free(slot.mamba_pool_idx.unsqueeze(0))
|
mamba_allocator.free(slot.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
slot.mamba_pool_idx = None
|
slot.kv.mamba_pool_idx = None
|
||||||
if slot.mamba_ping_pong_track_buffer is not None:
|
if slot.kv.mamba_ping_pong_track_buffer is not None:
|
||||||
mamba_allocator.free(slot.mamba_ping_pong_track_buffer)
|
mamba_allocator.free(slot.kv.mamba_ping_pong_track_buffer)
|
||||||
slot.mamba_ping_pong_track_buffer = None
|
slot.kv.mamba_ping_pong_track_buffer = None
|
||||||
|
|
||||||
# -- Internal helpers (streaming body bits) --
|
# -- Internal helpers (streaming body bits) --
|
||||||
|
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ class ScriptedReqHandle:
|
|||||||
@property
|
@property
|
||||||
def kv_pages(self) -> int:
|
def kv_pages(self) -> int:
|
||||||
req = self.req
|
req = self.req
|
||||||
if req is None or not req.kv.is_held:
|
if req is None or not req.kv.holds_kv:
|
||||||
return 0
|
return 0
|
||||||
page_size = self.context.scheduler.page_size
|
page_size = self.context.scheduler.page_size
|
||||||
return (req.kv.kv_allocated_len + page_size - 1) // page_size
|
return (req.kv.kv_allocated_len + page_size - 1) // page_size
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import unittest
|
|||||||
from collections import deque
|
from collections import deque
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci
|
from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
@@ -708,7 +709,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
req = FakeRequest()
|
req = FakeRequest()
|
||||||
runner._req_to_token_pool.alloc([req])
|
runner._req_to_token_pool.alloc([req])
|
||||||
runner._req_to_token_pool.auxiliary_state_pool.store_cache(
|
runner._req_to_token_pool.auxiliary_state_pool.store_cache(
|
||||||
req.mamba_pool_idx,
|
req.kv.mamba_pool_idx,
|
||||||
[FakeNativeCache(mx.array([42.0], dtype=mx.float32)), None],
|
[FakeNativeCache(mx.array([42.0], dtype=mx.float32)), None],
|
||||||
[0],
|
[0],
|
||||||
)
|
)
|
||||||
@@ -731,7 +732,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
self.assertIsInstance(pending.cache[1], ContiguousAttentionKVCache)
|
self.assertIsInstance(pending.cache[1], ContiguousAttentionKVCache)
|
||||||
restored = [FakeNativeCache(), None]
|
restored = [FakeNativeCache(), None]
|
||||||
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
|
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
|
||||||
req.mamba_pool_idx, restored, [0]
|
req.kv.mamba_pool_idx, restored, [0]
|
||||||
)
|
)
|
||||||
self.assertEqual(restored[0].state[0].tolist(), [1.0])
|
self.assertEqual(restored[0].state[0].tolist(), [1.0])
|
||||||
|
|
||||||
@@ -782,11 +783,11 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
runner.prefill_finalize(pending)
|
runner.prefill_finalize(pending)
|
||||||
tracked = [FakeNativeCache(), None]
|
tracked = [FakeNativeCache(), None]
|
||||||
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
|
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
|
||||||
req.mamba_ping_pong_track_buffer[0], tracked, [0]
|
req.kv.mamba_ping_pong_track_buffer[0], tracked, [0]
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual([len(x[0]) for x in runner.model.seen_inputs], [64, 6])
|
self.assertEqual([len(x[0]) for x in runner.model.seen_inputs], [64, 6])
|
||||||
self.assertEqual(req.mamba_last_track_seqlen, 64)
|
self.assertEqual(req.kv.mamba_last_track_seqlen, 64)
|
||||||
self.assertEqual(tracked[0].state[0].tolist(), [64.0])
|
self.assertEqual(tracked[0].state[0].tolist(), [64.0])
|
||||||
self.assertEqual(pending.synced_offset, 70)
|
self.assertEqual(pending.synced_offset, 70)
|
||||||
|
|
||||||
@@ -825,7 +826,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
req = FakeRequest()
|
req = FakeRequest()
|
||||||
runner._req_to_token_pool.alloc([req])
|
runner._req_to_token_pool.alloc([req])
|
||||||
runner._req_to_token_pool.auxiliary_state_pool.store_cache(
|
runner._req_to_token_pool.auxiliary_state_pool.store_cache(
|
||||||
req.mamba_pool_idx,
|
req.kv.mamba_pool_idx,
|
||||||
[FakeNativeCache(mx.array([64.0], dtype=mx.float32)), None],
|
[FakeNativeCache(mx.array([64.0], dtype=mx.float32)), None],
|
||||||
[0],
|
[0],
|
||||||
)
|
)
|
||||||
@@ -844,12 +845,12 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
runner.prefill_finalize(pending)
|
runner.prefill_finalize(pending)
|
||||||
tracked = [FakeNativeCache(), None]
|
tracked = [FakeNativeCache(), None]
|
||||||
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
|
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
|
||||||
req.mamba_ping_pong_track_buffer[0], tracked, [0]
|
req.kv.mamba_ping_pong_track_buffer[0], tracked, [0]
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual([len(x[0]) for x in runner.model.seen_inputs], [192, 1])
|
self.assertEqual([len(x[0]) for x in runner.model.seen_inputs], [192, 1])
|
||||||
self.assertEqual(runner.model.seen_auxiliary_states, [[64.0], [192.0]])
|
self.assertEqual(runner.model.seen_auxiliary_states, [[64.0], [192.0]])
|
||||||
self.assertEqual(req.mamba_last_track_seqlen, 256)
|
self.assertEqual(req.kv.mamba_last_track_seqlen, 256)
|
||||||
self.assertEqual(tracked[0].state[0].tolist(), [192.0])
|
self.assertEqual(tracked[0].state[0].tolist(), [192.0])
|
||||||
self.assertEqual(pending.synced_offset, 257)
|
self.assertEqual(pending.synced_offset, 257)
|
||||||
|
|
||||||
@@ -926,11 +927,11 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
self.assertIn(req_indices[0], range(1, pool.size + 1))
|
self.assertIn(req_indices[0], range(1, pool.size + 1))
|
||||||
self.assertIsNotNone(auxiliary_state_idx)
|
self.assertIsNotNone(auxiliary_state_idx)
|
||||||
self.assertIsNone(req.kv.req_pool_idx)
|
self.assertIsNone(req.kv.req_pool_idx)
|
||||||
self.assertIsNotNone(req.mamba_pool_idx)
|
self.assertIsNotNone(req.kv.mamba_pool_idx)
|
||||||
self.assertIs(pool.mamba_allocator, pool.mamba_pool)
|
self.assertIs(pool.mamba_allocator, pool.mamba_pool)
|
||||||
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
||||||
pool.free_auxiliary_state_cache(req)
|
pool.free_auxiliary_state_cache(req)
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertEqual(pool.available_size(), 2)
|
self.assertEqual(pool.available_size(), 2)
|
||||||
self.assertEqual(pool.auxiliary_state_pool.available_size(), 4)
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 4)
|
||||||
|
|
||||||
@@ -944,13 +945,13 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
req = FakeRequest()
|
req = FakeRequest()
|
||||||
pool.alloc([req])
|
pool.alloc([req])
|
||||||
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
req.kv.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
||||||
req.mamba_next_track_idx = 0
|
req.kv.mamba_next_track_idx = 0
|
||||||
|
|
||||||
pool.free_auxiliary_state_cache(req, track_buffer_to_keep=0)
|
pool.free_auxiliary_state_cache(req, track_buffer_to_keep=0)
|
||||||
|
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertIsNone(req.mamba_ping_pong_track_buffer)
|
self.assertIsNone(req.kv.mamba_ping_pong_track_buffer)
|
||||||
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
||||||
|
|
||||||
def test_auxiliary_state_component_inserts_tracked_slot_and_frees_live_slot(self):
|
def test_auxiliary_state_component_inserts_tracked_slot_and_frees_live_slot(self):
|
||||||
@@ -963,9 +964,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
req = FakeRequest()
|
req = FakeRequest()
|
||||||
pool.alloc([req])
|
pool.alloc([req])
|
||||||
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
req.kv.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
||||||
req.mamba_next_track_idx = 0
|
req.kv.mamba_next_track_idx = 0
|
||||||
req.mamba_last_track_seqlen = 64
|
req.kv.mamba_last_track_seqlen = 64
|
||||||
component = MlxAuxiliaryStateComponent(
|
component = MlxAuxiliaryStateComponent(
|
||||||
SimpleNamespace(req_to_token_pool=pool),
|
SimpleNamespace(req_to_token_pool=pool),
|
||||||
SimpleNamespace(enable_mamba_extra_buffer=False),
|
SimpleNamespace(enable_mamba_extra_buffer=False),
|
||||||
@@ -988,9 +989,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
self.assertEqual(cache_len, 64)
|
self.assertEqual(cache_len, 64)
|
||||||
self.assertTrue(getattr(insert_params, "mlx_auxiliary_state_uses_track_slot"))
|
self.assertTrue(getattr(insert_params, "mlx_auxiliary_state_uses_track_slot"))
|
||||||
self.assertEqual(insert_params.mamba_value.tolist(), [2])
|
self.assertEqual(insert_params.mamba_value.tolist(), [2])
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertIsNone(req.mamba_ping_pong_track_buffer)
|
self.assertIsNone(req.kv.mamba_ping_pong_track_buffer)
|
||||||
self.assertIsNone(req.mamba_last_track_seqlen)
|
self.assertIsNone(req.kv.mamba_last_track_seqlen)
|
||||||
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
||||||
|
|
||||||
def test_auxiliary_state_component_unfinished_frees_tracked_source_slot(self):
|
def test_auxiliary_state_component_unfinished_frees_tracked_source_slot(self):
|
||||||
@@ -1003,9 +1004,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
req = FakeRequest()
|
req = FakeRequest()
|
||||||
pool.alloc([req])
|
pool.alloc([req])
|
||||||
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
req.kv.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
||||||
req.mamba_next_track_idx = 0
|
req.kv.mamba_next_track_idx = 0
|
||||||
req.mamba_last_track_seqlen = 64
|
req.kv.mamba_last_track_seqlen = 64
|
||||||
component = MlxAuxiliaryStateComponent(
|
component = MlxAuxiliaryStateComponent(
|
||||||
SimpleNamespace(req_to_token_pool=pool),
|
SimpleNamespace(req_to_token_pool=pool),
|
||||||
SimpleNamespace(enable_mamba_extra_buffer=False),
|
SimpleNamespace(enable_mamba_extra_buffer=False),
|
||||||
@@ -1027,9 +1028,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
|
|
||||||
self.assertEqual(cache_len, 64)
|
self.assertEqual(cache_len, 64)
|
||||||
self.assertEqual(insert_params.mamba_value.tolist(), [3])
|
self.assertEqual(insert_params.mamba_value.tolist(), [3])
|
||||||
self.assertIsNotNone(req.mamba_pool_idx)
|
self.assertIsNotNone(req.kv.mamba_pool_idx)
|
||||||
self.assertIsNone(req.mamba_ping_pong_track_buffer)
|
self.assertIsNone(req.kv.mamba_ping_pong_track_buffer)
|
||||||
self.assertIsNone(req.mamba_last_track_seqlen)
|
self.assertIsNone(req.kv.mamba_last_track_seqlen)
|
||||||
self.assertEqual(pool.auxiliary_state_pool.available_size(), 2)
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 2)
|
||||||
|
|
||||||
def test_auxiliary_state_component_frees_stale_track_slot_when_live_slot_inserted(
|
def test_auxiliary_state_component_frees_stale_track_slot_when_live_slot_inserted(
|
||||||
@@ -1044,8 +1045,8 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
req = FakeRequest()
|
req = FakeRequest()
|
||||||
pool.alloc([req])
|
pool.alloc([req])
|
||||||
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
req.kv.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
||||||
req.mamba_next_track_idx = 0
|
req.kv.mamba_next_track_idx = 0
|
||||||
component = MlxAuxiliaryStateComponent(
|
component = MlxAuxiliaryStateComponent(
|
||||||
SimpleNamespace(req_to_token_pool=pool),
|
SimpleNamespace(req_to_token_pool=pool),
|
||||||
SimpleNamespace(enable_mamba_extra_buffer=False),
|
SimpleNamespace(enable_mamba_extra_buffer=False),
|
||||||
@@ -1068,9 +1069,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
self.assertEqual(cache_len, 7)
|
self.assertEqual(cache_len, 7)
|
||||||
self.assertFalse(getattr(insert_params, "mlx_auxiliary_state_uses_track_slot"))
|
self.assertFalse(getattr(insert_params, "mlx_auxiliary_state_uses_track_slot"))
|
||||||
self.assertEqual(insert_params.mamba_value.tolist(), [1])
|
self.assertEqual(insert_params.mamba_value.tolist(), [1])
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertIsNone(req.mamba_ping_pong_track_buffer)
|
self.assertIsNone(req.kv.mamba_ping_pong_track_buffer)
|
||||||
self.assertIsNone(req.mamba_next_track_idx)
|
self.assertIsNone(req.kv.mamba_next_track_idx)
|
||||||
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
||||||
|
|
||||||
def test_auxiliary_state_component_frees_duplicate_live_slot(self):
|
def test_auxiliary_state_component_frees_duplicate_live_slot(self):
|
||||||
@@ -1102,7 +1103,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|||||||
insert_params=insert_params,
|
insert_params=insert_params,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertEqual(pool.auxiliary_state_pool.available_size(), 4)
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 4)
|
||||||
|
|
||||||
|
|
||||||
@@ -1518,8 +1519,7 @@ if _HAS_MLX:
|
|||||||
|
|
||||||
class FakeRequest:
|
class FakeRequest:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.kv = SimpleNamespace(req_pool_idx=None)
|
self.kv = ReqKvInfo()
|
||||||
self.mamba_pool_idx = None
|
|
||||||
self.inflight_middle_chunks = 0
|
self.inflight_middle_chunks = 0
|
||||||
|
|
||||||
class FakeTpWorker:
|
class FakeTpWorker:
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ import unittest
|
|||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest import mock
|
from unittest import mock
|
||||||
|
|
||||||
|
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci
|
from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
@@ -120,9 +121,7 @@ def _hybrid_stub_for_initialize(
|
|||||||
def _fake_req():
|
def _fake_req():
|
||||||
return SimpleNamespace(
|
return SimpleNamespace(
|
||||||
inflight_middle_chunks=0,
|
inflight_middle_chunks=0,
|
||||||
kv=SimpleNamespace(req_pool_idx=None),
|
kv=ReqKvInfo(),
|
||||||
mamba_pool_idx=None,
|
|
||||||
mamba_ping_pong_track_buffer=None,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -284,13 +283,13 @@ class TestMlxHybridInitializeAllocation(CustomTestCase):
|
|||||||
req = _fake_req()
|
req = _fake_req()
|
||||||
self.assertIsNotNone(pool.alloc([req]))
|
self.assertIsNotNone(pool.alloc([req]))
|
||||||
pool.free(req) # as release_kv_cache does after ChunkCache
|
pool.free(req) # as release_kv_cache does after ChunkCache
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertEqual(pool.auxiliary_state_pool.available_size(), aux_capacity)
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), aux_capacity)
|
||||||
|
|
||||||
def test_radix_enabled_free_does_not_touch_aux_slot(self):
|
def test_radix_enabled_free_does_not_touch_aux_slot(self):
|
||||||
# Retention contract: with the radix cache enabled the tree component
|
# Retention contract: with the radix cache enabled the tree component
|
||||||
# owns auxiliary release (it frees or adopts the slot and nulls
|
# owns auxiliary release (it frees or adopts the slot and nulls
|
||||||
# req.mamba_pool_idx BEFORE the row is freed). pool.free(req) must
|
# req.kv.mamba_pool_idx BEFORE the row is freed). pool.free(req) must
|
||||||
# therefore never release auxiliary slots itself -- even if called
|
# therefore never release auxiliary slots itself -- even if called
|
||||||
# while mamba_pool_idx is still set -- or a tree-owned snapshot slot
|
# while mamba_pool_idx is still set -- or a tree-owned snapshot slot
|
||||||
# could be recycled under a live radix node.
|
# could be recycled under a live radix node.
|
||||||
@@ -306,7 +305,7 @@ class TestMlxHybridInitializeAllocation(CustomTestCase):
|
|||||||
req = _fake_req()
|
req = _fake_req()
|
||||||
pool.alloc([req])
|
pool.alloc([req])
|
||||||
pool.free(req)
|
pool.free(req)
|
||||||
self.assertIsNotNone(req.mamba_pool_idx) # slot NOT released by free()
|
self.assertIsNotNone(req.kv.mamba_pool_idx) # slot NOT released by free()
|
||||||
self.assertEqual(pool.auxiliary_state_pool.available_size(), free_before - 1)
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), free_before - 1)
|
||||||
|
|
||||||
def test_default_aux_sizing_uses_shared_ratio(self):
|
def test_default_aux_sizing_uses_shared_ratio(self):
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ from types import SimpleNamespace
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci
|
from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci
|
||||||
@@ -169,7 +170,7 @@ class _FakeReq:
|
|||||||
self.rid = rid
|
self.rid = rid
|
||||||
self.prefix_indices = torch.empty(0, dtype=torch.long)
|
self.prefix_indices = torch.empty(0, dtype=torch.long)
|
||||||
self.fill_ids = [0]
|
self.fill_ids = [0]
|
||||||
self.kv = SimpleNamespace(req_pool_idx=req_pool_idx)
|
self.kv = ReqKvInfo(req_pool_idx=req_pool_idx)
|
||||||
# Mirrors Req's chunk-finality contract read by
|
# Mirrors Req's chunk-finality contract read by
|
||||||
# MlxTpModelWorker._chunk_needs_logits: extend_range=None means
|
# MlxTpModelWorker._chunk_needs_logits: extend_range=None means
|
||||||
# "not truncated" (final chunk / plain prefill).
|
# "not truncated" (final chunk / plain prefill).
|
||||||
|
|||||||
@@ -45,8 +45,8 @@ def _track_seqlen(*, tree_page: int, prefix_len: int, extend_len: int) -> int:
|
|||||||
)
|
)
|
||||||
req.prefix_indices = torch.arange(prefix_len, dtype=torch.int64)
|
req.prefix_indices = torch.arange(prefix_len, dtype=torch.int64)
|
||||||
req.set_extend_range(prefix_len, prefix_len + extend_len)
|
req.set_extend_range(prefix_len, prefix_len + extend_len)
|
||||||
req.mamba_ping_pong_track_buffer = torch.tensor([0, 1], dtype=torch.int64)
|
req.kv.mamba_ping_pong_track_buffer = torch.tensor([0, 1], dtype=torch.int64)
|
||||||
req.mamba_next_track_idx = 0
|
req.kv.mamba_next_track_idx = 0
|
||||||
req.mamba_branching_seqlen = None
|
req.mamba_branching_seqlen = None
|
||||||
|
|
||||||
batch = ScheduleBatch(reqs=[req])
|
batch = ScheduleBatch(reqs=[req])
|
||||||
@@ -58,7 +58,7 @@ def _track_seqlen(*, tree_page: int, prefix_len: int, extend_len: int) -> int:
|
|||||||
batch.req_to_token_pool.get_mamba_ping_pong_other_idx.return_value = 1
|
batch.req_to_token_pool.get_mamba_ping_pong_other_idx.return_value = 1
|
||||||
|
|
||||||
batch._mamba_radix_cache_v2_req_prepare_for_extend(req)
|
batch._mamba_radix_cache_v2_req_prepare_for_extend(req)
|
||||||
return req.mamba_last_track_seqlen
|
return req.kv.mamba_last_track_seqlen
|
||||||
|
|
||||||
|
|
||||||
class TestMambaCheckpointDepth(unittest.TestCase):
|
class TestMambaCheckpointDepth(unittest.TestCase):
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ def _make_req(
|
|||||||
req.positional_embed_overrides = None
|
req.positional_embed_overrides = None
|
||||||
req.extra_key = None
|
req.extra_key = None
|
||||||
req.cache_salt = None
|
req.cache_salt = None
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
req.sampling_params = SimpleNamespace(max_new_tokens=128, ignore_eos=False)
|
req.sampling_params = SimpleNamespace(max_new_tokens=128, ignore_eos=False)
|
||||||
return req
|
return req
|
||||||
|
|
||||||
|
|||||||
@@ -142,7 +142,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
assert req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size
|
assert req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size
|
||||||
|
|
||||||
# alloc req without free mamba cache
|
# alloc req without free mamba cache
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
req_to_token_pool.alloc([req])
|
req_to_token_pool.alloc([req])
|
||||||
req_to_token_pool.free(req)
|
req_to_token_pool.free(req)
|
||||||
assert req_to_token_pool.available_size() == max_num_reqs
|
assert req_to_token_pool.available_size() == max_num_reqs
|
||||||
@@ -233,7 +233,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key,
|
key=key,
|
||||||
value=req1_kv_indices[: len(key)],
|
value=req1_kv_indices[: len(key)],
|
||||||
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
prefix_len = result.prefix_len
|
prefix_len = result.prefix_len
|
||||||
@@ -251,7 +251,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key,
|
key=key,
|
||||||
value=req2_kv_indices[: len(key)],
|
value=req2_kv_indices[: len(key)],
|
||||||
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
prefix_len = result.prefix_len
|
prefix_len = result.prefix_len
|
||||||
@@ -270,7 +270,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key,
|
key=key,
|
||||||
value=req3_kv_indices[: len(key)],
|
value=req3_kv_indices[: len(key)],
|
||||||
mamba_value=req3.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req3.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
prefix_len = result.prefix_len
|
prefix_len = result.prefix_len
|
||||||
@@ -288,7 +288,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key,
|
key=key,
|
||||||
value=req4_kv_indices[: len(key)],
|
value=req4_kv_indices[: len(key)],
|
||||||
mamba_value=req4.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req4.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
prefix_len = result.prefix_len
|
prefix_len = result.prefix_len
|
||||||
@@ -372,13 +372,13 @@ class TestMamba(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
kv_indices, last_node = result.device_indices, result.last_device_node
|
kv_indices, last_node = result.device_indices, result.last_device_node
|
||||||
assert req9.mamba_pool_idx is not None
|
assert req9.kv.holds_mamba
|
||||||
assert torch.all(
|
assert torch.all(
|
||||||
mamba_pool.mamba_cache.conv[0][:, req9.mamba_pool_idx]
|
mamba_pool.mamba_cache.conv[0][:, req9.kv.mamba_pool_idx]
|
||||||
== mamba_pool.mamba_cache.conv[0][:, last_node.mamba_value]
|
== mamba_pool.mamba_cache.conv[0][:, last_node.mamba_value]
|
||||||
)
|
)
|
||||||
assert torch.all(
|
assert torch.all(
|
||||||
mamba_pool.mamba_cache.temporal[:, req9.mamba_pool_idx]
|
mamba_pool.mamba_cache.temporal[:, req9.kv.mamba_pool_idx]
|
||||||
== mamba_pool.mamba_cache.temporal[:, last_node.mamba_value]
|
== mamba_pool.mamba_cache.temporal[:, last_node.mamba_value]
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -404,7 +404,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=RadixKey(array("q", token_ids)),
|
key=RadixKey(array("q", token_ids)),
|
||||||
value=kv,
|
value=kv,
|
||||||
mamba_value=req.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -459,7 +459,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key1,
|
key=key1,
|
||||||
value=allocator.alloc(3)[: len(key1)],
|
value=allocator.alloc(3)[: len(key1)],
|
||||||
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
events = tree.take_events()
|
events = tree.take_events()
|
||||||
@@ -476,7 +476,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key2,
|
key=key2,
|
||||||
value=allocator.alloc(5)[: len(key2)],
|
value=allocator.alloc(5)[: len(key2)],
|
||||||
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
events = tree.take_events()
|
events = tree.take_events()
|
||||||
@@ -515,7 +515,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key1,
|
key=key1,
|
||||||
value=allocator.alloc(4)[: len(key1)],
|
value=allocator.alloc(4)[: len(key1)],
|
||||||
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
first_insert_events = [
|
first_insert_events = [
|
||||||
@@ -530,7 +530,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key2,
|
key=key2,
|
||||||
value=allocator.alloc(4)[: len(key2)],
|
value=allocator.alloc(4)[: len(key2)],
|
||||||
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
second_insert_events = [
|
second_insert_events = [
|
||||||
@@ -771,7 +771,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key1,
|
key=key1,
|
||||||
value=allocator.alloc(3)[: len(key1)],
|
value=allocator.alloc(3)[: len(key1)],
|
||||||
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
assert allocator.available_size() == initial_avail - 3
|
assert allocator.available_size() == initial_avail - 3
|
||||||
@@ -784,7 +784,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key2,
|
key=key2,
|
||||||
value=allocator.alloc(7)[: len(key2)],
|
value=allocator.alloc(7)[: len(key2)],
|
||||||
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
prev_prefix_len=0,
|
prev_prefix_len=0,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -802,7 +802,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key3,
|
key=key3,
|
||||||
value=allocator.alloc(8)[: len(key3)],
|
value=allocator.alloc(8)[: len(key3)],
|
||||||
mamba_value=req3.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req3.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
prev_prefix_len=2,
|
prev_prefix_len=2,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -819,7 +819,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=key4,
|
key=key4,
|
||||||
value=allocator.alloc(9)[: len(key4)],
|
value=allocator.alloc(9)[: len(key4)],
|
||||||
mamba_value=req4.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req4.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
prev_prefix_len=8,
|
prev_prefix_len=8,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -39,7 +39,6 @@ class _StubReq:
|
|||||||
self.best_match_node = None
|
self.best_match_node = None
|
||||||
self.host_hit_length = None
|
self.host_hit_length = None
|
||||||
self.num_matched_prefix_tokens = 0
|
self.num_matched_prefix_tokens = 0
|
||||||
self.mamba_branching_seqlen = None
|
|
||||||
self.kv = SimpleNamespace(cache_protected_len=None)
|
self.kv = SimpleNamespace(cache_protected_len=None)
|
||||||
|
|
||||||
def _compute_max_prefix_len(self, input_len):
|
def _compute_max_prefix_len(self, input_len):
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ def _req_and_pool():
|
|||||||
req.kv = ReqKvInfo(req_pool_idx=0)
|
req.kv = ReqKvInfo(req_pool_idx=0)
|
||||||
req.origin_input_ids = [1, 2]
|
req.origin_input_ids = [1, 2]
|
||||||
req.output_ids = [3]
|
req.output_ids = [3]
|
||||||
req.mamba_pool_idx = torch.tensor(1)
|
req.kv.mamba_pool_idx = torch.tensor(1)
|
||||||
|
|
||||||
pool = object.__new__(HybridReqToTokenPool)
|
pool = object.__new__(HybridReqToTokenPool)
|
||||||
pool.req_to_token = torch.zeros(1, 8, dtype=torch.int64)
|
pool.req_to_token = torch.zeros(1, 8, dtype=torch.int64)
|
||||||
@@ -59,7 +59,7 @@ class TestRetractionMambaBackup(unittest.TestCase):
|
|||||||
allocator = _Allocator(carries_mamba=False)
|
allocator = _Allocator(carries_mamba=False)
|
||||||
|
|
||||||
req.offload_kv_cache(pool, allocator)
|
req.offload_kv_cache(pool, allocator)
|
||||||
self.assertIs(req.retraction_backup.mamba_cpu, MAMBA_STATE)
|
self.assertIs(req.kv.retraction_backup.mamba_cpu, MAMBA_STATE)
|
||||||
|
|
||||||
req.load_kv_cache(pool, allocator)
|
req.load_kv_cache(pool, allocator)
|
||||||
self.assertIs(pool.mamba_pool.loaded, MAMBA_STATE)
|
self.assertIs(pool.mamba_pool.loaded, MAMBA_STATE)
|
||||||
@@ -69,7 +69,7 @@ class TestRetractionMambaBackup(unittest.TestCase):
|
|||||||
allocator = _Allocator(carries_mamba=True)
|
allocator = _Allocator(carries_mamba=True)
|
||||||
|
|
||||||
req.offload_kv_cache(pool, allocator)
|
req.offload_kv_cache(pool, allocator)
|
||||||
self.assertIsNone(req.retraction_backup.mamba_cpu)
|
self.assertIsNone(req.kv.retraction_backup.mamba_cpu)
|
||||||
|
|
||||||
req.load_kv_cache(pool, allocator)
|
req.load_kv_cache(pool, allocator)
|
||||||
self.assertIsNone(pool.mamba_pool.loaded)
|
self.assertIsNone(pool.mamba_pool.loaded)
|
||||||
|
|||||||
@@ -81,12 +81,6 @@ class _FakeReq:
|
|||||||
self.last_node = None
|
self.last_node = None
|
||||||
self.swa_uuid_for_lock = None
|
self.swa_uuid_for_lock = None
|
||||||
self.skip_lock_node_ids = {}
|
self.skip_lock_node_ids = {}
|
||||||
self.mamba_pool_idx = None
|
|
||||||
self.mamba_ping_pong_track_buffer = None
|
|
||||||
self.mamba_next_track_idx = None
|
|
||||||
self.mamba_last_track_idx = None
|
|
||||||
self.mamba_last_track_seqlen = None
|
|
||||||
self.mamba_branching_seqlen = None
|
|
||||||
self.to_finish = None
|
self.to_finish = None
|
||||||
self.finished_reason = None
|
self.finished_reason = None
|
||||||
self.finished_len = None
|
self.finished_len = None
|
||||||
@@ -97,11 +91,12 @@ class _FakeReq:
|
|||||||
|
|
||||||
|
|
||||||
def test_session_slot_round_trip_preserves_mamba_state():
|
def test_session_slot_round_trip_preserves_mamba_state():
|
||||||
|
# The mamba state rides in the shared ReqKvInfo record. mamba_branching_seqlen
|
||||||
|
# is a per-turn match observation on the Req and is not preserved by the slot.
|
||||||
req = _FakeReq("session-a", req_pool_idx=0, committed=4, allocated=4)
|
req = _FakeReq("session-a", req_pool_idx=0, committed=4, allocated=4)
|
||||||
req.mamba_next_track_idx = 1
|
req.kv.mamba_next_track_idx = 1
|
||||||
req.mamba_last_track_idx = 0
|
req.kv.mamba_last_track_idx = 0
|
||||||
req.mamba_last_track_seqlen = 3
|
req.kv.mamba_last_track_seqlen = 3
|
||||||
req.mamba_branching_seqlen = 2
|
|
||||||
|
|
||||||
slot = SessionSlot()
|
slot = SessionSlot()
|
||||||
slot.save_from_req(req, is_first=True)
|
slot.save_from_req(req, is_first=True)
|
||||||
@@ -109,10 +104,9 @@ def test_session_slot_round_trip_preserves_mamba_state():
|
|||||||
next_req = _FakeReq("session-a", req_pool_idx=1, committed=0, allocated=0)
|
next_req = _FakeReq("session-a", req_pool_idx=1, committed=0, allocated=0)
|
||||||
slot.restore_to_req(next_req)
|
slot.restore_to_req(next_req)
|
||||||
|
|
||||||
assert next_req.mamba_next_track_idx == 1
|
assert next_req.kv.mamba_next_track_idx == 1
|
||||||
assert next_req.mamba_last_track_idx == 0
|
assert next_req.kv.mamba_last_track_idx == 0
|
||||||
assert next_req.mamba_last_track_seqlen == 3
|
assert next_req.kv.mamba_last_track_seqlen == 3
|
||||||
assert next_req.mamba_branching_seqlen == 2
|
|
||||||
|
|
||||||
|
|
||||||
def test_preabort_detaches_session_and_preserves_slot():
|
def test_preabort_detaches_session_and_preserves_slot():
|
||||||
|
|||||||
@@ -334,7 +334,7 @@ def _insert_seq(env, seq):
|
|||||||
mamba_val = None
|
mamba_val = None
|
||||||
if env.has_mamba:
|
if env.has_mamba:
|
||||||
req = env.make_req()
|
req = env.make_req()
|
||||||
mamba_val = req.mamba_pool_idx.unsqueeze(0)
|
mamba_val = req.kv.mamba_pool_idx.unsqueeze(0)
|
||||||
key = RadixKey(array("q", seq))
|
key = RadixKey(array("q", seq))
|
||||||
env.tree.insert(InsertParams(key=key, value=v[: len(key)], mamba_value=mamba_val))
|
env.tree.insert(InsertParams(key=key, value=v[: len(key)], mamba_value=mamba_val))
|
||||||
return True
|
return True
|
||||||
@@ -356,7 +356,7 @@ def _fill_no_evict(env):
|
|||||||
mamba_val = None
|
mamba_val = None
|
||||||
if env.has_mamba:
|
if env.has_mamba:
|
||||||
req = env.make_req()
|
req = env.make_req()
|
||||||
mamba_val = req.mamba_pool_idx.unsqueeze(0)
|
mamba_val = req.kv.mamba_pool_idx.unsqueeze(0)
|
||||||
key = RadixKey(array("q", seq))
|
key = RadixKey(array("q", seq))
|
||||||
env.tree.insert(
|
env.tree.insert(
|
||||||
InsertParams(key=key, value=v[: len(key)], mamba_value=mamba_val)
|
InsertParams(key=key, value=v[: len(key)], mamba_value=mamba_val)
|
||||||
|
|||||||
@@ -601,7 +601,7 @@ class TestUnifiedRadixAllocationEvictionRealComponents(CustomTestCase):
|
|||||||
sampling_params=SamplingParams(temperature=0, max_new_tokens=1),
|
sampling_params=SamplingParams(temperature=0, max_new_tokens=1),
|
||||||
)
|
)
|
||||||
req_to_token_pool.alloc([req])
|
req_to_token_pool.alloc([req])
|
||||||
params.mamba_value = req.mamba_pool_idx.unsqueeze(0)
|
params.mamba_value = req.kv.mamba_pool_idx.unsqueeze(0)
|
||||||
cache.insert(params)
|
cache.insert(params)
|
||||||
|
|
||||||
def _build_internal_chain(self, component_type, enable_session_radix_cache):
|
def _build_internal_chain(self, component_type, enable_session_radix_cache):
|
||||||
@@ -1075,7 +1075,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
params = InsertParams(key=key, value=value[: len(key)], priority=priority)
|
params = InsertParams(key=key, value=value[: len(key)], priority=priority)
|
||||||
if self.cfg.has_mamba:
|
if self.cfg.has_mamba:
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
params.mamba_value = req.mamba_pool_idx.unsqueeze(0)
|
params.mamba_value = req.kv.mamba_pool_idx.unsqueeze(0)
|
||||||
return cache.insert(params)
|
return cache.insert(params)
|
||||||
|
|
||||||
def test_insert_and_match_basic(self):
|
def test_insert_and_match_basic(self):
|
||||||
@@ -1189,7 +1189,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
)
|
)
|
||||||
if self.cfg.has_mamba:
|
if self.cfg.has_mamba:
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
params.mamba_value = req.mamba_pool_idx.unsqueeze(0)
|
params.mamba_value = req.kv.mamba_pool_idx.unsqueeze(0)
|
||||||
result = cache.insert(params)
|
result = cache.insert(params)
|
||||||
self.assertEqual(result.prefix_len, len(seq_1p))
|
self.assertEqual(result.prefix_len, len(seq_1p))
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
@@ -1213,7 +1213,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
)
|
)
|
||||||
if self.cfg.has_mamba:
|
if self.cfg.has_mamba:
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
params.mamba_value = req.mamba_pool_idx.unsqueeze(0)
|
params.mamba_value = req.kv.mamba_pool_idx.unsqueeze(0)
|
||||||
result = cache.insert(params)
|
result = cache.insert(params)
|
||||||
self.assertEqual(result.prefix_len, len(seq_2p))
|
self.assertEqual(result.prefix_len, len(seq_2p))
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
@@ -1246,7 +1246,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
|
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
|
||||||
)
|
)
|
||||||
if self.cfg.has_mamba:
|
if self.cfg.has_mamba:
|
||||||
req.mamba_last_track_seqlen = kv_len
|
req.kv.mamba_last_track_seqlen = kv_len
|
||||||
|
|
||||||
cache.cache_finished_req(
|
cache.cache_finished_req(
|
||||||
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
|
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
|
||||||
@@ -1283,7 +1283,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
req.swa_uuid_for_lock = None
|
req.swa_uuid_for_lock = None
|
||||||
req.extra_key = None
|
req.extra_key = None
|
||||||
if self.cfg.has_mamba:
|
if self.cfg.has_mamba:
|
||||||
req.mamba_last_track_seqlen = kv_len
|
req.kv.mamba_last_track_seqlen = kv_len
|
||||||
req.reasoning_tokens = 1
|
req.reasoning_tokens = 1
|
||||||
|
|
||||||
# cache_finished_req reads get_serving().strip_thinking_cache
|
# cache_finished_req reads get_serving().strip_thinking_cache
|
||||||
@@ -1362,7 +1362,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
req.swa_uuid_for_lock = None
|
req.swa_uuid_for_lock = None
|
||||||
req.extra_key = None
|
req.extra_key = None
|
||||||
if self.cfg.has_mamba:
|
if self.cfg.has_mamba:
|
||||||
req.mamba_last_track_seqlen = kv_len
|
req.kv.mamba_last_track_seqlen = kv_len
|
||||||
|
|
||||||
cache.cache_unfinished_req(req)
|
cache.cache_unfinished_req(req)
|
||||||
|
|
||||||
@@ -1507,7 +1507,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
|
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
|
||||||
)
|
)
|
||||||
if self.cfg.has_mamba:
|
if self.cfg.has_mamba:
|
||||||
req.mamba_last_track_seqlen = kv_len
|
req.kv.mamba_last_track_seqlen = kv_len
|
||||||
|
|
||||||
avail_before = allocator.available_size()
|
avail_before = allocator.available_size()
|
||||||
cache.cache_finished_req(
|
cache.cache_finished_req(
|
||||||
@@ -1577,12 +1577,12 @@ class UnifiedRadixCacheSuite:
|
|||||||
MatchPrefixParams(key=RadixKey(array("q", seq)), cow_mamba=True, req=req2)
|
MatchPrefixParams(key=RadixKey(array("q", seq)), cow_mamba=True, req=req2)
|
||||||
)
|
)
|
||||||
self.assertEqual(len(m.device_indices), len(seq))
|
self.assertEqual(len(m.device_indices), len(seq))
|
||||||
self.assertIsNotNone(req2.mamba_pool_idx)
|
self.assertIsNotNone(req2.kv.mamba_pool_idx)
|
||||||
|
|
||||||
src_value = _device_value(cache, m.last_device_node, ComponentType.MAMBA)
|
src_value = _device_value(cache, m.last_device_node, ComponentType.MAMBA)
|
||||||
self.assertTrue(
|
self.assertTrue(
|
||||||
torch.all(
|
torch.all(
|
||||||
mamba_pool.mamba_cache.conv[0][:, req2.mamba_pool_idx]
|
mamba_pool.mamba_cache.conv[0][:, req2.kv.mamba_pool_idx]
|
||||||
== mamba_pool.mamba_cache.conv[0][:, src_value]
|
== mamba_pool.mamba_cache.conv[0][:, src_value]
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -5469,7 +5469,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
|
|
||||||
# Simulate a request without its own mamba slot so load-back allocates one
|
# Simulate a request without its own mamba slot so load-back allocates one
|
||||||
# (that allocation is what a called-off load-back must free + not publish).
|
# (that allocation is what a called-off load-back must free + not publish).
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
avail_before = req_to_token_pool.mamba_allocator.available_size()
|
avail_before = req_to_token_pool.mamba_allocator.available_size()
|
||||||
new_indices, new_node = cache.init_load_back(
|
new_indices, new_node = cache.init_load_back(
|
||||||
InitLoadBackParams(
|
InitLoadBackParams(
|
||||||
@@ -5485,7 +5485,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
self.assertIsNone(_device_value(cache, leaf, ComponentType.FULL))
|
self.assertIsNone(_device_value(cache, leaf, ComponentType.FULL))
|
||||||
self.assertIsNone(_device_value(cache, leaf, ComponentType.MAMBA))
|
self.assertIsNone(_device_value(cache, leaf, ComponentType.MAMBA))
|
||||||
# A failed load-back must roll back the pre-allocated mamba slot.
|
# A failed load-back must roll back the pre-allocated mamba slot.
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
req_to_token_pool.mamba_allocator.available_size(), avail_before
|
req_to_token_pool.mamba_allocator.available_size(), avail_before
|
||||||
)
|
)
|
||||||
@@ -5510,7 +5510,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
self._apply_match_to_req(req, match)
|
self._apply_match_to_req(req, match)
|
||||||
|
|
||||||
# Simulate a request without its own mamba slot so load-back allocates one.
|
# Simulate a request without its own mamba slot so load-back allocates one.
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
avail_before = req_to_token_pool.mamba_allocator.available_size()
|
avail_before = req_to_token_pool.mamba_allocator.available_size()
|
||||||
# H->D load fails after the mamba slot is pre-allocated -> must free it.
|
# H->D load fails after the mamba slot is pre-allocated -> must free it.
|
||||||
with mock.patch.object(cache.cache_controller, "load", return_value=None):
|
with mock.patch.object(cache.cache_controller, "load", return_value=None):
|
||||||
@@ -5524,7 +5524,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(len(new_indices), 0)
|
self.assertEqual(len(new_indices), 0)
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
req_to_token_pool.mamba_allocator.available_size(), avail_before
|
req_to_token_pool.mamba_allocator.available_size(), avail_before
|
||||||
)
|
)
|
||||||
@@ -5549,7 +5549,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
self._apply_match_to_req(req, match)
|
self._apply_match_to_req(req, match)
|
||||||
|
|
||||||
# Simulate a request without its own mamba slot so load-back allocates one.
|
# Simulate a request without its own mamba slot so load-back allocates one.
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
avail_before = req_to_token_pool.mamba_allocator.available_size()
|
avail_before = req_to_token_pool.mamba_allocator.available_size()
|
||||||
# No device room and eviction frees nothing -> load-back bails after the
|
# No device room and eviction frees nothing -> load-back bails after the
|
||||||
# mamba pre-alloc, which must still be freed.
|
# mamba pre-alloc, which must still be freed.
|
||||||
@@ -5571,7 +5571,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(len(new_indices), 0)
|
self.assertEqual(len(new_indices), 0)
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
req_to_token_pool.mamba_allocator.available_size(), avail_before
|
req_to_token_pool.mamba_allocator.available_size(), avail_before
|
||||||
)
|
)
|
||||||
@@ -5592,8 +5592,8 @@ class UnifiedRadixCacheSuite:
|
|||||||
|
|
||||||
# A request whose mamba slot was released: load_back's CoW arm allocates one.
|
# A request whose mamba slot was released: load_back's CoW arm allocates one.
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
|
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
|
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
|
||||||
|
|
||||||
# Impossible quota -> load_back aborts after building the transfers.
|
# Impossible quota -> load_back aborts after building the transfers.
|
||||||
@@ -5601,7 +5601,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
|
|
||||||
self.assertFalse(loaded)
|
self.assertFalse(loaded)
|
||||||
# the aborted call must return its slot and not leave req pointing at it
|
# the aborted call must return its slot and not leave req pointing at it
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
|
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
|
||||||
)
|
)
|
||||||
@@ -5621,8 +5621,8 @@ class UnifiedRadixCacheSuite:
|
|||||||
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
|
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
|
||||||
|
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
|
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
|
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
|
||||||
|
|
||||||
# cache_controller.load() failing (device alloc / transfer resolution)
|
# cache_controller.load() failing (device alloc / transfer resolution)
|
||||||
@@ -5631,7 +5631,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
loaded = cache.load_back(leaf, req=req)
|
loaded = cache.load_back(leaf, req=req)
|
||||||
|
|
||||||
self.assertFalse(loaded)
|
self.assertFalse(loaded)
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
|
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
|
||||||
)
|
)
|
||||||
@@ -5652,14 +5652,14 @@ class UnifiedRadixCacheSuite:
|
|||||||
|
|
||||||
# The request already owns its slot: an aborted load-back must not free it.
|
# The request already owns its slot: an aborted load-back must not free it.
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
preexisting_slot = req.mamba_pool_idx
|
preexisting_slot = req.kv.mamba_pool_idx
|
||||||
self.assertIsNotNone(preexisting_slot)
|
self.assertIsNotNone(preexisting_slot)
|
||||||
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
|
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
|
||||||
|
|
||||||
loaded = cache.load_back(leaf, mem_quota=-(10**9), req=req)
|
loaded = cache.load_back(leaf, mem_quota=-(10**9), req=req)
|
||||||
|
|
||||||
self.assertFalse(loaded)
|
self.assertFalse(loaded)
|
||||||
self.assertIs(req.mamba_pool_idx, preexisting_slot)
|
self.assertIs(req.kv.mamba_pool_idx, preexisting_slot)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
|
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
|
||||||
)
|
)
|
||||||
@@ -5679,15 +5679,15 @@ class UnifiedRadixCacheSuite:
|
|||||||
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
|
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
|
||||||
|
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
|
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
|
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
|
||||||
|
|
||||||
loaded = cache.load_back(leaf, req=req)
|
loaded = cache.load_back(leaf, req=req)
|
||||||
|
|
||||||
self.assertTrue(loaded)
|
self.assertTrue(loaded)
|
||||||
# the successful load must keep the freshly allocated slot published
|
# the successful load must keep the freshly allocated slot published
|
||||||
self.assertIsNotNone(req.mamba_pool_idx)
|
self.assertIsNotNone(req.kv.mamba_pool_idx)
|
||||||
self.assertIsNotNone(_device_value(cache, leaf, ComponentType.MAMBA))
|
self.assertIsNotNone(_device_value(cache, leaf, ComponentType.MAMBA))
|
||||||
# one slot restores the node's mamba value, one is the request's CoW slot
|
# one slot restores the node's mamba value, one is the request's CoW slot
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
@@ -5717,17 +5717,17 @@ class UnifiedRadixCacheSuite:
|
|||||||
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
|
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
|
||||||
|
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
|
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
|
|
||||||
loaded = cache.load_back(leaf, req=req)
|
loaded = cache.load_back(leaf, req=req)
|
||||||
self.assertTrue(loaded)
|
self.assertTrue(loaded)
|
||||||
self.assertIsNotNone(req.mamba_pool_idx)
|
self.assertIsNotNone(req.kv.mamba_pool_idx)
|
||||||
self._finish_pending_loads(cache)
|
self._finish_pending_loads(cache)
|
||||||
|
|
||||||
# The CoW slot must actually hold the backed-up mamba state, not merely exist.
|
# The CoW slot must actually hold the backed-up mamba state, not merely exist.
|
||||||
actual_temporal, actual_conv = self._snapshot_mamba_state(
|
actual_temporal, actual_conv = self._snapshot_mamba_state(
|
||||||
req_to_token_pool, req.mamba_pool_idx.unsqueeze(0)
|
req_to_token_pool, req.kv.mamba_pool_idx.unsqueeze(0)
|
||||||
)
|
)
|
||||||
self.assertTrue(torch.equal(actual_temporal, expected_temporal))
|
self.assertTrue(torch.equal(actual_temporal, expected_temporal))
|
||||||
self.assertEqual(len(actual_conv), len(expected_conv))
|
self.assertEqual(len(actual_conv), len(expected_conv))
|
||||||
@@ -5747,12 +5747,12 @@ class UnifiedRadixCacheSuite:
|
|||||||
|
|
||||||
self._backup_node(cache, leaf)
|
self._backup_node(cache, leaf)
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
|
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
|
|
||||||
# device value still present -> nothing to prepare even though host-backed
|
# device value still present -> nothing to prepare even though host-backed
|
||||||
self.assertIsNone(comp.prepare_load_back(leaf, req=req).allocated_mamba_slot)
|
self.assertIsNone(comp.prepare_load_back(leaf, req=req).allocated_mamba_slot)
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
|
|
||||||
cache.evict(EvictParams(num_tokens=_node_key_length(cache, leaf)))
|
cache.evict(EvictParams(num_tokens=_node_key_length(cache, leaf)))
|
||||||
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
|
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
|
||||||
@@ -5769,16 +5769,16 @@ class UnifiedRadixCacheSuite:
|
|||||||
# fresh request + host-only mamba -> allocates and publishes onto req
|
# fresh request + host-only mamba -> allocates and publishes onto req
|
||||||
prep = comp.prepare_load_back(leaf, req=req)
|
prep = comp.prepare_load_back(leaf, req=req)
|
||||||
self.assertIsNotNone(prep.allocated_mamba_slot)
|
self.assertIsNotNone(prep.allocated_mamba_slot)
|
||||||
self.assertEqual(int(req.mamba_pool_idx), int(prep.allocated_mamba_slot[0]))
|
self.assertEqual(int(req.kv.mamba_pool_idx), int(prep.allocated_mamba_slot[0]))
|
||||||
|
|
||||||
# node without host-backed mamba -> nothing to prepare
|
# node without host-backed mamba -> nothing to prepare
|
||||||
req2 = self._make_req(req_to_token_pool)
|
req2 = self._make_req(req_to_token_pool)
|
||||||
req_to_token_pool.mamba_allocator.free(req2.mamba_pool_idx.unsqueeze(0))
|
req_to_token_pool.mamba_allocator.free(req2.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req2.mamba_pool_idx = None
|
req2.kv.mamba_pool_idx = None
|
||||||
root = cache.root_node_handle()
|
root = cache.root_node_handle()
|
||||||
self.assertIsNone(_host_value(cache, root, ComponentType.MAMBA))
|
self.assertIsNone(_host_value(cache, root, ComponentType.MAMBA))
|
||||||
self.assertIsNone(comp.prepare_load_back(root, req=req2).allocated_mamba_slot)
|
self.assertIsNone(comp.prepare_load_back(root, req=req2).allocated_mamba_slot)
|
||||||
self.assertIsNone(req2.mamba_pool_idx)
|
self.assertIsNone(req2.kv.mamba_pool_idx)
|
||||||
|
|
||||||
def test_prepare_load_back_skips_device_present_node(self):
|
def test_prepare_load_back_skips_device_present_node(self):
|
||||||
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
|
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
|
||||||
@@ -5796,12 +5796,12 @@ class UnifiedRadixCacheSuite:
|
|||||||
self.assertIsNotNone(_host_value(cache, leaf, ComponentType.MAMBA))
|
self.assertIsNotNone(_host_value(cache, leaf, ComponentType.MAMBA))
|
||||||
|
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
|
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
|
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
|
||||||
|
|
||||||
self.assertIsNone(comp.prepare_load_back(leaf, req=req).allocated_mamba_slot)
|
self.assertIsNone(comp.prepare_load_back(leaf, req=req).allocated_mamba_slot)
|
||||||
self.assertIsNone(req.mamba_pool_idx)
|
self.assertIsNone(req.kv.mamba_pool_idx)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
|
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
|
||||||
)
|
)
|
||||||
@@ -5820,8 +5820,8 @@ class UnifiedRadixCacheSuite:
|
|||||||
cache.evict(EvictParams(num_tokens=_node_key_length(cache, leaf)))
|
cache.evict(EvictParams(num_tokens=_node_key_length(cache, leaf)))
|
||||||
|
|
||||||
req = self._make_req(req_to_token_pool)
|
req = self._make_req(req_to_token_pool)
|
||||||
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
|
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
|
||||||
req.mamba_pool_idx = None
|
req.kv.mamba_pool_idx = None
|
||||||
retry_slot = req_to_token_pool.mamba_allocator.alloc(1)
|
retry_slot = req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
|
|
||||||
# first alloc fails -> prepare must evict a mamba slot and retry
|
# first alloc fails -> prepare must evict a mamba slot and retry
|
||||||
@@ -5838,7 +5838,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
prep = comp.prepare_load_back(leaf, req=req)
|
prep = comp.prepare_load_back(leaf, req=req)
|
||||||
evict_for_alloc.assert_called_once_with(EvictParams(num_tokens=0, mamba_num=1))
|
evict_for_alloc.assert_called_once_with(EvictParams(num_tokens=0, mamba_num=1))
|
||||||
self.assertIs(prep.allocated_mamba_slot, retry_slot)
|
self.assertIs(prep.allocated_mamba_slot, retry_slot)
|
||||||
self.assertEqual(int(req.mamba_pool_idx), int(retry_slot[0]))
|
self.assertEqual(int(req.kv.mamba_pool_idx), int(retry_slot[0]))
|
||||||
|
|
||||||
def test_hicache_swa_load_back_min_suffix(self):
|
def test_hicache_swa_load_back_min_suffix(self):
|
||||||
"""LOAD_BACK collects only the suffix nodes needed to cover sliding_window_size."""
|
"""LOAD_BACK collects only the suffix nodes needed to cover sliding_window_size."""
|
||||||
@@ -6694,7 +6694,7 @@ class TestUnifiedMambaLRUMatchRefresh(CustomTestCase):
|
|||||||
InsertParams(
|
InsertParams(
|
||||||
key=RadixKey(array("q", tokens)),
|
key=RadixKey(array("q", tokens)),
|
||||||
value=value[: len(tokens)],
|
value=value[: len(tokens)],
|
||||||
mamba_value=req.mamba_pool_idx.unsqueeze(0),
|
mamba_value=req.kv.mamba_pool_idx.unsqueeze(0),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -6785,7 +6785,7 @@ class TestUnifiedRadixCacheInt8MambaCheckpoint(CustomTestCase):
|
|||||||
req.kv.cache_protected_len = 0
|
req.kv.cache_protected_len = 0
|
||||||
req.swa_uuid_for_lock = None
|
req.swa_uuid_for_lock = None
|
||||||
req.extra_key = None
|
req.extra_key = None
|
||||||
req.mamba_last_track_seqlen = len(tokens)
|
req.kv.mamba_last_track_seqlen = len(tokens)
|
||||||
return req
|
return req
|
||||||
|
|
||||||
def _cache_finished(self, cache, allocator, req_to_token_pool, tokens):
|
def _cache_finished(self, cache, allocator, req_to_token_pool, tokens):
|
||||||
|
|||||||
Reference in New Issue
Block a user