[Fix] data race in req_to_token pool (#17850)

This commit is contained in:
cctry
2026-02-02 14:38:15 -08:00
committed by GitHub
parent cbf1500390
commit 027f314050
13 changed files with 107 additions and 113 deletions
+22 -15
View File
@@ -25,7 +25,7 @@ import time
from collections import deque from collections import deque
from dataclasses import dataclass from dataclasses import dataclass
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, List, Optional, Tuple, Type, Union from typing import TYPE_CHECKING, List, Optional, Tuple, Type
import torch import torch
from torch.distributed import ProcessGroup from torch.distributed import ProcessGroup
@@ -116,19 +116,31 @@ class DecodeReqToTokenPool:
def available_size(self): def available_size(self):
return len(self.free_slots) return len(self.free_slots)
def alloc(self, need_size: int) -> List[int]: def alloc(self, reqs: List["Req"]) -> Optional[List[int]]:
chunked = [i for i, r in enumerate(reqs) if r.req_pool_idx is not None]
assert (
len(chunked) <= 1
), "only one chunked request may reuse req_pool_idx in a batch"
assert all(
reqs[i].is_chunked > 0 or reqs[i].kv_committed_len > 0 for i in chunked
), "request has req_pool_idx but is not chunked"
need_size = len(reqs) - len(chunked)
if need_size > len(self.free_slots): if need_size > len(self.free_slots):
return None return None
select_index = self.free_slots[:need_size] select_index = self.free_slots[:need_size]
self.free_slots = self.free_slots[need_size:] self.free_slots = self.free_slots[need_size:]
return select_index offset = 0
for r in reqs:
if r.req_pool_idx is None:
r.req_pool_idx = select_index[offset]
offset += 1
return [r.req_pool_idx for r in reqs]
def free(self, free_index: Union[int, List[int]]): def free(self, req: "Req"):
if isinstance(free_index, (int,)): assert req.req_pool_idx is not None, "request must have req_pool_idx"
self.free_slots.append(free_index) self.free_slots.append(req.req_pool_idx)
else: req.req_pool_idx = None
self.free_slots.extend(free_index)
def clear(self): def clear(self):
self.free_slots = list(range(self.size + self.pre_alloc_size)) self.free_slots = list(range(self.size + self.pre_alloc_size))
@@ -652,17 +664,12 @@ class DecodePreallocQueue:
def _pre_alloc(self, req: Req) -> torch.Tensor: def _pre_alloc(self, req: Req) -> torch.Tensor:
"""Pre-allocate the memory for req_to_token and token_kv_pool""" """Pre-allocate the memory for req_to_token and token_kv_pool"""
if isinstance(self.req_to_token_pool, HybridMambaDecodeReqToTokenPool): req_pool_indices = self.req_to_token_pool.alloc([req])
req_pool_indices = self.req_to_token_pool.alloc(1, [req])
else:
req_pool_indices = self.req_to_token_pool.alloc(1)
assert ( assert (
req_pool_indices is not None req_pool_indices is not None
), "req_pool_indices is full! There is a bug in memory estimation." ), "req_pool_indices is full! There is a bug in memory estimation."
req.req_pool_idx = req_pool_indices[0]
# Alloc all tokens for the prebuilt req (except for the reserved input token for decoding) # Alloc all tokens for the prebuilt req (except for the reserved input token for decoding)
fill_len = len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0) fill_len = len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0)
req.kv_allocated_len = fill_len req.kv_allocated_len = fill_len
@@ -191,7 +191,7 @@ class DecodeKVCacheOffloadManager:
# Free the incremental part of the request # Free the incremental part of the request
self.token_to_kv_pool_allocator.free(kv_indices) self.token_to_kv_pool_allocator.free(kv_indices)
self.req_to_token_pool.free(req.req_pool_idx) self.req_to_token_pool.free(req)
self.tree_cache.protected_size_ -= len(req.prefix_indices) self.tree_cache.protected_size_ -= len(req.prefix_indices)
def _check_backup_progress(self, finish_count): def _check_backup_progress(self, finish_count):
@@ -632,13 +632,6 @@ class SchedulerDisaggregationPrefillMixin:
) )
else: else:
self.send_kv_chunk(self.chunked_req) self.send_kv_chunk(self.chunked_req)
# chunked request keeps its rid but will get a new req_pool_idx
if self.tp_worker.model_runner.mambaish_config is not None:
self.req_to_token_pool.free(
self.chunked_req.req_pool_idx, free_mamba_cache=False
)
else:
self.req_to_token_pool.free(self.chunked_req.req_pool_idx)
self.running_batch.batch_is_full = False self.running_batch.batch_is_full = False
if self.last_batch and self.last_batch.forward_mode.is_extend(): if self.last_batch and self.last_batch.forward_mode.is_extend():
-5
View File
@@ -1792,11 +1792,6 @@ class Scheduler(
def stash_chunked_request(self, req: Req): def stash_chunked_request(self, req: Req):
self.tree_cache.cache_unfinished_req(req, chunked=True) self.tree_cache.cache_unfinished_req(req, chunked=True)
# Chunked request keeps its rid but will get a new req_pool_idx
if self.tp_worker.model_runner.mambaish_config is not None:
self.req_to_token_pool.free(req.req_pool_idx, free_mamba_cache=False)
else:
self.req_to_token_pool.free(req.req_pool_idx)
def get_next_batch_to_run(self) -> Optional[ScheduleBatch]: def get_next_batch_to_run(self) -> Optional[ScheduleBatch]:
self._abort_on_queued_timeout() self._abort_on_queued_timeout()
@@ -644,7 +644,7 @@ class SchedulerPPMixin:
req.req_pool_idx, : len(req.fill_ids) req.req_pool_idx, : len(req.fill_ids)
] ]
self.token_to_kv_pool_allocator.free(kv_indices) self.token_to_kv_pool_allocator.free(kv_indices)
self.req_to_token_pool.free(req.req_pool_idx) self.req_to_token_pool.free(req)
logger.info( logger.info(
f"[PP Dynamic Chunk] [PP0] Profiled {len(seq_lens)} samples: " f"[PP Dynamic Chunk] [PP0] Profiled {len(seq_lens)} samples: "
@@ -68,7 +68,6 @@ class ChunkCache(BasePrefixCache):
kv_indices = self.req_to_token_pool.req_to_token[ kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, :kv_committed_len req.req_pool_idx, :kv_committed_len
] ]
self.req_to_token_pool.free(req.req_pool_idx)
self.token_to_kv_pool_allocator.free(kv_indices) self.token_to_kv_pool_allocator.free(kv_indices)
def cache_unfinished_req(self, req: Req, chunked=False): def cache_unfinished_req(self, req: Req, chunked=False):
+27 -17
View File
@@ -296,11 +296,11 @@ def alloc_paged_token_slots_extend(
def alloc_req_slots( def alloc_req_slots(
req_to_token_pool: ReqToTokenPool, req_to_token_pool: ReqToTokenPool,
num_reqs: int, reqs: list[Req],
reqs: list[Req] | None,
tree_cache: BasePrefixCache | None, tree_cache: BasePrefixCache | None,
) -> list[int]: ) -> list[int]:
"""Allocate request slots from the pool.""" """Allocate request slots from the pool."""
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_pool.available_size()
factor = ( factor = (
@@ -313,9 +313,7 @@ def alloc_req_slots(
if tree_cache is not None and tree_cache.supports_mamba(): if tree_cache is not None and tree_cache.supports_mamba():
mamba_num = max(0, mamba_state_needed - mamba_available_size) mamba_num = max(0, mamba_state_needed - mamba_available_size)
tree_cache.evict(EvictParams(num_tokens=0, mamba_num=mamba_num)) tree_cache.evict(EvictParams(num_tokens=0, mamba_num=mamba_num))
req_pool_indices = req_to_token_pool.alloc(num_reqs, reqs) req_pool_indices = req_to_token_pool.alloc(reqs)
else:
req_pool_indices = req_to_token_pool.alloc(num_reqs)
if req_pool_indices is None: if req_pool_indices is None:
raise RuntimeError( raise RuntimeError(
@@ -341,7 +339,6 @@ def alloc_for_extend(
# free out-of-window swa tokens # free out-of-window swa tokens
batch.maybe_evict_swa() batch.maybe_evict_swa()
bs = len(batch.reqs)
prefix_tensors = [r.prefix_indices for r in batch.reqs] prefix_tensors = [r.prefix_indices for r in batch.reqs]
# Create tensors for allocation # Create tensors for allocation
@@ -352,7 +349,7 @@ def alloc_for_extend(
# Allocate req slots # Allocate req slots
req_pool_indices = alloc_req_slots( req_pool_indices = alloc_req_slots(
batch.req_to_token_pool, bs, batch.reqs, batch.tree_cache batch.req_to_token_pool, batch.reqs, batch.tree_cache
) )
req_pool_indices_cpu = torch.tensor(req_pool_indices, dtype=torch.int64) req_pool_indices_cpu = torch.tensor(req_pool_indices, dtype=torch.int64)
req_pool_indices_device = req_pool_indices_cpu.to(batch.device, non_blocking=True) req_pool_indices_device = req_pool_indices_cpu.to(batch.device, non_blocking=True)
@@ -466,15 +463,21 @@ def alloc_for_decode(batch: ScheduleBatch, token_per_req: int) -> torch.Tensor:
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):
tree_cache.cache_finished_req(req, is_insert=is_insert)
# MambaRadixCache may alloc mamba state before alloc KV cache # MambaRadixCache may alloc mamba state before alloc KV cache
if req.req_pool_idx is None: if req.req_pool_idx is None:
assert ( assert (
tree_cache.supports_mamba() tree_cache.supports_mamba()
), "Only MambaRadixCache can handle abort with prefix cache hit before alloc" ), "Only MambaRadixCache allow freeing before alloc"
# TODO (csy, hanming): clean up this early allocation logic
if req.mamba_pool_idx is not None:
tree_cache.req_to_token_pool.mamba_pool.free(
req.mamba_pool_idx.unsqueeze(-1)
)
req.mamba_pool_idx = None
return return
tree_cache.cache_finished_req(req, is_insert=is_insert)
start_p, end_p = req.pop_overallocated_kv_cache() start_p, end_p = req.pop_overallocated_kv_cache()
global_server_args = get_global_server_args() global_server_args = get_global_server_args()
@@ -489,13 +492,20 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
if page_size > 1: if page_size > 1:
start_p = ceil_align(start_p, page_size) start_p = ceil_align(start_p, page_size)
if start_p >= end_p: if start_p < end_p:
return indices_to_free = tree_cache.req_to_token_pool.req_to_token[req.req_pool_idx][
start_p:end_p
indices_to_free = tree_cache.req_to_token_pool.req_to_token[req.req_pool_idx][ ]
start_p:end_p tree_cache.token_to_kv_pool_allocator.free(indices_to_free)
] # If the prefix cache doesn't manage mamba states, we must free them here.
tree_cache.token_to_kv_pool_allocator.free(indices_to_free) if isinstance(tree_cache.req_to_token_pool, HybridReqToTokenPool) and (
not tree_cache.supports_mamba()
):
assert (
req.mamba_pool_idx is not None
), "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(req)
def available_and_evictable_str(tree_cache) -> str: def available_and_evictable_str(tree_cache) -> str:
@@ -499,13 +499,6 @@ class MambaRadixCache(BasePrefixCache):
def cache_finished_req(self, req: Req, is_insert: bool = True) -> None: def cache_finished_req(self, req: Req, is_insert: bool = True) -> None:
"""Cache request when it finishes.""" """Cache request when it finishes."""
# for abort with prefix cache hit and before alloc is called
if req.req_pool_idx is None:
if req.mamba_pool_idx is not None:
self.req_to_token_pool.mamba_pool.free(req.mamba_pool_idx.unsqueeze(-1))
req.mamba_pool_idx = None
return
kv_committed_len = req.pop_committed_kv_cache() kv_committed_len = req.pop_committed_kv_cache()
if self.disable: if self.disable:
@@ -513,7 +506,7 @@ class MambaRadixCache(BasePrefixCache):
req.req_pool_idx, :kv_committed_len req.req_pool_idx, :kv_committed_len
] ]
self.token_to_kv_pool_allocator.free(kv_indices) self.token_to_kv_pool_allocator.free(kv_indices)
self.req_to_token_pool.free(req.req_pool_idx) self.req_to_token_pool.free_mamba_cache(req)
return return
token_ids = (req.origin_input_ids + req.output_ids)[:kv_committed_len] token_ids = (req.origin_input_ids + req.output_ids)[:kv_committed_len]
@@ -588,11 +581,11 @@ class MambaRadixCache(BasePrefixCache):
free_mamba_cache = True if self.enable_mamba_extra_buffer else mamba_exist free_mamba_cache = True if self.enable_mamba_extra_buffer else mamba_exist
self.req_to_token_pool.free( if free_mamba_cache:
req.req_pool_idx, self.req_to_token_pool.free_mamba_cache(
free_mamba_cache=free_mamba_cache, req,
mamba_ping_pong_track_buffer_to_keep=mamba_ping_pong_track_buffer_to_keep, mamba_ping_pong_track_buffer_to_keep=mamba_ping_pong_track_buffer_to_keep,
) )
self.dec_lock_ref(req.last_node) self.dec_lock_ref(req.last_node)
+43 -42
View File
@@ -133,7 +133,6 @@ class ReqToTokenPool:
device: str, device: str,
enable_memory_saver: bool, enable_memory_saver: bool,
): ):
memory_saver_adapter = TorchMemorySaverAdapter.create( memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=enable_memory_saver enable=enable_memory_saver
) )
@@ -145,7 +144,6 @@ class ReqToTokenPool:
self.req_to_token = torch.zeros( self.req_to_token = torch.zeros(
(size, max_context_len), dtype=torch.int32, device=device (size, max_context_len), dtype=torch.int32, device=device
) )
self.free_slots = list(range(size)) self.free_slots = list(range(size))
def write(self, indices, values): def write(self, indices, values):
@@ -154,20 +152,32 @@ class ReqToTokenPool:
def available_size(self): def available_size(self):
return len(self.free_slots) return len(self.free_slots)
def alloc(self, need_size: int) -> List[int]: def alloc(self, reqs: list[Req]) -> Optional[List[int]]:
chunked = [i for i, r in enumerate(reqs) if r.req_pool_idx is not None]
if not any(r.is_dllm() for r in reqs):
assert (
len(chunked) <= 1
), "only one chunked request may reuse req_pool_idx in a batch"
assert all(
reqs[i].is_chunked > 0 or reqs[i].kv_committed_len > 0 for i in chunked
), "request has req_pool_idx but is not chunked"
need_size = len(reqs) - len(chunked)
if need_size > len(self.free_slots): if need_size > len(self.free_slots):
return None return None
select_index = self.free_slots[:need_size] select_index = self.free_slots[:need_size]
self.free_slots = self.free_slots[need_size:] self.free_slots = self.free_slots[need_size:]
offset = 0
for r in reqs:
if r.req_pool_idx is None:
r.req_pool_idx = select_index[offset]
offset += 1
return [r.req_pool_idx for r in reqs]
return select_index def free(self, req: Req):
assert req.req_pool_idx is not None, "request must have req_pool_idx"
def free(self, free_index: Union[int, List[int]]): self.free_slots.append(req.req_pool_idx)
if isinstance(free_index, (int,)): req.req_pool_idx = None
self.free_slots.append(free_index)
else:
self.free_slots.extend(free_index)
def clear(self): def clear(self):
self.free_slots = list(range(self.size)) self.free_slots = list(range(self.size))
@@ -488,10 +498,9 @@ class HybridReqToTokenPool(ReqToTokenPool):
# For chunk prefill req, we do not need to allocate mamba cache, # For chunk prefill req, we do not need to allocate mamba cache,
# We could use allocated mamba cache instead. # We could use allocated mamba cache instead.
def alloc(self, need_size: int, reqs: Optional[List["Req"]]) -> Optional[List[int]]: def alloc(self, reqs: List["Req"]) -> Optional[List[int]]:
assert reqs is not None select_index = super().alloc(reqs)
select_index = super().alloc(need_size) if select_index is None:
if select_index == None:
return None return None
mamba_index = [] mamba_index = []
@@ -556,37 +565,29 @@ class HybridReqToTokenPool(ReqToTokenPool):
else: else:
return mamba_next_track_idx return mamba_next_track_idx
# For chunk prefill, we can not free mamba cache, we need use it in the future def free_mamba_cache(
def free( self, req: "Req", mamba_ping_pong_track_buffer_to_keep: Optional[int] = None
self,
free_index: Union[int, List[int]],
free_mamba_cache: bool = True,
mamba_ping_pong_track_buffer_to_keep: Optional[int] = None,
): ):
if isinstance(free_index, (int,)): mamba_index = req.mamba_pool_idx
free_index = [free_index] assert mamba_index is not None, "double free? mamba_index is None"
super().free(free_index) self.mamba_pool.free(mamba_index.unsqueeze(0))
if free_mamba_cache: req.mamba_pool_idx = None
mamba_index = self.req_index_to_mamba_index_mapping[free_index]
self.mamba_pool.free(mamba_index)
if self.enable_mamba_extra_buffer: if self.enable_mamba_extra_buffer:
mamba_ping_pong_track_buffer_to_free = (
self.req_index_to_mamba_ping_pong_track_buffer_mapping[req.req_pool_idx]
)
if mamba_ping_pong_track_buffer_to_keep is not None:
assert mamba_ping_pong_track_buffer_to_keep in [
0,
1,
], f"mamba_ping_pong_track_buffer_to_keep must be 0 or 1, {mamba_ping_pong_track_buffer_to_keep=}"
idx_to_free = list(range(self.mamba_ping_pong_track_buffer_size))
idx_to_free.remove(mamba_ping_pong_track_buffer_to_keep)
mamba_ping_pong_track_buffer_to_free = ( mamba_ping_pong_track_buffer_to_free = (
self.req_index_to_mamba_ping_pong_track_buffer_mapping[ mamba_ping_pong_track_buffer_to_free[idx_to_free]
free_index
].squeeze(0)
) )
if mamba_ping_pong_track_buffer_to_keep is not None: self.mamba_pool.free(mamba_ping_pong_track_buffer_to_free)
assert mamba_ping_pong_track_buffer_to_keep in [
0,
1,
], f"mamba_ping_pong_track_buffer_to_keep must be 0 or 1, {mamba_ping_pong_track_buffer_to_keep=}"
idx_to_free = list(range(self.mamba_ping_pong_track_buffer_size))
idx_to_free.remove(mamba_ping_pong_track_buffer_to_keep)
mamba_ping_pong_track_buffer_to_free = (
mamba_ping_pong_track_buffer_to_free[idx_to_free]
)
self.mamba_pool.free(mamba_ping_pong_track_buffer_to_free)
def clear(self): def clear(self):
logger.info("Reset HybridReqToTokenPool") logger.info("Reset HybridReqToTokenPool")
@@ -451,7 +451,6 @@ class RadixCache(BasePrefixCache):
req.req_pool_idx, :kv_committed_len req.req_pool_idx, :kv_committed_len
] ]
self.token_to_kv_pool_allocator.free(kv_indices) self.token_to_kv_pool_allocator.free(kv_indices)
self.req_to_token_pool.free(req.req_pool_idx)
return return
token_ids = (req.origin_input_ids + req.output_ids)[:kv_committed_len] token_ids = (req.origin_input_ids + req.output_ids)[:kv_committed_len]
@@ -485,7 +484,6 @@ class RadixCache(BasePrefixCache):
self.token_to_kv_pool_allocator.free(kv_indices[len(keys) :]) self.token_to_kv_pool_allocator.free(kv_indices[len(keys) :])
# Remove req slot release the cache lock # Remove req slot release the cache lock
self.req_to_token_pool.free(req.req_pool_idx)
self.dec_lock_ref(req.last_node) self.dec_lock_ref(req.last_node)
def cache_unfinished_req(self, req: Req, chunked=False): def cache_unfinished_req(self, req: Req, chunked=False):
@@ -198,7 +198,6 @@ class RadixCacheCpp(BasePrefixCache):
# Remove req slot release the cache lock # Remove req slot release the cache lock
self.dec_lock_ref(req.last_node) self.dec_lock_ref(req.last_node)
self.req_to_token_pool.free(req.req_pool_idx)
def cache_unfinished_req(self, req: Req, chunked=False): def cache_unfinished_req(self, req: Req, chunked=False):
"""Cache request when it is unfinished.""" """Cache request when it is unfinished."""
@@ -461,7 +461,6 @@ class SWARadixCache(BasePrefixCache):
req.req_pool_idx, :kv_committed_len req.req_pool_idx, :kv_committed_len
] ]
self.token_to_kv_pool_allocator.free(kv_indices) self.token_to_kv_pool_allocator.free(kv_indices)
self.req_to_token_pool.free(req.req_pool_idx)
return return
token_ids = (req.origin_input_ids + req.output_ids)[:kv_committed_len] token_ids = (req.origin_input_ids + req.output_ids)[:kv_committed_len]
@@ -512,7 +511,6 @@ class SWARadixCache(BasePrefixCache):
self.token_to_kv_pool_allocator.free(kv_indices[page_aligned_len:]) self.token_to_kv_pool_allocator.free(kv_indices[page_aligned_len:])
# Remove req slot release the cache lock # Remove req slot release the cache lock
self.req_to_token_pool.free(req.req_pool_idx)
self.dec_lock_ref(req.last_node, req.swa_uuid_for_lock) self.dec_lock_ref(req.last_node, req.swa_uuid_for_lock)
def cache_unfinished_req(self, req: Req, chunked=False) -> None: def cache_unfinished_req(self, req: Req, chunked=False) -> None:
@@ -116,24 +116,25 @@ class TestMamba(unittest.TestCase):
) )
# alloc req # alloc req
req_index = req_to_token_pool.alloc(1, [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_pool.available_size() == mamba_cache_size - 1
# free req # free req
req_to_token_pool.free(req_index) req_to_token_pool.free_mamba_cache(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_pool.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_index = req_to_token_pool.alloc(1, [req]) req_to_token_pool.alloc([req])
req_to_token_pool.free(req_index, free_mamba_cache=False) 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_pool.available_size() == mamba_cache_size - 1
# alloc again # alloc again
req_index = req_to_token_pool.alloc(1, [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_pool.available_size() == mamba_cache_size - 1
@@ -225,7 +226,7 @@ class TestMamba(unittest.TestCase):
origin_input_ids=[], origin_input_ids=[],
sampling_params=sampling_params, sampling_params=sampling_params,
) )
req_to_token_pool.alloc(1, reqs=[req]) req_to_token_pool.alloc([req])
return req return req
mamba_pool = req_to_token_pool.mamba_pool mamba_pool = req_to_token_pool.mamba_pool