[KDA] Enable ReplaySSM for GLM-5.3 Flash (#40517)
This commit is contained in:
@@ -133,6 +133,10 @@ def glm5_next_config(model_config: ModelConfig):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def hybrid_kda_config(model_config: ModelConfig):
|
||||||
|
return kimi_linear_config(model_config) or glm5_next_config(model_config)
|
||||||
|
|
||||||
|
|
||||||
def linear_attn_model_spec(model_config: ModelConfig):
|
def linear_attn_model_spec(model_config: ModelConfig):
|
||||||
result = _get_linear_attn_registry_result(model_config)
|
result = _get_linear_attn_registry_result(model_config)
|
||||||
return result[0] if result else None
|
return result[0] if result else None
|
||||||
@@ -142,8 +146,7 @@ def mambaish_config(model_config: ModelConfig):
|
|||||||
existing = (
|
existing = (
|
||||||
mamba2_config(model_config)
|
mamba2_config(model_config)
|
||||||
or hybrid_gdn_config(model_config)
|
or hybrid_gdn_config(model_config)
|
||||||
or kimi_linear_config(model_config)
|
or hybrid_kda_config(model_config)
|
||||||
or glm5_next_config(model_config)
|
|
||||||
or hybrid_lightning_config(model_config)
|
or hybrid_lightning_config(model_config)
|
||||||
)
|
)
|
||||||
if existing:
|
if existing:
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ import torch
|
|||||||
from sglang.srt.arg_groups.overrides import resolving_view
|
from sglang.srt.arg_groups.overrides import resolving_view
|
||||||
from sglang.srt.configs.hybrid_arch import (
|
from sglang.srt.configs.hybrid_arch import (
|
||||||
hybrid_gdn_config,
|
hybrid_gdn_config,
|
||||||
kimi_linear_config,
|
hybrid_kda_config,
|
||||||
mambaish_config,
|
mambaish_config,
|
||||||
)
|
)
|
||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import (
|
||||||
@@ -285,12 +285,14 @@ class KVCacheConfigurator:
|
|||||||
kv_cache_dtype_str: Optional[str] = None
|
kv_cache_dtype_str: Optional[str] = None
|
||||||
mambaish_config: Optional[Any] = field(init=False)
|
mambaish_config: Optional[Any] = field(init=False)
|
||||||
hybrid_gdn_config: Optional[Any] = field(init=False)
|
hybrid_gdn_config: Optional[Any] = field(init=False)
|
||||||
|
hybrid_kda_config: Optional[Any] = field(init=False)
|
||||||
is_hybrid_swa_mtp_draft: bool = field(init=False)
|
is_hybrid_swa_mtp_draft: bool = field(init=False)
|
||||||
draft_swa_full_capacity: bool = field(init=False)
|
draft_swa_full_capacity: bool = field(init=False)
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
self.mambaish_config = mambaish_config(self.model_config)
|
self.mambaish_config = mambaish_config(self.model_config)
|
||||||
self.hybrid_gdn_config = hybrid_gdn_config(self.model_config)
|
self.hybrid_gdn_config = hybrid_gdn_config(self.model_config)
|
||||||
|
self.hybrid_kda_config = hybrid_kda_config(self.model_config)
|
||||||
self.is_hybrid_swa_mtp_draft = (
|
self.is_hybrid_swa_mtp_draft = (
|
||||||
self.is_draft_worker
|
self.is_draft_worker
|
||||||
and self.draft_model_idx is not None
|
and self.draft_model_idx is not None
|
||||||
@@ -1091,7 +1093,7 @@ class KVCacheConfigurator:
|
|||||||
get_exec().mamba.enable_linear_replayssm_spec
|
get_exec().mamba.enable_linear_replayssm_spec
|
||||||
and (
|
and (
|
||||||
self.hybrid_gdn_config is not None
|
self.hybrid_gdn_config is not None
|
||||||
or kimi_linear_config(self.model_config) is not None
|
or self.hybrid_kda_config is not None
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
@@ -1132,11 +1134,11 @@ class KVCacheConfigurator:
|
|||||||
if (
|
if (
|
||||||
get_exec().mamba.enable_linear_replayssm_spec
|
get_exec().mamba.enable_linear_replayssm_spec
|
||||||
and _algo in ("DSPARK", "DFLASH")
|
and _algo in ("DSPARK", "DFLASH")
|
||||||
and kimi_linear_config(self.model_config) is None
|
and self.hybrid_kda_config is None
|
||||||
):
|
):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--enable-linear-replayssm-spec with DSPARK/DFLASH requires a KDA "
|
"--enable-linear-replayssm-spec with DSPARK/DFLASH requires a KDA "
|
||||||
"(kimi_linear) model; got a non-KDA model."
|
"model; got a non-KDA model."
|
||||||
)
|
)
|
||||||
req_to_token_pool = HybridReqToTokenPool(
|
req_to_token_pool = HybridReqToTokenPool(
|
||||||
size=max_num_reqs,
|
size=max_num_reqs,
|
||||||
@@ -1172,7 +1174,7 @@ class KVCacheConfigurator:
|
|||||||
get_exec().mamba.enable_linear_replayssm_spec
|
get_exec().mamba.enable_linear_replayssm_spec
|
||||||
and (
|
and (
|
||||||
self.hybrid_gdn_config is not None
|
self.hybrid_gdn_config is not None
|
||||||
or kimi_linear_config(self.model_config) is not None
|
or self.hybrid_kda_config is not None
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
@@ -2458,8 +2460,7 @@ class KVCacheConfigurator:
|
|||||||
# The ring is not part of mamba_cache_per_req. GDN replay is fixed-size
|
# The ring is not part of mamba_cache_per_req. GDN replay is fixed-size
|
||||||
# request scratch; KDA replay remains attached to each mamba slot.
|
# request scratch; KDA replay remains attached to each mamba slot.
|
||||||
replayssm_active = get_exec().mamba.enable_linear_replayssm_spec and (
|
replayssm_active = get_exec().mamba.enable_linear_replayssm_spec and (
|
||||||
self.hybrid_gdn_config is not None
|
self.hybrid_gdn_config is not None or self.hybrid_kda_config is not None
|
||||||
or kimi_linear_config(self.model_config) is not None
|
|
||||||
)
|
)
|
||||||
if replayssm_active:
|
if replayssm_active:
|
||||||
record_len = get_exec().mamba.linear_replayssm_cache_len
|
record_len = get_exec().mamba.linear_replayssm_cache_len
|
||||||
@@ -2471,7 +2472,7 @@ class KVCacheConfigurator:
|
|||||||
else:
|
else:
|
||||||
replayssm_ring_per_req = 0
|
replayssm_ring_per_req = 0
|
||||||
replayssm_ring_per_req = int(replayssm_ring_per_req * pp_layer_scale)
|
replayssm_ring_per_req = int(replayssm_ring_per_req * pp_layer_scale)
|
||||||
if replayssm_active and kimi_linear_config(self.model_config) is None:
|
if replayssm_active and self.hybrid_kda_config is None:
|
||||||
replay_req_slots = (
|
replay_req_slots = (
|
||||||
get_schedule().max_running_requests // self.ps.attn_dp_size + 1
|
get_schedule().max_running_requests // self.ps.attn_dp_size + 1
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user