[mem_cache][4/N] refactor: extract MambaTokenToKVPoolAllocator into allocator/ (#27256)
This commit is contained in:
@@ -241,7 +241,7 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool):
|
|||||||
|
|
||||||
def clear(self):
|
def clear(self):
|
||||||
self.free_slots = list(range(1, self._alloc_size))
|
self.free_slots = list(range(1, self._alloc_size))
|
||||||
self.mamba_pool.clear()
|
self.mamba_allocator.clear()
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -2417,10 +2417,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
if envs.SGLANG_TEST_MAMBA_LAZY_ALLOC_FAIL.get():
|
if envs.SGLANG_TEST_MAMBA_LAZY_ALLOC_FAIL.get():
|
||||||
new_slot = None
|
new_slot = None
|
||||||
else:
|
else:
|
||||||
new_slot = pool.mamba_pool.alloc(1)
|
new_slot = pool.mamba_allocator.alloc(1)
|
||||||
if new_slot is None:
|
if new_slot is None:
|
||||||
self.tree_cache.evict(EvictParams(num_tokens=0, mamba_num=1))
|
self.tree_cache.evict(EvictParams(num_tokens=0, mamba_num=1))
|
||||||
new_slot = pool.mamba_pool.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.mamba_next_track_idx = other_idx
|
||||||
|
|||||||
@@ -2662,9 +2662,9 @@ class Scheduler(
|
|||||||
self.running_batch.reqs,
|
self.running_batch.reqs,
|
||||||
)
|
)
|
||||||
|
|
||||||
mamba_pool = getattr(self.req_to_token_pool, "mamba_pool", None)
|
mamba_allocator = getattr(self.req_to_token_pool, "mamba_allocator", None)
|
||||||
if mamba_pool is not None:
|
if mamba_allocator is not None:
|
||||||
mamba_pool.alloc_group_begin(len(self.waiting_queue))
|
mamba_allocator.alloc_group_begin(len(self.waiting_queue))
|
||||||
# Get requests from the waiting queue to a new prefill batch
|
# Get requests from the waiting queue to a new prefill batch
|
||||||
for req in self.waiting_queue:
|
for req in self.waiting_queue:
|
||||||
if self.enable_lora and not self._can_schedule_lora_req(req, running_loras):
|
if self.enable_lora and not self._can_schedule_lora_req(req, running_loras):
|
||||||
@@ -2725,14 +2725,14 @@ class Scheduler(
|
|||||||
and req.mamba_pool_idx is not None
|
and req.mamba_pool_idx is not None
|
||||||
and not getattr(req, "session", None)
|
and not getattr(req, "session", None)
|
||||||
):
|
):
|
||||||
self.tree_cache.req_to_token_pool.mamba_pool.free(
|
self.tree_cache.req_to_token_pool.mamba_allocator.free(
|
||||||
req.mamba_pool_idx.unsqueeze(-1)
|
req.mamba_pool_idx.unsqueeze(-1)
|
||||||
)
|
)
|
||||||
req.mamba_pool_idx = None
|
req.mamba_pool_idx = None
|
||||||
break
|
break
|
||||||
|
|
||||||
if mamba_pool is not None:
|
if mamba_allocator is not None:
|
||||||
mamba_pool.alloc_group_end()
|
mamba_allocator.alloc_group_end()
|
||||||
|
|
||||||
# Update waiting queue
|
# Update waiting queue
|
||||||
can_run_list: List[Req] = adder.can_run_list
|
can_run_list: List[Req] = adder.can_run_list
|
||||||
@@ -2867,9 +2867,13 @@ class Scheduler(
|
|||||||
):
|
):
|
||||||
old_available_tokens = self.token_to_kv_pool_allocator.available_size()
|
old_available_tokens = self.token_to_kv_pool_allocator.available_size()
|
||||||
old_ratio = self.new_token_ratio_tracker.current
|
old_ratio = self.new_token_ratio_tracker.current
|
||||||
mamba_pool = getattr(self.tree_cache.req_to_token_pool, "mamba_pool", None)
|
mamba_allocator = getattr(
|
||||||
|
self.tree_cache.req_to_token_pool, "mamba_allocator", None
|
||||||
|
)
|
||||||
old_mamba_available = (
|
old_mamba_available = (
|
||||||
mamba_pool.available_size() if mamba_pool is not None else None
|
mamba_allocator.available_size()
|
||||||
|
if mamba_allocator is not None
|
||||||
|
else None
|
||||||
)
|
)
|
||||||
retracted_reqs, new_token_ratio, reqs_to_abort = batch.retract_decode(
|
retracted_reqs, new_token_ratio, reqs_to_abort = batch.retract_decode(
|
||||||
self.server_args
|
self.server_args
|
||||||
@@ -2877,8 +2881,8 @@ class Scheduler(
|
|||||||
new_available_tokens = self.token_to_kv_pool_allocator.available_size()
|
new_available_tokens = self.token_to_kv_pool_allocator.available_size()
|
||||||
new_token_gained = new_available_tokens - old_available_tokens
|
new_token_gained = new_available_tokens - old_available_tokens
|
||||||
mamba_num_gained = (
|
mamba_num_gained = (
|
||||||
mamba_pool.available_size() - old_mamba_available
|
mamba_allocator.available_size() - old_mamba_available
|
||||||
if mamba_pool is not None
|
if mamba_allocator is not None
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -938,7 +938,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
other_val = req.mamba_ping_pong_track_buffer[other_idx].item()
|
other_val = req.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_pool.free(
|
pool.mamba_allocator.free(
|
||||||
req.mamba_ping_pong_track_buffer[other_idx].unsqueeze(0)
|
req.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)
|
||||||
|
|||||||
@@ -135,13 +135,12 @@ class SchedulerInvariantChecker:
|
|||||||
leaked_full_pages = (
|
leaked_full_pages = (
|
||||||
expected_full_pages - free_full_pages - cached_full_pages
|
expected_full_pages - free_full_pages - cached_full_pages
|
||||||
)
|
)
|
||||||
free_mamba_pages = set(
|
mamba_allocator = self.req_to_token_pool.mamba_allocator
|
||||||
self.req_to_token_pool.mamba_pool.free_slots.tolist()
|
free_mamba_pages = set(mamba_allocator.free_slots.tolist())
|
||||||
)
|
|
||||||
cached_mamba_pages = set(
|
cached_mamba_pages = set(
|
||||||
self.tree_cache.all_mamba_values_flatten().tolist()
|
self.tree_cache.all_mamba_values_flatten().tolist()
|
||||||
)
|
)
|
||||||
expected_mamba_pages = set(range(self.req_to_token_pool.mamba_pool.size))
|
expected_mamba_pages = set(range(1, mamba_allocator.size + 1))
|
||||||
leaked_mamba_pages = (
|
leaked_mamba_pages = (
|
||||||
expected_mamba_pages - free_mamba_pages - cached_mamba_pages
|
expected_mamba_pages - free_mamba_pages - cached_mamba_pages
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -246,7 +246,7 @@ class SchedulerPoolStatsObserver:
|
|||||||
full_evictable_size = (
|
full_evictable_size = (
|
||||||
self.tree_cache.full_evictable_size() if is_mamba_radix_cache else 0
|
self.tree_cache.full_evictable_size() if is_mamba_radix_cache else 0
|
||||||
)
|
)
|
||||||
mamba_available_size = self.req_to_token_pool.mamba_pool.available_size()
|
mamba_available_size = self.req_to_token_pool.mamba_allocator.available_size()
|
||||||
mamba_evictable_size = (
|
mamba_evictable_size = (
|
||||||
self.tree_cache.mamba_evictable_size() if is_mamba_radix_cache else 0
|
self.tree_cache.mamba_evictable_size() if is_mamba_radix_cache else 0
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,90 @@
|
|||||||
|
"""
|
||||||
|
Copyright 2026 SGLang Team
|
||||||
|
Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
you may not use this file except in compliance with the License.
|
||||||
|
You may obtain a copy of the License at
|
||||||
|
|
||||||
|
http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
|
||||||
|
Unless required by applicable law or agreed to in writing, software
|
||||||
|
distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
See the License for the specific language governing permissions and
|
||||||
|
limitations under the License.
|
||||||
|
|
||||||
|
Slot allocator for the Mamba state pool.
|
||||||
|
|
||||||
|
Mamba caches one whole state tensor per request, so the allocator hands out
|
||||||
|
fixed-size slots (1 per request) rather than paged token KV indices. The
|
||||||
|
underlying tensor storage lives in ``MambaPool``; this class owns only the
|
||||||
|
free-slot bookkeeping.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Iterator, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
class MambaSlotAllocator:
|
||||||
|
"""Manages the free-list of Mamba pool slot indices.
|
||||||
|
|
||||||
|
Unlike ``BaseTokenToKVPoolAllocator`` which is designed for per-token KV
|
||||||
|
pages, Mamba slots are request-level (typically 1 slot per request).
|
||||||
|
We keep the interface minimal and do NOT inherit the KV base class.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, size: int, device: str):
|
||||||
|
self.size = size
|
||||||
|
self.device = device
|
||||||
|
# Active preallocated batch for `alloc_group_begin` / `alloc_group_end`.
|
||||||
|
# When non-None, `alloc(1)` consumes the next slot from this iterator
|
||||||
|
# instead of calling `_do_alloc(1)` per request. Reset to None outside
|
||||||
|
# a group window so `alloc` falls through to the per-call path.
|
||||||
|
self._alloc_iter: Optional[Iterator] = None
|
||||||
|
self.clear()
|
||||||
|
|
||||||
|
def available_size(self) -> int:
|
||||||
|
return len(self.free_slots)
|
||||||
|
|
||||||
|
def alloc_group_begin(self, num_reqs: int):
|
||||||
|
"""Pre-allocate a batch of slots for match_prefix to amortize overhead."""
|
||||||
|
self._alloc_iter = None
|
||||||
|
if num_reqs > 0:
|
||||||
|
result = self._do_alloc(num_reqs)
|
||||||
|
if result is not None:
|
||||||
|
self._alloc_iter = iter(result.split(1))
|
||||||
|
|
||||||
|
def alloc_group_end(self):
|
||||||
|
"""Return any unused pre-allocated slots from the current group."""
|
||||||
|
if self._alloc_iter is not None:
|
||||||
|
remaining = list(self._alloc_iter)
|
||||||
|
if remaining:
|
||||||
|
self.free(torch.cat(remaining))
|
||||||
|
self._alloc_iter = None
|
||||||
|
|
||||||
|
def alloc(self, need_size: int) -> Optional[torch.Tensor]:
|
||||||
|
if self._alloc_iter is not None and need_size == 1:
|
||||||
|
slot = next(self._alloc_iter, None)
|
||||||
|
if slot is not None:
|
||||||
|
return slot
|
||||||
|
return self._do_alloc(need_size)
|
||||||
|
|
||||||
|
def _do_alloc(self, need_size: int) -> Optional[torch.Tensor]:
|
||||||
|
if need_size > len(self.free_slots):
|
||||||
|
return None
|
||||||
|
select_index = self.free_slots[:need_size]
|
||||||
|
self.free_slots = self.free_slots[need_size:]
|
||||||
|
return select_index
|
||||||
|
|
||||||
|
def free(self, free_index: torch.Tensor):
|
||||||
|
if free_index.numel() == 0:
|
||||||
|
return
|
||||||
|
self.free_slots = torch.cat((self.free_slots, free_index))
|
||||||
|
|
||||||
|
def clear(self):
|
||||||
|
# Slot 0 is reserved as a dummy write target for padded tokens.
|
||||||
|
self.free_slots = torch.arange(
|
||||||
|
1, self.size + 1, dtype=torch.int64, device=self.device
|
||||||
|
)
|
||||||
@@ -314,7 +314,7 @@ def alloc_req_slots(
|
|||||||
"""Allocate request slots from the pool."""
|
"""Allocate request slots from the pool."""
|
||||||
num_reqs = len(reqs)
|
num_reqs = len(reqs)
|
||||||
if isinstance(req_to_token_pool, HybridReqToTokenPool):
|
if isinstance(req_to_token_pool, HybridReqToTokenPool):
|
||||||
mamba_available_size = req_to_token_pool.mamba_pool.available_size()
|
mamba_available_size = req_to_token_pool.mamba_allocator.available_size()
|
||||||
if tree_cache.supports_mamba():
|
if tree_cache.supports_mamba():
|
||||||
factor = (
|
factor = (
|
||||||
MAMBA_STATE_PER_REQ_PREFIX_CACHE_LAZY
|
MAMBA_STATE_PER_REQ_PREFIX_CACHE_LAZY
|
||||||
@@ -486,7 +486,7 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
|
|||||||
), "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.mamba_pool_idx is not None:
|
||||||
tree_cache.req_to_token_pool.mamba_pool.free(
|
tree_cache.req_to_token_pool.mamba_allocator.free(
|
||||||
req.mamba_pool_idx.unsqueeze(-1)
|
req.mamba_pool_idx.unsqueeze(-1)
|
||||||
)
|
)
|
||||||
req.mamba_pool_idx = None
|
req.mamba_pool_idx = None
|
||||||
|
|||||||
@@ -508,7 +508,7 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
if node.mamba_value is None:
|
if node.mamba_value is None:
|
||||||
return 0
|
return 0
|
||||||
mamba_num = len(node.mamba_value)
|
mamba_num = len(node.mamba_value)
|
||||||
self.req_to_token_pool.mamba_pool.free(node.mamba_value)
|
self.req_to_token_pool.mamba_allocator.free(node.mamba_value)
|
||||||
if node.mamba_lock_ref > 0:
|
if node.mamba_lock_ref > 0:
|
||||||
self.mamba_protected_size_ -= mamba_num
|
self.mamba_protected_size_ -= mamba_num
|
||||||
node.mamba_lock_ref = 0
|
node.mamba_lock_ref = 0
|
||||||
@@ -797,7 +797,7 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
# Internal: free device mamba only, KV stays on device (tombstone)
|
# Internal: free device mamba only, KV stays on device (tombstone)
|
||||||
x_next = self.mamba_lru_list.get_prev_no_lock(x)
|
x_next = self.mamba_lru_list.get_prev_no_lock(x)
|
||||||
mamba_num_evicted += len(x.mamba_value)
|
mamba_num_evicted += len(x.mamba_value)
|
||||||
self.req_to_token_pool.mamba_pool.free(x.mamba_value)
|
self.req_to_token_pool.mamba_allocator.free(x.mamba_value)
|
||||||
self.mamba_lru_list.remove_node(x)
|
self.mamba_lru_list.remove_node(x)
|
||||||
self._tombstone_internal_node(x)
|
self._tombstone_internal_node(x)
|
||||||
else:
|
else:
|
||||||
@@ -1047,7 +1047,7 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
if cow_mamba and mamba_node.mamba_value is not None:
|
if cow_mamba and mamba_node.mamba_value is not None:
|
||||||
if req.mamba_pool_idx is None:
|
if req.mamba_pool_idx is None:
|
||||||
dst_index = self._alloc_with_evict(
|
dst_index = self._alloc_with_evict(
|
||||||
self.req_to_token_pool.mamba_pool,
|
self.req_to_token_pool.mamba_allocator,
|
||||||
1,
|
1,
|
||||||
self.evict_mamba,
|
self.evict_mamba,
|
||||||
lock_node=mamba_node,
|
lock_node=mamba_node,
|
||||||
@@ -2085,7 +2085,7 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
):
|
):
|
||||||
if req.mamba_pool_idx is None:
|
if req.mamba_pool_idx is None:
|
||||||
req.mamba_pool_idx = self._alloc_with_evict(
|
req.mamba_pool_idx = self._alloc_with_evict(
|
||||||
self.req_to_token_pool.mamba_pool,
|
self.req_to_token_pool.mamba_allocator,
|
||||||
len(last_hit_node.mamba_host_value),
|
len(last_hit_node.mamba_host_value),
|
||||||
self.evict_mamba,
|
self.evict_mamba,
|
||||||
lock_node=last_hit_node,
|
lock_node=last_hit_node,
|
||||||
|
|||||||
@@ -515,6 +515,7 @@ def build_hybrid_mamba_stack(
|
|||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping)
|
transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping)
|
||||||
|
mamba_allocator = params.req_to_token_pool.mamba_allocator
|
||||||
kv_host_pool = build_kv_host_pool(
|
kv_host_pool = build_kv_host_pool(
|
||||||
kv_pool=kv_pool,
|
kv_pool=kv_pool,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
@@ -545,6 +546,8 @@ def build_hybrid_mamba_stack(
|
|||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num,
|
||||||
host_evict_fn=host_mamba_evict_fn,
|
host_evict_fn=host_mamba_evict_fn,
|
||||||
device_evict_fn=device_mamba_evict_fn,
|
device_evict_fn=device_mamba_evict_fn,
|
||||||
|
device_alloc_fn=mamba_allocator.alloc,
|
||||||
|
device_free_fn=mamba_allocator.free,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
host_pool_group = HostPoolGroup(entries)
|
host_pool_group = HostPoolGroup(entries)
|
||||||
|
|||||||
@@ -671,7 +671,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
)
|
)
|
||||||
new_prefix_len, mamba_exist = result.prefix_len, result.mamba_exist
|
new_prefix_len, mamba_exist = result.prefix_len, result.mamba_exist
|
||||||
if mamba_exist:
|
if mamba_exist:
|
||||||
self.req_to_token_pool.mamba_pool.free(mamba_value_donated)
|
self.req_to_token_pool.mamba_allocator.free(mamba_value_donated)
|
||||||
|
|
||||||
# The prefix indices could be updated, reuse it
|
# The prefix indices could be updated, reuse it
|
||||||
match_result = self.match_prefix(
|
match_result = self.match_prefix(
|
||||||
@@ -729,7 +729,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
self._record_remove_event(x)
|
self._record_remove_event(x)
|
||||||
self.token_to_kv_pool_allocator.free(x.value)
|
self.token_to_kv_pool_allocator.free(x.value)
|
||||||
full_num_evicted = len(x.value)
|
full_num_evicted = len(x.value)
|
||||||
self.req_to_token_pool.mamba_pool.free(x.mamba_value)
|
self.req_to_token_pool.mamba_allocator.free(x.mamba_value)
|
||||||
mamba_num_evicted = len(x.mamba_value)
|
mamba_num_evicted = len(x.mamba_value)
|
||||||
|
|
||||||
# 2. get the next node, update the lru lists
|
# 2. get the next node, update the lru lists
|
||||||
@@ -782,7 +782,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
|
|
||||||
if len(x.children) > 0:
|
if len(x.children) > 0:
|
||||||
# 1. an internal node, free mamba tokens.
|
# 1. an internal node, free mamba tokens.
|
||||||
self.req_to_token_pool.mamba_pool.free(x.mamba_value)
|
self.req_to_token_pool.mamba_allocator.free(x.mamba_value)
|
||||||
mamba_num_evicted += len(x.mamba_value)
|
mamba_num_evicted += len(x.mamba_value)
|
||||||
|
|
||||||
# 2. get the next node, update the lru lists
|
# 2. get the next node, update the lru lists
|
||||||
@@ -947,10 +947,10 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
|
|
||||||
def _alloc_mamba_slot(self) -> torch.Tensor:
|
def _alloc_mamba_slot(self) -> torch.Tensor:
|
||||||
"""Allocate one mamba pool slot, evicting if necessary."""
|
"""Allocate one mamba pool slot, evicting if necessary."""
|
||||||
slot = self.req_to_token_pool.mamba_pool.alloc(1)
|
slot = self.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
if slot is None:
|
if slot is None:
|
||||||
self.evict(EvictParams(num_tokens=0, mamba_num=1))
|
self.evict(EvictParams(num_tokens=0, mamba_num=1))
|
||||||
slot = self.req_to_token_pool.mamba_pool.alloc(1)
|
slot = self.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
assert slot is not None, "Can not alloc mamba cache"
|
assert slot is not None, "Can not alloc mamba cache"
|
||||||
return slot
|
return slot
|
||||||
|
|
||||||
@@ -1046,11 +1046,11 @@ class MambaRadixCache(KVCacheEventMixin, 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 req.mamba_pool_idx is None:
|
||||||
dst_index = self.req_to_token_pool.mamba_pool.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)
|
||||||
self.evict(EvictParams(num_tokens=0, mamba_num=1))
|
self.evict(EvictParams(num_tokens=0, mamba_num=1))
|
||||||
dst_index = self.req_to_token_pool.mamba_pool.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.mamba_pool_idx = dst_index[0]
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ import dataclasses
|
|||||||
import logging
|
import logging
|
||||||
from contextlib import contextmanager, nullcontext
|
from contextlib import contextmanager, nullcontext
|
||||||
from dataclasses import dataclass, fields
|
from dataclasses import dataclass, fields
|
||||||
from typing import TYPE_CHECKING, Any, Iterator, List, Optional, Tuple, Union
|
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
@@ -44,6 +44,7 @@ from sglang.srt.layers.attention.dsa.quant_k_cache import (
|
|||||||
from sglang.srt.layers.attention.dsa.utils import aiter_can_use_preshuffle_paged_mqa
|
from sglang.srt.layers.attention.dsa.utils import aiter_can_use_preshuffle_paged_mqa
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz
|
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
|
from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator
|
||||||
from sglang.srt.mem_cache.triton_ops.cache_move import (
|
from sglang.srt.mem_cache.triton_ops.cache_move import (
|
||||||
copy_all_layer_kv_cache_tiled,
|
copy_all_layer_kv_cache_tiled,
|
||||||
)
|
)
|
||||||
@@ -360,19 +361,9 @@ class MambaPool:
|
|||||||
f"conv_state size: {get_tensor_size_bytes(conv_state) / GB:.2f}GB, "
|
f"conv_state size: {get_tensor_size_bytes(conv_state) / GB:.2f}GB, "
|
||||||
f"ssm_state size: {get_tensor_size_bytes(temporal_state) / GB:.2f}GB "
|
f"ssm_state size: {get_tensor_size_bytes(temporal_state) / GB:.2f}GB "
|
||||||
)
|
)
|
||||||
# The padded slot 0 is used for writing dummy outputs from padded tokens.
|
|
||||||
self.free_slots = torch.arange(
|
|
||||||
1, self.size + 1, dtype=torch.int64, device=self.device
|
|
||||||
)
|
|
||||||
self.mem_usage = self.mamba_cache.mem_usage_bytes() / GB
|
self.mem_usage = self.mamba_cache.mem_usage_bytes() / GB
|
||||||
self.num_mamba_layers = num_mamba_layers
|
self.num_mamba_layers = num_mamba_layers
|
||||||
|
|
||||||
# Active preallocated batch for `alloc_group_begin` / `alloc_group_end`.
|
|
||||||
# When non-None, `alloc(1)` consumes the next slot from this iterator
|
|
||||||
# instead of calling `_do_alloc(1)` per request. Reset to None outside
|
|
||||||
# a group window so `alloc` falls through to the per-call path.
|
|
||||||
self._alloc_iter: Optional[Iterator] = None
|
|
||||||
|
|
||||||
def get_speculative_mamba2_params_all_layers(self) -> SpeculativeState:
|
def get_speculative_mamba2_params_all_layers(self) -> SpeculativeState:
|
||||||
assert isinstance(self.mamba_cache, self.SpeculativeState)
|
assert isinstance(self.mamba_cache, self.SpeculativeState)
|
||||||
return self.mamba_cache
|
return self.mamba_cache
|
||||||
@@ -380,39 +371,6 @@ class MambaPool:
|
|||||||
def mamba2_layer_cache(self, layer_id: int):
|
def mamba2_layer_cache(self, layer_id: int):
|
||||||
return self.mamba_cache.at_layer_idx(layer_id)
|
return self.mamba_cache.at_layer_idx(layer_id)
|
||||||
|
|
||||||
def available_size(self):
|
|
||||||
return len(self.free_slots)
|
|
||||||
|
|
||||||
# -- Batched alloc for match_prefix --
|
|
||||||
def alloc_group_begin(self, num_reqs: int):
|
|
||||||
self._alloc_iter = None
|
|
||||||
if num_reqs > 0:
|
|
||||||
result = self._do_alloc(num_reqs)
|
|
||||||
if result is not None:
|
|
||||||
self._alloc_iter = iter(result.split(1))
|
|
||||||
|
|
||||||
def alloc_group_end(self):
|
|
||||||
if self._alloc_iter is not None:
|
|
||||||
remaining = list(self._alloc_iter)
|
|
||||||
if remaining:
|
|
||||||
self.free(torch.cat(remaining))
|
|
||||||
self._alloc_iter = None
|
|
||||||
|
|
||||||
def alloc(self, need_size: int) -> Optional[torch.Tensor]:
|
|
||||||
if self._alloc_iter is not None and need_size == 1:
|
|
||||||
slot = next(self._alloc_iter, None)
|
|
||||||
if slot is not None:
|
|
||||||
return slot
|
|
||||||
return self._do_alloc(need_size)
|
|
||||||
|
|
||||||
def _do_alloc(self, need_size: int) -> Optional[torch.Tensor]:
|
|
||||||
if need_size > len(self.free_slots):
|
|
||||||
return None
|
|
||||||
|
|
||||||
select_index = self.free_slots[:need_size]
|
|
||||||
self.free_slots = self.free_slots[need_size:]
|
|
||||||
return select_index
|
|
||||||
|
|
||||||
def clear_slots(self, indices: torch.Tensor):
|
def clear_slots(self, indices: torch.Tensor):
|
||||||
"""Zero out mamba state at the given pool indices. Must run on forward stream."""
|
"""Zero out mamba state at the given pool indices. Must run on forward stream."""
|
||||||
need_size = len(indices)
|
need_size = len(indices)
|
||||||
@@ -428,16 +386,6 @@ class MambaPool:
|
|||||||
)
|
)
|
||||||
t[:, indices] = z
|
t[:, indices] = z
|
||||||
|
|
||||||
def free(self, free_index: torch.Tensor):
|
|
||||||
if free_index.numel() == 0:
|
|
||||||
return
|
|
||||||
self.free_slots = torch.cat((self.free_slots, free_index))
|
|
||||||
|
|
||||||
def clear(self):
|
|
||||||
self.free_slots = torch.arange(
|
|
||||||
1, self.size + 1, dtype=torch.int64, device=self.device
|
|
||||||
)
|
|
||||||
|
|
||||||
def copy_from(self, src_indices: torch.Tensor, dst_indices: torch.Tensor):
|
def copy_from(self, src_indices: torch.Tensor, dst_indices: torch.Tensor):
|
||||||
for i in range(len(self.mamba_cache.conv)):
|
for i in range(len(self.mamba_cache.conv)):
|
||||||
self.mamba_cache.conv[i][:, dst_indices] = self.mamba_cache.conv[i][
|
self.mamba_cache.conv[i][:, dst_indices] = self.mamba_cache.conv[i][
|
||||||
@@ -588,6 +536,10 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
enable_memory_saver=self.enable_memory_saver,
|
enable_memory_saver=self.enable_memory_saver,
|
||||||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||||||
)
|
)
|
||||||
|
self.mamba_allocator = MambaSlotAllocator(
|
||||||
|
size=mamba_size,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
self.mamba_map = {layer_id: i for i, layer_id in enumerate(mamba_layer_ids)}
|
self.mamba_map = {layer_id: i for i, layer_id in enumerate(mamba_layer_ids)}
|
||||||
|
|
||||||
self.device = device
|
self.device = device
|
||||||
@@ -622,10 +574,10 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
if req.mamba_pool_idx is not None: # for radix cache / continuing chunked
|
if req.mamba_pool_idx is not None: # for radix cache / continuing chunked
|
||||||
pass
|
pass
|
||||||
else:
|
else:
|
||||||
mid = self.mamba_pool.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_pool.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.mamba_pool_idx = mid[0]
|
||||||
req.mamba_needs_clear = True
|
req.mamba_needs_clear = True
|
||||||
mamba_indices.append(req.mamba_pool_idx)
|
mamba_indices.append(req.mamba_pool_idx)
|
||||||
@@ -694,7 +646,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
if self.enable_mamba_extra_buffer_lazy
|
if self.enable_mamba_extra_buffer_lazy
|
||||||
else self.mamba_ping_pong_track_buffer_size
|
else self.mamba_ping_pong_track_buffer_size
|
||||||
)
|
)
|
||||||
slots = self.mamba_pool.alloc(n)
|
slots = self.mamba_allocator.alloc(n)
|
||||||
assert slots is not None, (
|
assert slots is not None, (
|
||||||
"Not enough space for mamba ping pong idx, "
|
"Not enough space for mamba ping pong idx, "
|
||||||
"try to increase --mamba-full-memory-ratio."
|
"try to increase --mamba-full-memory-ratio."
|
||||||
@@ -749,7 +701,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
):
|
):
|
||||||
mamba_index = req.mamba_pool_idx
|
mamba_index = req.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_pool.free(mamba_index.unsqueeze(0))
|
self.mamba_allocator.free(mamba_index.unsqueeze(0))
|
||||||
req.mamba_pool_idx = None
|
req.mamba_pool_idx = None
|
||||||
|
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.enable_mamba_extra_buffer:
|
||||||
@@ -789,7 +741,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
mamba_ping_pong_track_buffer_to_free != -1
|
mamba_ping_pong_track_buffer_to_free != -1
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.mamba_pool.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.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
|
||||||
@@ -800,7 +752,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
def clear(self):
|
def clear(self):
|
||||||
logger.info("Reset HybridReqToTokenPool")
|
logger.info("Reset HybridReqToTokenPool")
|
||||||
super().clear()
|
super().clear()
|
||||||
self.mamba_pool.clear()
|
self.mamba_allocator.clear()
|
||||||
self.req_index_to_mamba_index_mapping.zero_()
|
self.req_index_to_mamba_index_mapping.zero_()
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.enable_mamba_extra_buffer:
|
||||||
self.req_index_to_mamba_ping_pong_track_buffer_mapping.zero_()
|
self.req_index_to_mamba_ping_pong_track_buffer_mapping.zero_()
|
||||||
|
|||||||
@@ -2574,8 +2574,8 @@ class PoolEntry:
|
|||||||
device_evict_fn: Optional[Callable] = None
|
device_evict_fn: Optional[Callable] = None
|
||||||
# Optional alloc/free overrides for the device side, used by
|
# Optional alloc/free overrides for the device side, used by
|
||||||
# _resolve_pool_transfers_allocation. Set when entry.device_pool is the
|
# _resolve_pool_transfers_allocation. Set when entry.device_pool is the
|
||||||
# raw KV pool (layout) rather than an allocator (e.g. SWA, where alloc
|
# raw KV/state pool (layout) rather than an allocator (e.g. SWA/Mamba,
|
||||||
# lives on a separate sub-allocator inside SWATokenToKVPoolAllocator).
|
# where alloc lives on a separate allocator object).
|
||||||
# When None, fall back to entry.device_pool.alloc/free.
|
# When None, fall back to entry.device_pool.alloc/free.
|
||||||
device_alloc_fn: Optional[Callable] = None
|
device_alloc_fn: Optional[Callable] = None
|
||||||
device_free_fn: Optional[Callable] = None
|
device_free_fn: Optional[Callable] = None
|
||||||
|
|||||||
@@ -96,7 +96,7 @@ class MambaComponent(TreeComponent):
|
|||||||
if cow_mamba and mamba_value is not None:
|
if cow_mamba and mamba_value is not None:
|
||||||
assert req is not None
|
assert req is not None
|
||||||
if req.mamba_pool_idx is None:
|
if req.mamba_pool_idx is None:
|
||||||
dst_index = self.cache.req_to_token_pool.mamba_pool.alloc(1)
|
dst_index = self.cache.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
if dst_index is None:
|
if dst_index is None:
|
||||||
# Capture the inc result and thread swa_uuid_for_lock back
|
# Capture the inc result and thread swa_uuid_for_lock back
|
||||||
# into dec. Without it, SWA's release walks past this
|
# into dec. Without it, SWA's release walks past this
|
||||||
@@ -105,7 +105,7 @@ class MambaComponent(TreeComponent):
|
|||||||
# on ancestor nodes.
|
# on ancestor nodes.
|
||||||
lock_result = self.cache.inc_lock_ref(last_node)
|
lock_result = self.cache.inc_lock_ref(last_node)
|
||||||
self.cache.evict(EvictParams(num_tokens=0, mamba_num=1))
|
self.cache.evict(EvictParams(num_tokens=0, mamba_num=1))
|
||||||
dst_index = self.cache.req_to_token_pool.mamba_pool.alloc(1)
|
dst_index = self.cache.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
self.cache.dec_lock_ref(last_node, lock_result.to_dec_params())
|
self.cache.dec_lock_ref(last_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.mamba_pool_idx = dst_index[0]
|
||||||
@@ -172,7 +172,7 @@ class MambaComponent(TreeComponent):
|
|||||||
|
|
||||||
# Device layer
|
# Device layer
|
||||||
if EvictLayer.DEVICE in target and cd.value is not None:
|
if EvictLayer.DEVICE in target and cd.value is not None:
|
||||||
self.cache.req_to_token_pool.mamba_pool.free(cd.value)
|
self.cache.req_to_token_pool.mamba_allocator.free(cd.value)
|
||||||
freed = len(cd.value)
|
freed = len(cd.value)
|
||||||
self.cache.component_evictable_size_[self.component_type] -= freed
|
self.cache.component_evictable_size_[self.component_type] -= freed
|
||||||
cd.value = None
|
cd.value = None
|
||||||
@@ -285,10 +285,10 @@ class MambaComponent(TreeComponent):
|
|||||||
|
|
||||||
def _alloc_mamba_slot(self) -> torch.Tensor:
|
def _alloc_mamba_slot(self) -> torch.Tensor:
|
||||||
"""Allocate one mamba pool slot, evicting if necessary."""
|
"""Allocate one mamba pool slot, evicting if necessary."""
|
||||||
slot = self.cache.req_to_token_pool.mamba_pool.alloc(1)
|
slot = self.cache.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
if slot is None:
|
if slot is None:
|
||||||
self.cache.evict(EvictParams(num_tokens=0, mamba_num=1))
|
self.cache.evict(EvictParams(num_tokens=0, mamba_num=1))
|
||||||
slot = self.cache.req_to_token_pool.mamba_pool.alloc(1)
|
slot = self.cache.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
assert slot is not None, "Can not alloc mamba cache"
|
assert slot is not None, "Can not alloc mamba cache"
|
||||||
return slot
|
return slot
|
||||||
|
|
||||||
@@ -365,7 +365,9 @@ class MambaComponent(TreeComponent):
|
|||||||
if insert_params.mamba_value is not None and (
|
if insert_params.mamba_value is not None and (
|
||||||
insert_result is None or insert_result.mamba_exist
|
insert_result is None or insert_result.mamba_exist
|
||||||
):
|
):
|
||||||
self.cache.req_to_token_pool.mamba_pool.free(insert_params.mamba_value)
|
self.cache.req_to_token_pool.mamba_allocator.free(
|
||||||
|
insert_params.mamba_value
|
||||||
|
)
|
||||||
req.mamba_last_track_seqlen = None
|
req.mamba_last_track_seqlen = None
|
||||||
|
|
||||||
# ---- HiCache Hooks ----
|
# ---- HiCache Hooks ----
|
||||||
@@ -414,10 +416,10 @@ class MambaComponent(TreeComponent):
|
|||||||
cd = node.component_data[ct]
|
cd = node.component_data[ct]
|
||||||
if req is not None and cd.host_value is not None:
|
if req is not None and cd.host_value is not None:
|
||||||
if req.mamba_pool_idx is None:
|
if req.mamba_pool_idx is None:
|
||||||
dst = self.cache.req_to_token_pool.mamba_pool.alloc(1)
|
dst = self.cache.req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
if dst is None:
|
if dst is None:
|
||||||
self.cache.evict(EvictParams(num_tokens=0, mamba_num=1))
|
self.cache.evict(EvictParams(num_tokens=0, mamba_num=1))
|
||||||
dst = self.cache.req_to_token_pool.mamba_pool.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.mamba_pool_idx = dst[0]
|
||||||
transfers.append(
|
transfers.append(
|
||||||
|
|||||||
@@ -2511,7 +2511,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
if ct.is_swa:
|
if ct.is_swa:
|
||||||
available_size = self.token_to_kv_pool_allocator.swa_available_size()
|
available_size = self.token_to_kv_pool_allocator.swa_available_size()
|
||||||
elif ct.is_mamba:
|
elif ct.is_mamba:
|
||||||
available_size = self.req_to_token_pool.mamba_pool.available_size()
|
available_size = self.req_to_token_pool.mamba_allocator.available_size()
|
||||||
else:
|
else:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
|||||||
@@ -500,14 +500,14 @@ class StreamingSession(BasePrefixCache):
|
|||||||
|
|
||||||
def _free_slot_mamba(self, slot: SessionSlot) -> None:
|
def _free_slot_mamba(self, slot: SessionSlot) -> None:
|
||||||
"""Return a session slot's mamba pool state to the allocator."""
|
"""Return a session slot's mamba pool state to the allocator."""
|
||||||
mamba_pool = getattr(self.req_to_token_pool, "mamba_pool", None)
|
mamba_allocator = getattr(self.req_to_token_pool, "mamba_allocator", None)
|
||||||
if mamba_pool is None:
|
if mamba_allocator is None:
|
||||||
return
|
return
|
||||||
if slot.mamba_pool_idx is not None:
|
if slot.mamba_pool_idx is not None:
|
||||||
mamba_pool.free(slot.mamba_pool_idx.unsqueeze(0))
|
mamba_allocator.free(slot.mamba_pool_idx.unsqueeze(0))
|
||||||
slot.mamba_pool_idx = None
|
slot.mamba_pool_idx = None
|
||||||
if slot.mamba_ping_pong_track_buffer is not None:
|
if slot.mamba_ping_pong_track_buffer is not None:
|
||||||
mamba_pool.free(slot.mamba_ping_pong_track_buffer)
|
mamba_allocator.free(slot.mamba_ping_pong_track_buffer)
|
||||||
slot.mamba_ping_pong_track_buffer = None
|
slot.mamba_ping_pong_track_buffer = None
|
||||||
|
|
||||||
# -- Internal helpers (streaming body bits) --
|
# -- Internal helpers (streaming body bits) --
|
||||||
|
|||||||
@@ -109,7 +109,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert req_to_token_pool.available_size() == max_num_reqs
|
assert req_to_token_pool.available_size() == max_num_reqs
|
||||||
assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size
|
assert req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size
|
||||||
|
|
||||||
sampling_params = SamplingParams(
|
sampling_params = SamplingParams(
|
||||||
temperature=0,
|
temperature=0,
|
||||||
@@ -125,34 +125,41 @@ class TestMamba(unittest.TestCase):
|
|||||||
# alloc req
|
# alloc req
|
||||||
req_to_token_pool.alloc([req])
|
req_to_token_pool.alloc([req])
|
||||||
assert req_to_token_pool.available_size() == max_num_reqs - 1
|
assert req_to_token_pool.available_size() == max_num_reqs - 1
|
||||||
assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size - 1
|
assert (
|
||||||
|
req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1
|
||||||
|
)
|
||||||
|
|
||||||
# free req
|
# free req
|
||||||
req_to_token_pool.free_mamba_cache(req)
|
req_to_token_pool.free_mamba_cache(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
|
||||||
assert req_to_token_pool.mamba_pool.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.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
|
||||||
assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size - 1
|
assert (
|
||||||
|
req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1
|
||||||
|
)
|
||||||
|
|
||||||
# alloc again
|
# alloc again
|
||||||
req_to_token_pool.alloc([req])
|
req_to_token_pool.alloc([req])
|
||||||
assert req_to_token_pool.available_size() == max_num_reqs - 1
|
assert req_to_token_pool.available_size() == max_num_reqs - 1
|
||||||
assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size - 1
|
assert (
|
||||||
|
req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1
|
||||||
|
)
|
||||||
|
|
||||||
def test_mamba_radix_cache_1(self):
|
def test_mamba_radix_cache_1(self):
|
||||||
tree, allocator, req_to_token_pool, make_dummy_req = (
|
tree, allocator, req_to_token_pool, make_dummy_req = (
|
||||||
self._setup_tree_and_allocator()
|
self._setup_tree_and_allocator()
|
||||||
)
|
)
|
||||||
|
mamba_allocator = req_to_token_pool.mamba_allocator
|
||||||
mamba_pool = req_to_token_pool.mamba_pool
|
mamba_pool = req_to_token_pool.mamba_pool
|
||||||
# test
|
# test
|
||||||
print(
|
print(
|
||||||
f"[Start] allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
|
f"[Start] allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
|
||||||
)
|
)
|
||||||
req1 = make_dummy_req()
|
req1 = make_dummy_req()
|
||||||
req1_token_ids, req1_kv_indices = [1, 2, 3], allocator.alloc(3)
|
req1_token_ids, req1_kv_indices = [1, 2, 3], allocator.alloc(3)
|
||||||
@@ -170,7 +177,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
prefix_len = result.prefix_len
|
prefix_len = result.prefix_len
|
||||||
print(
|
print(
|
||||||
f"req1: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
|
f"req1: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
|
||||||
)
|
)
|
||||||
req2 = make_dummy_req()
|
req2 = make_dummy_req()
|
||||||
req2_token_ids, req2_kv_indices = [1, 2, 3, 4, 5, 6, 7], allocator.alloc(7)
|
req2_token_ids, req2_kv_indices = [1, 2, 3, 4, 5, 6, 7], allocator.alloc(7)
|
||||||
@@ -188,7 +195,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
prefix_len = result.prefix_len
|
prefix_len = result.prefix_len
|
||||||
print(
|
print(
|
||||||
f"req2: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
|
f"req2: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
|
||||||
)
|
)
|
||||||
|
|
||||||
req3 = make_dummy_req()
|
req3 = make_dummy_req()
|
||||||
@@ -207,7 +214,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
prefix_len = result.prefix_len
|
prefix_len = result.prefix_len
|
||||||
print(
|
print(
|
||||||
f"req3: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
|
f"req3: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
|
||||||
)
|
)
|
||||||
req4 = make_dummy_req()
|
req4 = make_dummy_req()
|
||||||
req4_token_ids, req4_kv_indices = [1, 2, 3, 4, 5, 60, 70], allocator.alloc(7)
|
req4_token_ids, req4_kv_indices = [1, 2, 3, 4, 5, 60, 70], allocator.alloc(7)
|
||||||
@@ -225,7 +232,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
prefix_len = result.prefix_len
|
prefix_len = result.prefix_len
|
||||||
print(
|
print(
|
||||||
f"req4: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
|
f"req4: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
|
||||||
)
|
)
|
||||||
|
|
||||||
tree.pretty_print()
|
tree.pretty_print()
|
||||||
@@ -553,7 +560,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
_, _, req_to_token_pool, _ = self._setup_tree_and_allocator()
|
_, _, req_to_token_pool, _ = self._setup_tree_and_allocator()
|
||||||
mamba_pool = req_to_token_pool.mamba_pool
|
mamba_pool = req_to_token_pool.mamba_pool
|
||||||
n = 3
|
n = 3
|
||||||
indices = mamba_pool.alloc(n)
|
indices = req_to_token_pool.mamba_allocator.alloc(n)
|
||||||
self.assertIsNotNone(indices)
|
self.assertIsNotNone(indices)
|
||||||
|
|
||||||
# Write known sentinel values at the allocated slots.
|
# Write known sentinel values at the allocated slots.
|
||||||
@@ -608,7 +615,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
n_tokens = 4
|
n_tokens = 4
|
||||||
kv_indices = allocator.alloc(n_tokens)
|
kv_indices = allocator.alloc(n_tokens)
|
||||||
self.assertIsNotNone(kv_indices)
|
self.assertIsNotNone(kv_indices)
|
||||||
mamba_indices = mamba_pool.alloc(1)
|
mamba_indices = req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
self.assertIsNotNone(mamba_indices)
|
self.assertIsNotNone(mamba_indices)
|
||||||
|
|
||||||
# Write sentinel values into KV buffers (all full-attention layers).
|
# Write sentinel values into KV buffers (all full-attention layers).
|
||||||
|
|||||||
@@ -3545,7 +3545,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
xfer = tree.components[ComponentType.MAMBA].build_hicache_transfers(
|
xfer = tree.components[ComponentType.MAMBA].build_hicache_transfers(
|
||||||
node, CacheTransferPhase.LOAD_BACK
|
node, CacheTransferPhase.LOAD_BACK
|
||||||
)[0]
|
)[0]
|
||||||
new_mamba = req_to_token_pool.mamba_pool.alloc(1)
|
new_mamba = req_to_token_pool.mamba_allocator.alloc(1)
|
||||||
self.assertIsNotNone(new_mamba)
|
self.assertIsNotNone(new_mamba)
|
||||||
xfer.device_indices = new_mamba
|
xfer.device_indices = new_mamba
|
||||||
tree.components[ComponentType.MAMBA].commit_hicache_transfer(
|
tree.components[ComponentType.MAMBA].commit_hicache_transfer(
|
||||||
|
|||||||
Reference in New Issue
Block a user