[Qwen3-next] Add PD disaggregation support for mamba with extra_buffer (#15180)

Signed-off-by: Shangming Cai <csmthu@gmail.com>
Co-authored-by: ybyang <10629930+whybeyoung@users.noreply.github.com>
This commit is contained in:
Shangming Cai
2025-12-16 14:36:00 +08:00
committed by GitHub
parent 6292d97135
commit 36fcf71fff
5 changed files with 79 additions and 7 deletions

View File

@@ -144,6 +144,7 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool):
enable_memory_saver: bool,
cache_params: "Mamba2CacheParams",
speculative_num_draft_tokens: int,
enable_mamba_extra_buffer: bool,
pre_alloc_size: int,
):
DecodeReqToTokenPool.__init__(
@@ -154,10 +155,11 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool):
enable_memory_saver=enable_memory_saver,
pre_alloc_size=pre_alloc_size,
)
self.enable_memory_saver = enable_memory_saver
self.enable_mamba_extra_buffer = (
False # TODO: add PD support for mamba cache extra_buffer
self.mamba_ping_pong_track_buffer_size = (
2 if speculative_num_draft_tokens is None else 1
)
self.enable_mamba_extra_buffer = enable_mamba_extra_buffer
self.enable_memory_saver = enable_memory_saver
self._init_mamba_pool(
size=size + pre_alloc_size,
mamba_spec_state_size=size + pre_alloc_size,

View File

@@ -1797,6 +1797,7 @@ class ModelRunner:
enable_memory_saver=self.server_args.enable_memory_saver,
cache_params=config.mamba2_cache_params,
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
pre_alloc_size=pre_alloc_size,
)
else:

View File

@@ -1382,9 +1382,6 @@ class ServerArgs:
assert (
is_cuda()
), "Mamba extra_buffer is only supported on CUDA devices with FLA backend"
assert (
self.disaggregation_mode == "null"
), "Mamba extra_buffer is not compatible with disaggregation mode yet."
if self.speculative_num_draft_tokens is not None:
assert (
self.mamba_track_interval >= self.speculative_num_draft_tokens