Add Ling_2_6 (#23837)

Signed-off-by: vito.yy <vito.yy@antgroup.com>
Co-authored-by: 得泽 <zhangkaihong.zkh@antgroup.com>
This commit is contained in:
ant-yy
2026-05-27 14:57:23 +08:00
committed by GitHub
co-authored by 得泽
parent d6032c04b6
commit dea85c30f4
12 changed files with 812 additions and 67 deletions
@@ -823,6 +823,11 @@ class HybridLinearAttnBackend(AttentionBackend):
return layer_id in self.full_attn_layers
def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_draft_extend_v2():
# DRAFT_EXTEND_V2 only runs full-attn layers in the draft model,
# so skip linear/mamba backend metadata which requires query_start_loc.
self.full_attn_backend.init_forward_metadata(forward_batch)
return
for attn_backend in self.attn_backend_list:
attn_backend.init_forward_metadata(forward_batch)
@@ -38,6 +38,10 @@ class LightningAttentionBackend(MambaAttnBackendBase):
def __init__(self, model_runner: ModelRunner):
super().__init__(model_runner)
# lightning attn does not need conv cache, but to keep the interface for mamba cache
self.conv_states_shape = (
model_runner.req_to_token_pool.mamba_pool.mamba_cache.conv[0].shape
)
assert not (
model_runner.sliding_window_size is not None
@@ -279,8 +283,6 @@ class LightningAttentionBackend(MambaAttnBackendBase):
save_kv_cache=True,
**kwargs,
):
q_rope = kwargs["q_rope"] if "q_rope" in kwargs else None
k_rope = kwargs["k_rope"] if "k_rope" in kwargs else None
layer_id = layer.layer_id if layer else kwargs["layer_id"]
metadata = self.forward_metadata
@@ -288,7 +290,6 @@ class LightningAttentionBackend(MambaAttnBackendBase):
if self.kv_cache_dtype_str != "auto" and layer.k_scale is not None:
q = q.to(self.kv_cache_dtype)
query_start_loc = self.forward_metadata.query_start_loc
cache_indices = self.forward_metadata.mamba_cache_indices
mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id)
ssm_states = mamba_cache_params.temporal
@@ -332,6 +333,18 @@ class LightningAttentionBackend(MambaAttnBackendBase):
raise ValueError(
f"linear backend: {self.linear_backend} is not support for now"
)
if (
not forward_batch.forward_mode.is_target_verify()
and forward_batch.mamba_track_mask is not None
):
# save mamba cache for extra buffer
mamba_track_mask = forward_batch.mamba_track_mask
mamba_track_indices = forward_batch.mamba_track_indices
dst_masked = mamba_track_indices[mamba_track_mask]
src_masked = metadata.mamba_cache_indices[mamba_track_mask]
ssm_states[dst_masked] = ssm_states[src_masked]
return o.view(-1, layer.tp_q_head_num * layer.v_head_dim)
def forward_decode(
@@ -344,8 +357,6 @@ class LightningAttentionBackend(MambaAttnBackendBase):
save_kv_cache=True,
**kwargs,
) -> torch.Tensor:
q_rope = kwargs["q_rope"] if "q_rope" in kwargs else None
k_rope = kwargs["k_rope"] if "k_rope" in kwargs else None
layer_id = layer.layer_id if layer else kwargs["layer_id"]
# Use precomputed metadata across all layers
@@ -355,7 +366,6 @@ class LightningAttentionBackend(MambaAttnBackendBase):
q = q.to(self.kv_cache_dtype)
# Do linear attention
query_start_loc = self.forward_metadata.query_start_loc
cache_indices = self.forward_metadata.mamba_cache_indices
mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id)
ssm_states = mamba_cache_params.temporal
@@ -0,0 +1,146 @@
{
"1": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"2": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 5
},
"4": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"8": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3
},
"16": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"24": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3
},
"32": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3
},
"48": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"64": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3
},
"96": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 32,
"num_warps": 4,
"num_stages": 3
},
"128": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 3
},
"256": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 5
},
"512": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"1024": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"1536": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 32,
"num_warps": 4,
"num_stages": 5
},
"2048": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"3072": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 32,
"num_warps": 4,
"num_stages": 4
},
"4096": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
}
}
@@ -0,0 +1,164 @@
{
"1": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"2": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"4": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"8": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"16": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"24": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"32": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"48": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"64": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 16,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"96": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"128": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"256": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"512": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 256,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 32,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"1024": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"1536": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"2048": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"3072": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4,
"USE_TMA": true
},
"4096": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
}
}
@@ -0,0 +1,146 @@
{
"1": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"2": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 5
},
"4": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"8": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3
},
"16": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"24": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3
},
"32": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3
},
"48": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"64": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3
},
"96": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 32,
"num_warps": 4,
"num_stages": 3
},
"128": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 3
},
"256": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 5
},
"512": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"1024": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"1536": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 32,
"num_warps": 4,
"num_stages": 5
},
"2048": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"3072": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 32,
"num_warps": 4,
"num_stages": 4
},
"4096": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
}
}
@@ -0,0 +1,164 @@
{
"1": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"2": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"4": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"8": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"16": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"24": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"32": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"48": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"64": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 16,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"96": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"128": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"256": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"512": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 256,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 32,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"1024": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"1536": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"2048": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"3072": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4,
"USE_TMA": true
},
"4096": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
}
}
@@ -159,6 +159,8 @@ def build_kv_cache(
tp_worker.model_runner.hybrid_gdn_config is not None
or tp_worker.model_runner.mamba2_config is not None
or _registry_needs_mamba
or tp_worker.model_runner.kimi_linear_config is not None
or tp_worker.model_runner.hybrid_lightning_config is not None
)
sliding_window_size = None
@@ -118,9 +118,10 @@ class ForwardBatchDeepSeekMHAMixin:
MLATokenToKVPool,
)
assert isinstance(get_token_to_kv_pool(), MLATokenToKVPool) or (
isinstance(get_token_to_kv_pool(), HybridLinearKVPool)
and isinstance(get_token_to_kv_pool().full_kv_pool, MLATokenToKVPool)
token_to_kv_pool = get_token_to_kv_pool()
assert isinstance(token_to_kv_pool, MLATokenToKVPool) or (
isinstance(token_to_kv_pool, HybridLinearKVPool)
and isinstance(token_to_kv_pool.full_kv_pool, MLATokenToKVPool)
), "Currently chunked prefix cache can only be used by Deepseek models"
if not any(self.extend_prefix_lens_cpu):
@@ -80,28 +80,30 @@ class ModelRunnerKVCacheMixin:
server_args = self.server_args
assert config is not None
# reserve the memory for the intermediate mamba states used for spec dec
if not self.spec_algorithm.is_none():
has_spec_dec = not self.spec_algorithm.is_none()
if has_spec_dec:
assert server_args.speculative_num_draft_tokens is not None
assert server_args.max_running_requests is not None
max_running_requests = server_args.max_running_requests // (
self.dp_size if server_args.enable_dp_attention else 1
)
mamba_state_intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req
* max_running_requests
* server_args.speculative_num_draft_tokens
)
total_rest_memory = total_rest_memory - (
mamba_state_intermediate_size / (1 << 30)
)
if server_args.max_mamba_cache_size is not None:
# Use explicitly set max_mamba_cache_size
server_args.max_mamba_cache_size = server_args.max_mamba_cache_size // (
server_args.dp_size if server_args.enable_dp_attention else 1
)
# Reserve intermediate memory based on capped max_num_reqs
if has_spec_dec:
ratio = self._calculate_mamba_ratio()
capped_reqs = min(
server_args.max_running_requests
// (self.dp_size if server_args.enable_dp_attention else 1),
server_args.max_mamba_cache_size // ratio,
)
intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req
* capped_reqs
* server_args.speculative_num_draft_tokens
)
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
elif (
server_args.disable_radix_cache
and server_args.max_running_requests is not None
@@ -110,23 +112,64 @@ class ModelRunnerKVCacheMixin:
server_args.max_mamba_cache_size = server_args.max_running_requests // (
server_args.dp_size if server_args.enable_dp_attention else 1
)
# Reserve intermediate memory based on capped max_num_reqs
if has_spec_dec:
intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req
* server_args.max_mamba_cache_size
* server_args.speculative_num_draft_tokens
)
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
else:
# Use ratio-based calculation to auto-fit available memory
assert config.mamba2_cache_params.mamba_cache_per_req > 0
per_req = config.mamba2_cache_params.mamba_cache_per_req
# allocate the memory based on the ratio between mamba state memory vs. full kv cache memory
# solve the equations:
# 1. mamba_state_memory + full_kv_cache_memory == total_rest_memory
# 2. mamba_state_memory / full_kv_cache_memory == server_args.mamba_full_memory_ratio
mamba_state_memory_raw = (
# Solve jointly for max_mamba_cache_size accounting for intermediate memory.
# The mamba budget (from the ratio split) must cover both:
# 1. main mamba state: max_mamba_cache_size * per_req
# 2. intermediate states: (max_mamba_cache_size / ratio) * D * per_req
# So: max_mamba_cache_size * per_req * (1 + D/ratio) = mamba_budget_bytes
mamba_budget = (
total_rest_memory
* server_args.mamba_full_memory_ratio
/ (1 + server_args.mamba_full_memory_ratio)
)
# calculate the max_mamba_cache_size based on the given total mamba memory
server_args.max_mamba_cache_size = int(
(mamba_state_memory_raw * (1 << 30))
// config.mamba2_cache_params.mamba_cache_per_req
mamba_budget_bytes = mamba_budget * (1 << 30)
if has_spec_dec:
ratio = self._calculate_mamba_ratio()
D = server_args.speculative_num_draft_tokens
# Joint solve: main_state + intermediate = mamba_budget
server_args.max_mamba_cache_size = int(
mamba_budget_bytes // (per_req * (1 + D / ratio))
)
# Intermediate memory is included in mamba_budget, subtract it
# so the return value only has main_state subtracted from total
capped_reqs = min(
server_args.max_running_requests
// (self.dp_size if server_args.enable_dp_attention else 1),
server_args.max_mamba_cache_size // ratio,
)
intermediate_size = per_req * capped_reqs * D
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
else:
server_args.max_mamba_cache_size = int(mamba_budget_bytes // per_req)
# Validate: max_mamba_cache_size must be positive after memory allocation.
# A non-positive value means GPU memory is insufficient for the requested
# configuration. Fail fast with actionable advice instead of silently
# producing garbled output at runtime.
if server_args.max_mamba_cache_size <= 0:
raise RuntimeError(
f"Not enough GPU memory for hybrid (mamba/linear-attention) state cache. "
f"Computed max_mamba_cache_size={server_args.max_mamba_cache_size} "
f"(total_rest_memory={total_rest_memory:.2f} GB, "
f"mamba_cache_per_req={config.mamba2_cache_params.mamba_cache_per_req / (1 << 20):.2f} MB). "
f"Try: (1) reduce --max-running-requests, "
f"(2) increase --mem-fraction-static, "
f"(3) reduce --speculative-num-draft-tokens, or "
f"(4) use GPUs with more memory."
)
mamba_state_memory = (
@@ -836,6 +879,19 @@ class ModelRunnerKVCacheMixin:
max_num_reqs, self.server_args.max_mamba_cache_size // ratio
)
if max_num_reqs <= 0:
raise RuntimeError(
f"Hybrid (mamba/linear-attention) state cache is too small to serve "
f"any requests. max_mamba_cache_size={self.server_args.max_mamba_cache_size}, "
f"mamba_ratio={ratio}, resulting max_num_reqs={max_num_reqs}. "
f"Try: (1) reduce --max-running-requests, "
f"(2) increase --mem-fraction-static, or "
f"(3) use GPUs with more memory."
)
logger.info(
f"Max concurrent requests (per dp worker) from the finalized token capacity: "
f"max_num_reqs={max_num_reqs}."
)
return max_num_reqs
def _apply_memory_pool_config(self: ModelRunner, config: MemoryPoolConfig):
+78 -34
View File
@@ -57,6 +57,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
ParallelLMHead,
VocabParallelEmbedding,
)
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
@@ -243,9 +244,11 @@ class BailingMoE(nn.Module):
quant_config: Optional[QuantizationConfig] = None,
layer_id: int = 0,
prefix: str = "moe",
alt_stream=None,
):
super().__init__()
self.alt_stream = alt_stream
self.layer_id = layer_id
self.tp_size = get_tensor_model_parallel_world_size()
@@ -338,15 +341,35 @@ class BailingMoE(nn.Module):
) -> torch.Tensor:
num_tokens, hidden_size = hidden_states.shape
hidden_states = hidden_states.view(-1, hidden_size)
if self.num_shared_experts > 0:
shared_output = self.shared_experts(hidden_states)
router_logits = self.gate(hidden_states)
topk_output = self.topk(hidden_states, router_logits)
final_hidden_states = self.experts(hidden_states, topk_output)
if (
self.alt_stream is not None
and self.num_shared_experts > 0
and hidden_states.shape[0] > 0
and get_is_capture_mode()
):
with torch.no_grad():
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
# Main stream: shared experts (smaller computation)
shared_output = self.shared_experts(hidden_states)
# Alt stream: gate + topk + routed experts
with torch.cuda.stream(self.alt_stream):
router_logits = self.gate(hidden_states)
topk_output = self.topk(hidden_states, router_logits)
final_hidden_states = self.experts(hidden_states, topk_output)
current_stream.wait_stream(self.alt_stream)
final_hidden_states = final_hidden_states + shared_output
else:
if self.num_shared_experts > 0:
shared_output = self.shared_experts(hidden_states)
if self.num_shared_experts > 0:
final_hidden_states = final_hidden_states + shared_output
router_logits = self.gate(hidden_states)
topk_output = self.topk(hidden_states, router_logits)
final_hidden_states = self.experts(hidden_states, topk_output)
if self.num_shared_experts > 0:
final_hidden_states = final_hidden_states + shared_output
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=True,
@@ -398,9 +421,11 @@ class BailingMoELinearAttention(nn.Module):
quant_config: Optional[QuantizationConfig] = None,
layer_id: int = 0,
prefix: str = "linear_attn",
alt_stream=None,
):
super().__init__()
self.alt_stream = alt_stream
self.layer_id = layer_id
self.hidden_size = config.hidden_size
self.total_num_heads = config.num_attention_heads
@@ -530,8 +555,6 @@ class BailingMoELinearAttention(nn.Module):
**kwargs,
) -> torch.Tensor:
qkv, _ = self.query_key_value(hidden_states)
# logger.warning(f"===={self.layer_id=}, 1-1 {qkv.shape=}")
# use rotary_emb support fp32
qkv = qkv.to(torch.float32)
if self.linear_silu:
qkv = F.silu(qkv)
@@ -544,20 +567,40 @@ class BailingMoELinearAttention(nn.Module):
if self.use_qk_norm:
q = q.reshape(-1, self.tp_heads, self.head_dim)
k = k.reshape(-1, self.tp_kv_heads, self.head_dim)
q = layernorm_fn(
q,
self.query_layernorm.weight.data,
bias=None,
eps=self.rms_norm_eps,
is_rms_norm=True,
)
k = layernorm_fn(
k,
self.key_layernorm.weight.data,
bias=None,
eps=self.rms_norm_eps,
is_rms_norm=True,
)
if self.alt_stream is not None and get_is_capture_mode():
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
q = layernorm_fn(
q,
self.query_layernorm.weight.data,
bias=None,
eps=self.rms_norm_eps,
is_rms_norm=True,
)
with torch.cuda.stream(self.alt_stream):
k = layernorm_fn(
k,
self.key_layernorm.weight.data,
bias=None,
eps=self.rms_norm_eps,
is_rms_norm=True,
)
current_stream.wait_stream(self.alt_stream)
else:
q = layernorm_fn(
q,
self.query_layernorm.weight.data,
bias=None,
eps=self.rms_norm_eps,
is_rms_norm=True,
)
k = layernorm_fn(
k,
self.key_layernorm.weight.data,
bias=None,
eps=self.rms_norm_eps,
is_rms_norm=True,
)
q = q.reshape(-1, self.q_size_per_rank)
k = k.reshape(-1, self.kv_size_per_rank)
@@ -571,23 +614,18 @@ class BailingMoELinearAttention(nn.Module):
if self.linear_scale:
q = q * self.scaling
# q = q.to(torch.float32)
# k = k.to(torch.float32)
# v = v.to(torch.float32)
hidden = self.attn(q, k, v, forward_batch).to(hidden_states.dtype)
gate, _ = self.g_proj(hidden_states)
# logger.warning(
# f"===={self.layer_id=}, 1-3 {gate.shape=}, {hidden.shape=}, {gate.dtype=}, {hidden_states.dtype=}, {hidden.dtype=}"
# )
if self.group_norm_size > 1:
hidden = self.g_norm(hidden, gate)
else:
hidden = self.g_norm(hidden)
hidden = F.sigmoid(gate) * hidden
# logger.warning(f"===={self.layer_id=}, 1-4 {hidden.shape=}")
hidden = hidden.data.to(hidden_states.dtype)
hidden, _ = self.dense(hidden)
# logger.warning(f"===={self.layer_id=}, 1-5 {hidden.shape=}")
return hidden
@@ -709,12 +747,11 @@ class BailingMoELinearDecoderLayer(nn.Module):
layer_id: int = 0,
prefix: str = "layer",
is_nextn: bool = False,
alt_stream=None,
) -> None:
super().__init__()
self.layer_id = layer_id
self.use_mla = getattr(config, "full_attention_type", "mla") == "mla"
alt_stream = None # tptest
# todo nextn
if config.attention_type == 0: # Linear layer
self.attention = BailingMoELinearAttention(
@@ -722,6 +759,7 @@ class BailingMoELinearDecoderLayer(nn.Module):
quant_config=quant_config,
layer_id=self.layer_id,
prefix=prefix + ".attention",
alt_stream=alt_stream,
)
elif config.attention_type == 1: # softmax layer
if self.use_mla:
@@ -776,6 +814,7 @@ class BailingMoELinearDecoderLayer(nn.Module):
quant_config=quant_config,
layer_id=self.layer_id,
prefix=add_prefix("mlp", prefix),
alt_stream=alt_stream,
)
else:
# dense layer
@@ -921,6 +960,8 @@ class BailingMoELinearModel(nn.Module):
else:
self.word_embeddings = PPMissingLayer()
self.alt_stream = torch.cuda.Stream() if _is_cuda else None
def layer_fn(idx, prefix):
layer_idx = idx
layer_config = copy.deepcopy(config)
@@ -928,7 +969,10 @@ class BailingMoELinearModel(nn.Module):
decoder_kwargs = {"quant_config": quant_config, "layer_id": layer_idx}
return BailingMoELinearDecoderLayer(
layer_config, **decoder_kwargs, prefix=prefix
layer_config,
**decoder_kwargs,
prefix=prefix,
alt_stream=self.alt_stream,
)
self.layers, self.start_layer, self.end_layer = make_layers(
+7 -1
View File
@@ -2345,11 +2345,17 @@ class ServerArgs:
logger.info(
f"Using {self.attention_backend} as attention backend for {model_arch}."
)
elif model_arch in ["KimiLinearForCausalLM", "BailingMoeV2_5ForCausalLM"]:
elif model_arch in ["KimiLinearForCausalLM"]:
self._handle_mamba_radix_cache(
model_arch=model_arch,
support_mamba_cache=False,
)
elif model_arch in ["BailingMoeV2_5ForCausalLM"]:
self._handle_mamba_radix_cache(
model_arch=model_arch,
support_mamba_cache=True,
support_mamba_cache_extra_buffer=True,
)
elif model_arch in ["NemotronHForCausalLM"]:
from sglang.srt.arg_groups.nemotron_h_hook import (
apply_nemotron_h_defaults,
@@ -1112,6 +1112,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
if (
self.target_worker.model_runner.hybrid_gdn_config is not None
or self.target_worker.model_runner.mamba2_config is not None
or self.target_worker.model_runner.hybrid_lightning_config is not None
):
self._mamba_verify_update(
batch, verify_input, accept_lens, accept_index, bs