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:
@@ -823,6 +823,11 @@ class HybridLinearAttnBackend(AttentionBackend):
|
|||||||
return layer_id in self.full_attn_layers
|
return layer_id in self.full_attn_layers
|
||||||
|
|
||||||
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
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:
|
for attn_backend in self.attn_backend_list:
|
||||||
attn_backend.init_forward_metadata(forward_batch)
|
attn_backend.init_forward_metadata(forward_batch)
|
||||||
|
|
||||||
|
|||||||
@@ -38,6 +38,10 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
|
|
||||||
def __init__(self, model_runner: ModelRunner):
|
def __init__(self, model_runner: ModelRunner):
|
||||||
super().__init__(model_runner)
|
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 (
|
assert not (
|
||||||
model_runner.sliding_window_size is not None
|
model_runner.sliding_window_size is not None
|
||||||
@@ -279,8 +283,6 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
save_kv_cache=True,
|
save_kv_cache=True,
|
||||||
**kwargs,
|
**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"]
|
layer_id = layer.layer_id if layer else kwargs["layer_id"]
|
||||||
|
|
||||||
metadata = self.forward_metadata
|
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:
|
if self.kv_cache_dtype_str != "auto" and layer.k_scale is not None:
|
||||||
q = q.to(self.kv_cache_dtype)
|
q = q.to(self.kv_cache_dtype)
|
||||||
|
|
||||||
query_start_loc = self.forward_metadata.query_start_loc
|
|
||||||
cache_indices = self.forward_metadata.mamba_cache_indices
|
cache_indices = self.forward_metadata.mamba_cache_indices
|
||||||
mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id)
|
mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id)
|
||||||
ssm_states = mamba_cache_params.temporal
|
ssm_states = mamba_cache_params.temporal
|
||||||
@@ -332,6 +333,18 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"linear backend: {self.linear_backend} is not support for now"
|
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)
|
return o.view(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||||
|
|
||||||
def forward_decode(
|
def forward_decode(
|
||||||
@@ -344,8 +357,6 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
save_kv_cache=True,
|
save_kv_cache=True,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> torch.Tensor:
|
) -> 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"]
|
layer_id = layer.layer_id if layer else kwargs["layer_id"]
|
||||||
|
|
||||||
# Use precomputed metadata across all layers
|
# Use precomputed metadata across all layers
|
||||||
@@ -355,7 +366,6 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
q = q.to(self.kv_cache_dtype)
|
q = q.to(self.kv_cache_dtype)
|
||||||
|
|
||||||
# Do linear attention
|
# Do linear attention
|
||||||
query_start_loc = self.forward_metadata.query_start_loc
|
|
||||||
cache_indices = self.forward_metadata.mamba_cache_indices
|
cache_indices = self.forward_metadata.mamba_cache_indices
|
||||||
mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id)
|
mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id)
|
||||||
ssm_states = mamba_cache_params.temporal
|
ssm_states = mamba_cache_params.temporal
|
||||||
|
|||||||
+146
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
+164
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
+146
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
+164
@@ -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
|
tp_worker.model_runner.hybrid_gdn_config is not None
|
||||||
or tp_worker.model_runner.mamba2_config is not None
|
or tp_worker.model_runner.mamba2_config is not None
|
||||||
or _registry_needs_mamba
|
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
|
sliding_window_size = None
|
||||||
|
|||||||
@@ -118,9 +118,10 @@ class ForwardBatchDeepSeekMHAMixin:
|
|||||||
MLATokenToKVPool,
|
MLATokenToKVPool,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert isinstance(get_token_to_kv_pool(), MLATokenToKVPool) or (
|
token_to_kv_pool = get_token_to_kv_pool()
|
||||||
isinstance(get_token_to_kv_pool(), HybridLinearKVPool)
|
assert isinstance(token_to_kv_pool, MLATokenToKVPool) or (
|
||||||
and isinstance(get_token_to_kv_pool().full_kv_pool, MLATokenToKVPool)
|
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"
|
), "Currently chunked prefix cache can only be used by Deepseek models"
|
||||||
|
|
||||||
if not any(self.extend_prefix_lens_cpu):
|
if not any(self.extend_prefix_lens_cpu):
|
||||||
|
|||||||
@@ -80,28 +80,30 @@ class ModelRunnerKVCacheMixin:
|
|||||||
server_args = self.server_args
|
server_args = self.server_args
|
||||||
assert config is not None
|
assert config is not None
|
||||||
|
|
||||||
# reserve the memory for the intermediate mamba states used for spec dec
|
has_spec_dec = not self.spec_algorithm.is_none()
|
||||||
if not self.spec_algorithm.is_none():
|
if has_spec_dec:
|
||||||
assert server_args.speculative_num_draft_tokens is not None
|
assert server_args.speculative_num_draft_tokens is not None
|
||||||
assert server_args.max_running_requests 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:
|
if server_args.max_mamba_cache_size is not None:
|
||||||
# Use explicitly set max_mamba_cache_size
|
# Use explicitly set max_mamba_cache_size
|
||||||
server_args.max_mamba_cache_size = server_args.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
|
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 (
|
elif (
|
||||||
server_args.disable_radix_cache
|
server_args.disable_radix_cache
|
||||||
and server_args.max_running_requests is not None
|
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.max_mamba_cache_size = server_args.max_running_requests // (
|
||||||
server_args.dp_size if server_args.enable_dp_attention else 1
|
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:
|
else:
|
||||||
# Use ratio-based calculation to auto-fit available memory
|
# Use ratio-based calculation to auto-fit available memory
|
||||||
assert config.mamba2_cache_params.mamba_cache_per_req > 0
|
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 jointly for max_mamba_cache_size accounting for intermediate memory.
|
||||||
# solve the equations:
|
# The mamba budget (from the ratio split) must cover both:
|
||||||
# 1. mamba_state_memory + full_kv_cache_memory == total_rest_memory
|
# 1. main mamba state: max_mamba_cache_size * per_req
|
||||||
# 2. mamba_state_memory / full_kv_cache_memory == server_args.mamba_full_memory_ratio
|
# 2. intermediate states: (max_mamba_cache_size / ratio) * D * per_req
|
||||||
mamba_state_memory_raw = (
|
# So: max_mamba_cache_size * per_req * (1 + D/ratio) = mamba_budget_bytes
|
||||||
|
mamba_budget = (
|
||||||
total_rest_memory
|
total_rest_memory
|
||||||
* server_args.mamba_full_memory_ratio
|
* server_args.mamba_full_memory_ratio
|
||||||
/ (1 + 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
|
mamba_budget_bytes = mamba_budget * (1 << 30)
|
||||||
server_args.max_mamba_cache_size = int(
|
|
||||||
(mamba_state_memory_raw * (1 << 30))
|
if has_spec_dec:
|
||||||
// config.mamba2_cache_params.mamba_cache_per_req
|
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 = (
|
mamba_state_memory = (
|
||||||
@@ -836,6 +879,19 @@ class ModelRunnerKVCacheMixin:
|
|||||||
max_num_reqs, self.server_args.max_mamba_cache_size // ratio
|
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
|
return max_num_reqs
|
||||||
|
|
||||||
def _apply_memory_pool_config(self: ModelRunner, config: MemoryPoolConfig):
|
def _apply_memory_pool_config(self: ModelRunner, config: MemoryPoolConfig):
|
||||||
|
|||||||
@@ -57,6 +57,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
ParallelLMHead,
|
ParallelLMHead,
|
||||||
VocabParallelEmbedding,
|
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_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
|
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
|
||||||
@@ -243,9 +244,11 @@ class BailingMoE(nn.Module):
|
|||||||
quant_config: Optional[QuantizationConfig] = None,
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
layer_id: int = 0,
|
layer_id: int = 0,
|
||||||
prefix: str = "moe",
|
prefix: str = "moe",
|
||||||
|
alt_stream=None,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
|
self.alt_stream = alt_stream
|
||||||
self.layer_id = layer_id
|
self.layer_id = layer_id
|
||||||
|
|
||||||
self.tp_size = get_tensor_model_parallel_world_size()
|
self.tp_size = get_tensor_model_parallel_world_size()
|
||||||
@@ -338,15 +341,35 @@ class BailingMoE(nn.Module):
|
|||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
num_tokens, hidden_size = hidden_states.shape
|
num_tokens, hidden_size = hidden_states.shape
|
||||||
hidden_states = hidden_states.view(-1, hidden_size)
|
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)
|
if (
|
||||||
topk_output = self.topk(hidden_states, router_logits)
|
self.alt_stream is not None
|
||||||
final_hidden_states = self.experts(hidden_states, topk_output)
|
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:
|
router_logits = self.gate(hidden_states)
|
||||||
final_hidden_states = final_hidden_states + shared_output
|
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(
|
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
|
||||||
is_tp_path=True,
|
is_tp_path=True,
|
||||||
@@ -398,9 +421,11 @@ class BailingMoELinearAttention(nn.Module):
|
|||||||
quant_config: Optional[QuantizationConfig] = None,
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
layer_id: int = 0,
|
layer_id: int = 0,
|
||||||
prefix: str = "linear_attn",
|
prefix: str = "linear_attn",
|
||||||
|
alt_stream=None,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
|
self.alt_stream = alt_stream
|
||||||
self.layer_id = layer_id
|
self.layer_id = layer_id
|
||||||
self.hidden_size = config.hidden_size
|
self.hidden_size = config.hidden_size
|
||||||
self.total_num_heads = config.num_attention_heads
|
self.total_num_heads = config.num_attention_heads
|
||||||
@@ -530,8 +555,6 @@ class BailingMoELinearAttention(nn.Module):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
qkv, _ = self.query_key_value(hidden_states)
|
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)
|
qkv = qkv.to(torch.float32)
|
||||||
if self.linear_silu:
|
if self.linear_silu:
|
||||||
qkv = F.silu(qkv)
|
qkv = F.silu(qkv)
|
||||||
@@ -544,20 +567,40 @@ class BailingMoELinearAttention(nn.Module):
|
|||||||
if self.use_qk_norm:
|
if self.use_qk_norm:
|
||||||
q = q.reshape(-1, self.tp_heads, self.head_dim)
|
q = q.reshape(-1, self.tp_heads, self.head_dim)
|
||||||
k = k.reshape(-1, self.tp_kv_heads, self.head_dim)
|
k = k.reshape(-1, self.tp_kv_heads, self.head_dim)
|
||||||
q = layernorm_fn(
|
if self.alt_stream is not None and get_is_capture_mode():
|
||||||
q,
|
current_stream = torch.cuda.current_stream()
|
||||||
self.query_layernorm.weight.data,
|
self.alt_stream.wait_stream(current_stream)
|
||||||
bias=None,
|
q = layernorm_fn(
|
||||||
eps=self.rms_norm_eps,
|
q,
|
||||||
is_rms_norm=True,
|
self.query_layernorm.weight.data,
|
||||||
)
|
bias=None,
|
||||||
k = layernorm_fn(
|
eps=self.rms_norm_eps,
|
||||||
k,
|
is_rms_norm=True,
|
||||||
self.key_layernorm.weight.data,
|
)
|
||||||
bias=None,
|
with torch.cuda.stream(self.alt_stream):
|
||||||
eps=self.rms_norm_eps,
|
k = layernorm_fn(
|
||||||
is_rms_norm=True,
|
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)
|
q = q.reshape(-1, self.q_size_per_rank)
|
||||||
k = k.reshape(-1, self.kv_size_per_rank)
|
k = k.reshape(-1, self.kv_size_per_rank)
|
||||||
|
|
||||||
@@ -571,23 +614,18 @@ class BailingMoELinearAttention(nn.Module):
|
|||||||
|
|
||||||
if self.linear_scale:
|
if self.linear_scale:
|
||||||
q = q * self.scaling
|
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)
|
hidden = self.attn(q, k, v, forward_batch).to(hidden_states.dtype)
|
||||||
gate, _ = self.g_proj(hidden_states)
|
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:
|
if self.group_norm_size > 1:
|
||||||
hidden = self.g_norm(hidden, gate)
|
hidden = self.g_norm(hidden, gate)
|
||||||
else:
|
else:
|
||||||
hidden = self.g_norm(hidden)
|
hidden = self.g_norm(hidden)
|
||||||
hidden = F.sigmoid(gate) * hidden
|
hidden = F.sigmoid(gate) * hidden
|
||||||
# logger.warning(f"===={self.layer_id=}, 1-4 {hidden.shape=}")
|
|
||||||
hidden = hidden.data.to(hidden_states.dtype)
|
hidden = hidden.data.to(hidden_states.dtype)
|
||||||
hidden, _ = self.dense(hidden)
|
hidden, _ = self.dense(hidden)
|
||||||
# logger.warning(f"===={self.layer_id=}, 1-5 {hidden.shape=}")
|
|
||||||
return hidden
|
return hidden
|
||||||
|
|
||||||
|
|
||||||
@@ -709,12 +747,11 @@ class BailingMoELinearDecoderLayer(nn.Module):
|
|||||||
layer_id: int = 0,
|
layer_id: int = 0,
|
||||||
prefix: str = "layer",
|
prefix: str = "layer",
|
||||||
is_nextn: bool = False,
|
is_nextn: bool = False,
|
||||||
|
alt_stream=None,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.layer_id = layer_id
|
self.layer_id = layer_id
|
||||||
self.use_mla = getattr(config, "full_attention_type", "mla") == "mla"
|
self.use_mla = getattr(config, "full_attention_type", "mla") == "mla"
|
||||||
alt_stream = None # tptest
|
|
||||||
# todo nextn
|
|
||||||
|
|
||||||
if config.attention_type == 0: # Linear layer
|
if config.attention_type == 0: # Linear layer
|
||||||
self.attention = BailingMoELinearAttention(
|
self.attention = BailingMoELinearAttention(
|
||||||
@@ -722,6 +759,7 @@ class BailingMoELinearDecoderLayer(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
layer_id=self.layer_id,
|
layer_id=self.layer_id,
|
||||||
prefix=prefix + ".attention",
|
prefix=prefix + ".attention",
|
||||||
|
alt_stream=alt_stream,
|
||||||
)
|
)
|
||||||
elif config.attention_type == 1: # softmax layer
|
elif config.attention_type == 1: # softmax layer
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
@@ -776,6 +814,7 @@ class BailingMoELinearDecoderLayer(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
layer_id=self.layer_id,
|
layer_id=self.layer_id,
|
||||||
prefix=add_prefix("mlp", prefix),
|
prefix=add_prefix("mlp", prefix),
|
||||||
|
alt_stream=alt_stream,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# dense layer
|
# dense layer
|
||||||
@@ -921,6 +960,8 @@ class BailingMoELinearModel(nn.Module):
|
|||||||
else:
|
else:
|
||||||
self.word_embeddings = PPMissingLayer()
|
self.word_embeddings = PPMissingLayer()
|
||||||
|
|
||||||
|
self.alt_stream = torch.cuda.Stream() if _is_cuda else None
|
||||||
|
|
||||||
def layer_fn(idx, prefix):
|
def layer_fn(idx, prefix):
|
||||||
layer_idx = idx
|
layer_idx = idx
|
||||||
layer_config = copy.deepcopy(config)
|
layer_config = copy.deepcopy(config)
|
||||||
@@ -928,7 +969,10 @@ class BailingMoELinearModel(nn.Module):
|
|||||||
|
|
||||||
decoder_kwargs = {"quant_config": quant_config, "layer_id": layer_idx}
|
decoder_kwargs = {"quant_config": quant_config, "layer_id": layer_idx}
|
||||||
return BailingMoELinearDecoderLayer(
|
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(
|
self.layers, self.start_layer, self.end_layer = make_layers(
|
||||||
|
|||||||
@@ -2345,11 +2345,17 @@ class ServerArgs:
|
|||||||
logger.info(
|
logger.info(
|
||||||
f"Using {self.attention_backend} as attention backend for {model_arch}."
|
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(
|
self._handle_mamba_radix_cache(
|
||||||
model_arch=model_arch,
|
model_arch=model_arch,
|
||||||
support_mamba_cache=False,
|
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"]:
|
elif model_arch in ["NemotronHForCausalLM"]:
|
||||||
from sglang.srt.arg_groups.nemotron_h_hook import (
|
from sglang.srt.arg_groups.nemotron_h_hook import (
|
||||||
apply_nemotron_h_defaults,
|
apply_nemotron_h_defaults,
|
||||||
|
|||||||
@@ -1112,6 +1112,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
if (
|
if (
|
||||||
self.target_worker.model_runner.hybrid_gdn_config is not None
|
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.mamba2_config is not None
|
||||||
|
or self.target_worker.model_runner.hybrid_lightning_config is not None
|
||||||
):
|
):
|
||||||
self._mamba_verify_update(
|
self._mamba_verify_update(
|
||||||
batch, verify_input, accept_lens, accept_index, bs
|
batch, verify_input, accept_lens, accept_index, bs
|
||||||
|
|||||||
Reference in New Issue
Block a user