Enable RadixCache for Mamba2 models (#13584)
This commit is contained in:
@@ -362,7 +362,7 @@ class PrefillAdder:
|
|||||||
self.is_hybrid_swa = isinstance(
|
self.is_hybrid_swa = isinstance(
|
||||||
self.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator
|
self.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator
|
||||||
)
|
)
|
||||||
self.is_hybrid_gdn_cache = isinstance(self.tree_cache, MambaRadixCache)
|
self.is_ssm_radix_cache = isinstance(self.tree_cache, MambaRadixCache)
|
||||||
|
|
||||||
self.priority_scheduling_preemption_threshold = (
|
self.priority_scheduling_preemption_threshold = (
|
||||||
priority_scheduling_preemption_threshold
|
priority_scheduling_preemption_threshold
|
||||||
@@ -387,7 +387,7 @@ class PrefillAdder:
|
|||||||
self.token_to_kv_pool_allocator.swa_available_size()
|
self.token_to_kv_pool_allocator.swa_available_size()
|
||||||
+ self.tree_cache.swa_evictable_size(),
|
+ self.tree_cache.swa_evictable_size(),
|
||||||
)
|
)
|
||||||
elif self.is_hybrid_gdn_cache:
|
elif self.is_ssm_radix_cache:
|
||||||
available_and_evictable = (
|
available_and_evictable = (
|
||||||
self.token_to_kv_pool_allocator.available_size()
|
self.token_to_kv_pool_allocator.available_size()
|
||||||
+ self.tree_cache.full_evictable_size()
|
+ self.tree_cache.full_evictable_size()
|
||||||
@@ -409,7 +409,7 @@ class PrefillAdder:
|
|||||||
self.token_to_kv_pool_allocator.swa_available_size()
|
self.token_to_kv_pool_allocator.swa_available_size()
|
||||||
+ self.tree_cache.swa_evictable_size(),
|
+ self.tree_cache.swa_evictable_size(),
|
||||||
)
|
)
|
||||||
elif self.is_hybrid_gdn_cache:
|
elif self.is_ssm_radix_cache:
|
||||||
available_and_evictable = (
|
available_and_evictable = (
|
||||||
self.token_to_kv_pool_allocator.available_size()
|
self.token_to_kv_pool_allocator.available_size()
|
||||||
+ self.tree_cache.full_evictable_size()
|
+ self.tree_cache.full_evictable_size()
|
||||||
|
|||||||
@@ -398,7 +398,10 @@ class Scheduler(
|
|||||||
|
|
||||||
# Hybrid memory pool
|
# Hybrid memory pool
|
||||||
self.is_hybrid_swa = self.tp_worker.is_hybrid_swa
|
self.is_hybrid_swa = self.tp_worker.is_hybrid_swa
|
||||||
self.is_hybrid_gdn = self.tp_worker.model_runner.hybrid_gdn_config is not None
|
self.is_ssm_model = (
|
||||||
|
self.tp_worker.model_runner.hybrid_gdn_config is not None
|
||||||
|
or self.tp_worker.model_runner.mamba2_config is not None
|
||||||
|
)
|
||||||
|
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
self.sliding_window_size = self.tp_worker.sliding_window_size
|
self.sliding_window_size = self.tp_worker.sliding_window_size
|
||||||
@@ -762,7 +765,7 @@ class Scheduler(
|
|||||||
self.tree_cache = SWARadixCache(
|
self.tree_cache = SWARadixCache(
|
||||||
params=params, sliding_window_size=self.sliding_window_size
|
params=params, sliding_window_size=self.sliding_window_size
|
||||||
)
|
)
|
||||||
elif self.is_hybrid_gdn:
|
elif self.is_ssm_model:
|
||||||
from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache
|
from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache
|
||||||
|
|
||||||
self.tree_cache = MambaRadixCache(params)
|
self.tree_cache = MambaRadixCache(params)
|
||||||
|
|||||||
@@ -112,7 +112,7 @@ class SchedulerMetricsMixin:
|
|||||||
f"full token usage: {full_token_usage:.2f}, "
|
f"full token usage: {full_token_usage:.2f}, "
|
||||||
f"swa token usage: {swa_token_usage:.2f}, "
|
f"swa token usage: {swa_token_usage:.2f}, "
|
||||||
)
|
)
|
||||||
elif self.is_hybrid_gdn:
|
elif self.is_ssm_model:
|
||||||
(
|
(
|
||||||
full_num_used,
|
full_num_used,
|
||||||
_,
|
_,
|
||||||
@@ -166,7 +166,7 @@ class SchedulerMetricsMixin:
|
|||||||
self.stats.token_usage = token_usage
|
self.stats.token_usage = token_usage
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
self.stats.swa_token_usage = swa_token_usage
|
self.stats.swa_token_usage = swa_token_usage
|
||||||
if self.is_hybrid_gdn:
|
if self.is_ssm_model:
|
||||||
self.stats.mamba_usage = mamba_usage
|
self.stats.mamba_usage = mamba_usage
|
||||||
self.stats.num_queue_reqs = len(self.waiting_queue)
|
self.stats.num_queue_reqs = len(self.waiting_queue)
|
||||||
self.stats.num_grammar_queue_reqs = len(self.grammar_queue)
|
self.stats.num_grammar_queue_reqs = len(self.grammar_queue)
|
||||||
@@ -238,7 +238,7 @@ class SchedulerMetricsMixin:
|
|||||||
f"#swa token: {swa_num_used}, "
|
f"#swa token: {swa_num_used}, "
|
||||||
f"swa token usage: {swa_token_usage:.2f}, "
|
f"swa token usage: {swa_token_usage:.2f}, "
|
||||||
)
|
)
|
||||||
elif self.is_hybrid_gdn:
|
elif self.is_ssm_model:
|
||||||
(
|
(
|
||||||
full_num_used,
|
full_num_used,
|
||||||
mamba_used,
|
mamba_used,
|
||||||
@@ -315,7 +315,7 @@ class SchedulerMetricsMixin:
|
|||||||
self.stats.token_usage = token_usage
|
self.stats.token_usage = token_usage
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
self.stats.swa_token_usage = swa_token_usage
|
self.stats.swa_token_usage = swa_token_usage
|
||||||
if self.is_hybrid_gdn:
|
if self.is_ssm_model:
|
||||||
self.stats.mamba_usage = mamba_usage
|
self.stats.mamba_usage = mamba_usage
|
||||||
self.stats.gen_throughput = self.last_gen_throughput
|
self.stats.gen_throughput = self.last_gen_throughput
|
||||||
self.stats.num_queue_reqs = len(self.waiting_queue)
|
self.stats.num_queue_reqs = len(self.waiting_queue)
|
||||||
@@ -401,7 +401,7 @@ class SchedulerMetricsMixin:
|
|||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
full_num_used, swa_num_used, *_ = self._get_swa_token_info()
|
full_num_used, swa_num_used, *_ = self._get_swa_token_info()
|
||||||
num_tokens = max(full_num_used, swa_num_used)
|
num_tokens = max(full_num_used, swa_num_used)
|
||||||
elif self.is_hybrid_gdn:
|
elif self.is_ssm_model:
|
||||||
num_tokens = self._get_mamba_token_info()[0]
|
num_tokens = self._get_mamba_token_info()[0]
|
||||||
else:
|
else:
|
||||||
num_tokens = self._get_token_info()[0]
|
num_tokens = self._get_token_info()[0]
|
||||||
|
|||||||
@@ -204,7 +204,7 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
def check_memory(self: Scheduler):
|
def check_memory(self: Scheduler):
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
memory_leak, token_msg = self._check_hybrid_memory()
|
memory_leak, token_msg = self._check_hybrid_memory()
|
||||||
elif self.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache):
|
elif self.is_ssm_model and isinstance(self.tree_cache, MambaRadixCache):
|
||||||
memory_leak, token_msg = self._check_mamba_memory()
|
memory_leak, token_msg = self._check_mamba_memory()
|
||||||
else:
|
else:
|
||||||
memory_leak, token_msg = self._check_radix_cache_memory()
|
memory_leak, token_msg = self._check_radix_cache_memory()
|
||||||
@@ -239,7 +239,7 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
) = self._get_swa_token_info()
|
) = self._get_swa_token_info()
|
||||||
num_used = max(full_num_used, swa_num_used)
|
num_used = max(full_num_used, swa_num_used)
|
||||||
token_usage = max(full_token_usage, swa_token_usage)
|
token_usage = max(full_token_usage, swa_token_usage)
|
||||||
elif self.is_hybrid_gdn:
|
elif self.is_ssm_model:
|
||||||
(
|
(
|
||||||
num_used,
|
num_used,
|
||||||
_,
|
_,
|
||||||
@@ -278,7 +278,7 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
|
|
||||||
def check_tree_cache(self: Scheduler):
|
def check_tree_cache(self: Scheduler):
|
||||||
if (self.is_hybrid_swa and isinstance(self.tree_cache, SWARadixCache)) or (
|
if (self.is_hybrid_swa and isinstance(self.tree_cache, SWARadixCache)) or (
|
||||||
self.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache)
|
self.is_ssm_model and isinstance(self.tree_cache, MambaRadixCache)
|
||||||
):
|
):
|
||||||
self.tree_cache.sanity_check()
|
self.tree_cache.sanity_check()
|
||||||
|
|
||||||
@@ -322,7 +322,7 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
# Print batch size and memory pool info to check whether there are de-sync issues.
|
# Print batch size and memory pool info to check whether there are de-sync issues.
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
_, info_msg = self._check_hybrid_memory()
|
_, info_msg = self._check_hybrid_memory()
|
||||||
elif self.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache):
|
elif self.is_ssm_model and isinstance(self.tree_cache, MambaRadixCache):
|
||||||
_, info_msg = self._check_mamba_memory()
|
_, info_msg = self._check_mamba_memory()
|
||||||
else:
|
else:
|
||||||
_, info_msg = self._check_radix_cache_memory()
|
_, info_msg = self._check_radix_cache_memory()
|
||||||
|
|||||||
@@ -441,11 +441,6 @@ class ModelRunner:
|
|||||||
if architectures and not any("Llama4" in arch for arch in architectures):
|
if architectures and not any("Llama4" in arch for arch in architectures):
|
||||||
self.is_hybrid_swa = self.model_config.is_hybrid_swa = True
|
self.is_hybrid_swa = self.model_config.is_hybrid_swa = True
|
||||||
|
|
||||||
if config := self.mamba2_config:
|
|
||||||
class_name = config.__class__.__name__
|
|
||||||
logger.warning(f"{class_name} model detected, disable radix cache")
|
|
||||||
self.server_args.disable_radix_cache = True
|
|
||||||
|
|
||||||
# For MTP models like DeepSeek-V3 or GLM-4.5, the MTP layer(s) are used separately as draft
|
# For MTP models like DeepSeek-V3 or GLM-4.5, the MTP layer(s) are used separately as draft
|
||||||
# models for speculative decoding. In those cases, `num_nextn_predict_layers` is used to
|
# models for speculative decoding. In those cases, `num_nextn_predict_layers` is used to
|
||||||
# determine the number of layers.
|
# determine the number of layers.
|
||||||
|
|||||||
@@ -1278,6 +1278,18 @@ class ServerArgs:
|
|||||||
logger.info(
|
logger.info(
|
||||||
"Use flashinfer_trtllm as MoE runner backend on sm100 for Qwen3NextForCausalLM"
|
"Use flashinfer_trtllm as MoE runner backend on sm100 for Qwen3NextForCausalLM"
|
||||||
)
|
)
|
||||||
|
elif model_arch in [
|
||||||
|
"NemotronHForCausalLM",
|
||||||
|
"FalconH1ForCausalLM",
|
||||||
|
"JetNemotronForCausalLM",
|
||||||
|
"JetVLMForConditionalGeneration",
|
||||||
|
]:
|
||||||
|
if not self.disable_radix_cache:
|
||||||
|
logger.warning(
|
||||||
|
"Disabling overlap schedule since MambaRadixCache is not compatible with "
|
||||||
|
"overlap schedule currently, try to use --disable-radix-cache if overlap schedule is necessary"
|
||||||
|
)
|
||||||
|
self.disable_overlap_schedule = True
|
||||||
|
|
||||||
def _handle_sampling_backend(self):
|
def _handle_sampling_backend(self):
|
||||||
if self.sampling_backend is None:
|
if self.sampling_backend is None:
|
||||||
|
|||||||
Reference in New Issue
Block a user