[PD][Bugfix] fix mamba cache capping (#22462)
Co-authored-by: hzh0425 <hzh0425@apache.org> Co-authored-by: yizhang2077 <1109276519@qq.com>
This commit is contained in:
co-authored by
hzh0425
yizhang2077
parent
fdfe53e872
commit
2d2be5d7b2
@@ -197,18 +197,26 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool):
|
||||
self.mamba_ping_pong_track_buffer_size = 2 if enable_overlap_schedule else 1
|
||||
self.enable_mamba_extra_buffer = enable_mamba_extra_buffer
|
||||
self.enable_memory_saver = enable_memory_saver
|
||||
# Each request needs 1 main mamba slot + ping-pong slots when extra_buffer is enabled.
|
||||
# Cap the pool at max concurrent requests * slots_per_req to avoid allocating failed.
|
||||
slots_per_req = 1 + (
|
||||
self.mamba_ping_pong_track_buffer_size if enable_mamba_extra_buffer else 0
|
||||
)
|
||||
max_slots_needed = (size + pre_alloc_size) * slots_per_req
|
||||
if mamba_size is not None:
|
||||
effective_mamba_size = min(mamba_size, size + pre_alloc_size)
|
||||
if mamba_size > size + pre_alloc_size:
|
||||
effective_mamba_size = max(mamba_size, max_slots_needed)
|
||||
if mamba_size < max_slots_needed:
|
||||
logger.warning(
|
||||
"mamba_size (%d) exceeds size + pre_alloc_size (%d), "
|
||||
"capping effective_mamba_size to %d",
|
||||
"mamba_size (%d) is less than decode side's max_slots_needed (%d = %d reqs * %d slots/req), "
|
||||
"raising effective_mamba_size to %d",
|
||||
mamba_size,
|
||||
max_slots_needed,
|
||||
size + pre_alloc_size,
|
||||
slots_per_req,
|
||||
effective_mamba_size,
|
||||
)
|
||||
else:
|
||||
effective_mamba_size = size + pre_alloc_size
|
||||
effective_mamba_size = max_slots_needed
|
||||
self.start_layer = start_layer if start_layer is not None else 0
|
||||
self.layer_transfer_counter = None
|
||||
self._init_mamba_pool(
|
||||
|
||||
@@ -3683,6 +3683,12 @@ class ServerArgs:
|
||||
if self.disaggregation_mode == "decode":
|
||||
self.disable_radix_cache = True
|
||||
logger.warning("KV cache is forced as chunk cache for decode server")
|
||||
if self.enable_mamba_extra_buffer():
|
||||
logger.warning(
|
||||
"Mamba extra_buffer is disabled because decode disaggregation "
|
||||
"currently forces chunk cache. Falling back to no_buffer."
|
||||
)
|
||||
self.mamba_scheduler_strategy = "no_buffer"
|
||||
|
||||
elif self.disaggregation_mode == "prefill":
|
||||
assert (
|
||||
|
||||
Reference in New Issue
Block a user