Extract hybrid-arch helpers into configs.hybrid_arch and ModelConfig (#31150)

This commit is contained in:
fzyzcjy
2026-07-14 15:55:32 +08:00
committed by GitHub
parent c9b4081016
commit e20c346541
16 changed files with 219 additions and 178 deletions
+116
View File
@@ -0,0 +1,116 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any
from sglang.srt.configs import (
BailingHybridConfig,
FalconH1Config,
GraniteMoeHybridConfig,
InternS2PreviewConfig,
JetNemotronConfig,
JetVLMConfig,
KimiLinearConfig,
Lfm2Config,
Lfm2MoeConfig,
Lfm2VlConfig,
NemotronH_Nano_VL_V2_Config,
NemotronHConfig,
Qwen3_5Config,
Qwen3_5MoeConfig,
Qwen3NextConfig,
ZayaConfig,
)
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
def _get_linear_attn_registry_result(model_config: ModelConfig) -> Any:
return model_config.linear_attn_registry_result
def qwen3_next_config(model_config: ModelConfig):
config = model_config.hf_config
if isinstance(config, Qwen3NextConfig):
return config
return None
def hybrid_lightning_config(model_config: ModelConfig):
config = model_config.hf_config
if isinstance(config, BailingHybridConfig):
return config
return None
def hybrid_gdn_config(model_config: ModelConfig):
config = model_config.hf_config.get_text_config()
if isinstance(
config,
Qwen3NextConfig
| Qwen3_5Config
| Qwen3_5MoeConfig
| InternS2PreviewConfig
| JetNemotronConfig
| JetVLMConfig,
):
return config
return None
def mamba2_config(model_config: ModelConfig):
config = model_config.hf_config
if isinstance(config, NemotronHConfig) and model_config.is_draft_model:
# NemotronH MTP draft models have no Mamba layers (pattern like "*E")
# so they shouldn't use HybridLinearAttnBackend
pattern = getattr(config, "mtp_hybrid_override_pattern", None)
if pattern is not None and "M" not in pattern:
return None
if isinstance(
config,
FalconH1Config
| NemotronHConfig
| Lfm2Config
| Lfm2MoeConfig
| Lfm2VlConfig
| ZayaConfig,
):
return config
if isinstance(config, NemotronH_Nano_VL_V2_Config):
return config.llm_config
if isinstance(config, GraniteMoeHybridConfig):
has_mamba = any(
layer_type == "mamba" for layer_type in getattr(config, "layer_types", [])
)
if not has_mamba:
return None
else:
return config
return None
def kimi_linear_config(model_config: ModelConfig):
config = model_config.hf_config
if isinstance(config, KimiLinearConfig):
return config
return None
def linear_attn_model_spec(model_config: ModelConfig):
result = _get_linear_attn_registry_result(model_config)
return result[0] if result else None
def mambaish_config(model_config: ModelConfig):
existing = (
mamba2_config(model_config)
or hybrid_gdn_config(model_config)
or kimi_linear_config(model_config)
or hybrid_lightning_config(model_config)
)
if existing:
return existing
result = _get_linear_attn_registry_result(model_config)
return result[1] if result else None
+14
View File
@@ -18,12 +18,14 @@ import logging
import math import math
import os import os
from enum import Enum, IntEnum, auto from enum import Enum, IntEnum, auto
from functools import cached_property
from pathlib import Path from pathlib import Path
from typing import Any, List, Optional, Set, Union from typing import Any, List, Optional, Set, Union
import torch import torch
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_config
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.quantization import QUANTIZATION_METHODS from sglang.srt.layers.quantization import QUANTIZATION_METHODS
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -287,6 +289,14 @@ class ModelConfig:
) )
) )
self.hf_text_config = get_hf_text_config(self.hf_config) self.hf_text_config = get_hf_text_config(self.hf_config)
rope_scaling = getattr(self.hf_text_config, "rope_parameters", None) or getattr(
self.hf_text_config, "rope_scaling", {}
)
self.model_is_mrope = (
rope_scaling is not None and "mrope_section" in rope_scaling
)
self.hf_generation_config = get_generation_config( self.hf_generation_config = get_generation_config(
self.model_path, self.model_path,
trust_remote_code=trust_remote_code, trust_remote_code=trust_remote_code,
@@ -685,6 +695,10 @@ class ModelConfig:
"Gemma4UnifiedForConditionalGeneration", "Gemma4UnifiedForConditionalGeneration",
] ]
@cached_property
def linear_attn_registry_result(self) -> Any:
return get_linear_attn_config(self.hf_config)
def _detect_attention_sinks(self) -> bool: def _detect_attention_sinks(self) -> bool:
"""Check whether the model uses learned attention sinks. """Check whether the model uses learned attention sinks.
@@ -9,6 +9,7 @@ from typing import Tuple
import torch import torch
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import ( from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import (
MlxAuxiliaryStateReqToTokenPool, MlxAuxiliaryStateReqToTokenPool,
) )
@@ -160,7 +161,7 @@ class MlxModelRunnerStub(ModelRunner):
self.is_hybrid_swa = False self.is_hybrid_swa = False
# Create minimal pools # Create minimal pools
if self.mambaish_config is not None: if mambaish_config(self.model_config) is not None:
auxiliary_state_size = self.server_args.max_mamba_cache_size auxiliary_state_size = self.server_args.max_mamba_cache_size
if auxiliary_state_size is None: if auxiliary_state_size is None:
auxiliary_state_size = self.max_running_requests * 4 auxiliary_state_size = self.max_running_requests * 4
@@ -2,6 +2,13 @@ import logging
import warnings import warnings
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from sglang.srt.configs.hybrid_arch import (
hybrid_gdn_config,
hybrid_lightning_config,
kimi_linear_config,
mamba2_config,
mambaish_config,
)
from sglang.srt.configs.linear_attn_model_registry import ( from sglang.srt.configs.linear_attn_model_registry import (
get_linear_attn_config, get_linear_attn_config,
import_backend_class, import_backend_class,
@@ -255,7 +262,7 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
need to change the code of the original attention backend. need to change the code of the original attention backend.
""" """
assert not ( assert not (
runner.hybrid_gdn_config is not None and runner.use_mla_backend hybrid_gdn_config(runner.model_config) is not None and runner.use_mla_backend
), "hybrid_gdn can only be used with non-MLA models." ), "hybrid_gdn can only be used with non-MLA models."
from sglang.srt.configs.model_config import is_minimax_sparse from sglang.srt.configs.model_config import is_minimax_sparse
@@ -271,7 +278,7 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
full_attn_backend, sparse_backend, sparse_backend.sparse_layer_ids full_attn_backend, sparse_backend, sparse_backend.sparse_layer_ids
) )
if cfg := runner.mambaish_config: if cfg := mambaish_config(runner.model_config):
from sglang.srt.layers.attention.fla.utils import check_environments from sglang.srt.layers.attention.fla.utils import check_environments
from sglang.srt.layers.attention.linear.kda_backend import KDAAttnBackend from sglang.srt.layers.attention.linear.kda_backend import KDAAttnBackend
from sglang.srt.layers.attention.linear.lightning_backend import ( from sglang.srt.layers.attention.linear.lightning_backend import (
@@ -303,11 +310,11 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
) )
check_environments() check_environments()
if runner.hybrid_gdn_config is not None and not is_npu(): if hybrid_gdn_config(runner.model_config) is not None and not is_npu():
maybe_set_default_flashinfer_gdn_prefill(runner) maybe_set_default_flashinfer_gdn_prefill(runner)
initialize_linear_attn_config(runner.server_args) initialize_linear_attn_config(runner.server_args)
hybrid_backend_cls = HybridLinearAttnBackend hybrid_backend_cls = HybridLinearAttnBackend
if runner.hybrid_gdn_config is not None: if hybrid_gdn_config(runner.model_config) is not None:
if is_blackwell(): if is_blackwell():
assert ( assert (
runner.server_args.attention_backend == "triton" runner.server_args.attention_backend == "triton"
@@ -321,7 +328,7 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
), "ascend backend is the only supported backend on NPU for hybrid GDN models, use --attention-backend ascend to specify the backend." ), "ascend backend is the only supported backend on NPU for hybrid GDN models, use --attention-backend ascend to specify the backend."
logger.info(f"Using hybrid linear attention backend for hybrid GDN models.") logger.info(f"Using hybrid linear attention backend for hybrid GDN models.")
linear_attn_backend = GDNAttnBackend(runner) linear_attn_backend = GDNAttnBackend(runner)
elif runner.mamba2_config is not None: elif mamba2_config(runner.model_config) is not None:
from sglang.srt.configs.lfm2 import Lfm2Config from sglang.srt.configs.lfm2 import Lfm2Config
from sglang.srt.configs.lfm2_moe import Lfm2MoeConfig from sglang.srt.configs.lfm2_moe import Lfm2MoeConfig
from sglang.srt.configs.lfm2_vl import Lfm2VlConfig from sglang.srt.configs.lfm2_vl import Lfm2VlConfig
@@ -337,7 +344,7 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
Lfm2MoeConfig, Lfm2MoeConfig,
Lfm2VlConfig, Lfm2VlConfig,
) )
if isinstance(runner.mamba2_config, short_conv_cfgs): if isinstance(mamba2_config(runner.model_config), short_conv_cfgs):
if is_npu(): if is_npu():
# The model conv layers call # The model conv layers call
# get_attn_backend().conv_state_metadata() unconditionally, # get_attn_backend().conv_state_metadata() unconditionally,
@@ -362,9 +369,9 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
hybrid_backend_cls = ShortConvHybridAttnBackend hybrid_backend_cls = ShortConvHybridAttnBackend
else: else:
linear_attn_backend = Mamba2AttnBackend(runner) linear_attn_backend = Mamba2AttnBackend(runner)
elif runner.kimi_linear_config is not None: elif kimi_linear_config(runner.model_config) is not None:
linear_attn_backend = KDAAttnBackend(runner) linear_attn_backend = KDAAttnBackend(runner)
elif runner.hybrid_lightning_config is not None: elif hybrid_lightning_config(runner.model_config) is not None:
linear_attn_backend = LightningAttentionBackend(runner) linear_attn_backend = LightningAttentionBackend(runner)
else: else:
spec_result = get_linear_attn_config(runner.model_config.hf_config) spec_result = get_linear_attn_config(runner.model_config.hf_config)
@@ -3,6 +3,7 @@ from typing import Optional, Union
import torch import torch
from sglang.srt.configs.hybrid_arch import mamba2_config
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.mamba.causal_conv1d_triton import PAD_SLOT_ID from sglang.srt.layers.attention.mamba.causal_conv1d_triton import PAD_SLOT_ID
from sglang.srt.layers.attention.mamba.mamba import MambaMixer2 from sglang.srt.layers.attention.mamba.mamba import MambaMixer2
@@ -684,7 +685,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
def __init__(self, model_runner: ModelRunner): def __init__(self, model_runner: ModelRunner):
super().__init__(model_runner) super().__init__(model_runner)
config = model_runner.mamba2_config config = mamba2_config(model_runner.model_config)
assert config is not None assert config is not None
self.mamba_chunk_size = config.mamba_chunk_size self.mamba_chunk_size = config.mamba_chunk_size
self.conv_states_shape = ( self.conv_states_shape = (
@@ -2,6 +2,7 @@ from typing import Optional, Tuple, Union
import torch import torch
from sglang.srt.configs.hybrid_arch import hybrid_gdn_config
from sglang.srt.layers.attention.fla.fused_gdn_gating import fused_gdn_gating from sglang.srt.layers.attention.fla.fused_gdn_gating import fused_gdn_gating
from sglang.srt.layers.attention.hybrid_linear_attn_backend import MambaAttnBackendBase from sglang.srt.layers.attention.hybrid_linear_attn_backend import MambaAttnBackendBase
from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel
@@ -73,7 +74,7 @@ def maybe_set_default_flashinfer_gdn_prefill(model_runner: ModelRunner) -> None:
cuda_version = torch.version.cuda cuda_version = torch.version.cuda
chunk_size = args.chunked_prefill_size chunk_size = args.chunked_prefill_size
config = model_runner.hybrid_gdn_config config = hybrid_gdn_config(model_runner.model_config)
if ( if (
cuda_version is None cuda_version is None
or int(cuda_version.split(".", 1)[0]) < 13 or int(cuda_version.split(".", 1)[0]) < 13
@@ -10,6 +10,11 @@ from sglang.kernels.ops.attention.metadata import get_num_kv_splits_triton
from sglang.kernels.ops.kvcache.kv_indices import ( from sglang.kernels.ops.kvcache.kv_indices import (
create_flashinfer_kv_indices_triton, create_flashinfer_kv_indices_triton,
) )
from sglang.srt.configs.hybrid_arch import (
hybrid_gdn_config,
kimi_linear_config,
linear_attn_model_spec,
)
from sglang.srt.configs.model_config import AttentionArch from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory, use_symmetric_memory,
@@ -182,9 +187,9 @@ class TritonAttnBackend(AttentionBackend):
self.v_head_dim = full_v_head_dim self.v_head_dim = full_v_head_dim
self.swa_v_head_dim = swa_v_head_dim self.swa_v_head_dim = swa_v_head_dim
elif ( elif (
model_runner.hybrid_gdn_config is not None hybrid_gdn_config(model_runner.model_config) is not None
or model_runner.kimi_linear_config is not None or kimi_linear_config(model_runner.model_config) is not None
or model_runner.linear_attn_model_spec is not None or linear_attn_model_spec(model_runner.model_config) is not None
): ):
# For hybrid linear models, layer_id = 0 may not be full attention # For hybrid linear models, layer_id = 0 may not be full attention
self.v_head_dim = model_runner.token_to_kv_pool.get_v_head_dim() self.v_head_dim = model_runner.token_to_kv_pool.get_v_head_dim()
@@ -23,6 +23,13 @@ class KVCacheBuildResult:
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from sglang.srt.configs.hybrid_arch import (
hybrid_gdn_config,
hybrid_lightning_config,
kimi_linear_config,
linear_attn_model_spec,
mamba2_config,
)
from sglang.srt.configs.model_config import ModelImpl, is_deepseek_dsa from sglang.srt.configs.model_config import ModelImpl, is_deepseek_dsa
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.mm_utils import init_mm_embedding_cache from sglang.srt.managers.mm_utils import init_mm_embedding_cache
@@ -146,14 +153,14 @@ def build_kv_cache(
# Hybrid memory pool # Hybrid memory pool
is_hybrid_swa = tp_worker.is_hybrid_swa is_hybrid_swa = tp_worker.is_hybrid_swa
_spec = tp_worker.model_runner.linear_attn_model_spec _spec = linear_attn_model_spec(tp_worker.model_runner.model_config)
_registry_needs_mamba = _spec.uses_mamba_radix_cache if _spec is not None else False _registry_needs_mamba = _spec.uses_mamba_radix_cache if _spec is not None else False
is_hybrid_ssm = ( is_hybrid_ssm = (
tp_worker.model_runner.hybrid_gdn_config is not None hybrid_gdn_config(tp_worker.model_runner.model_config) is not None
or tp_worker.model_runner.mamba2_config is not None or mamba2_config(tp_worker.model_runner.model_config) is not None
or _registry_needs_mamba or _registry_needs_mamba
or tp_worker.model_runner.kimi_linear_config is not None or kimi_linear_config(tp_worker.model_runner.model_config) is not None
or tp_worker.model_runner.hybrid_lightning_config is not None or hybrid_lightning_config(tp_worker.model_runner.model_config) is not None
) )
is_dsa = is_deepseek_dsa(model_config.hf_config) is_dsa = is_deepseek_dsa(model_config.hf_config)
@@ -37,6 +37,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Set, Tuple, Union
import torch import torch
from sglang.kernels.ops.attention.position import compute_position_triton from sglang.kernels.ops.attention.position import compute_position_triton
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.kv_canary.req_to_expected_token_ids_manager import ( from sglang.srt.kv_canary.req_to_expected_token_ids_manager import (
compute_req_all_ids_info, compute_req_all_ids_info,
@@ -850,7 +851,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
if model_runner.use_ngram_embedding: if model_runner.use_ngram_embedding:
ret._init_ngram_embedding_info(batch, device) ret._init_ngram_embedding_info(batch, device)
if model_runner.model_is_mrope: if model_runner.model_config.model_is_mrope:
if ( if (
ret.spec_info is not None ret.spec_info is not None
and getattr(ret.spec_info, "positions", None) is not None and getattr(ret.spec_info, "positions", None) is not None
@@ -1229,7 +1230,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# Mamba-hybrid families need the fabricated-row idle conversion # Mamba-hybrid families need the fabricated-row idle conversion
# below; this includes their MTP draft workers, whose mamba-less # below; this includes their MTP draft workers, whose mamba-less
# "*E" pattern makes mambaish_config return None. # "*E" pattern makes mambaish_config return None.
hybrid_ssm = model_runner.mambaish_config is not None or ( hybrid_ssm = mambaish_config(model_runner.model_config) is not None or (
model_runner.is_draft_worker model_runner.is_draft_worker
and getattr( and getattr(
model_runner.model_config.hf_config, model_runner.model_config.hf_config,
@@ -25,31 +25,12 @@ import threading
import time import time
from collections import defaultdict from collections import defaultdict
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Optional, Union from typing import Optional, Union
import torch import torch
import torch.distributed as dist import torch.distributed as dist
from sglang.srt.configs import (
BailingHybridConfig,
FalconH1Config,
GraniteMoeHybridConfig,
InternS2PreviewConfig,
JetNemotronConfig,
JetVLMConfig,
KimiLinearConfig,
Lfm2Config,
Lfm2MoeConfig,
Lfm2VlConfig,
NemotronH_Nano_VL_V2_Config,
NemotronHConfig,
Qwen3_5Config,
Qwen3_5MoeConfig,
Qwen3NextConfig,
ZayaConfig,
)
from sglang.srt.configs.device_config import DeviceConfig from sglang.srt.configs.device_config import DeviceConfig
from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_config
from sglang.srt.configs.load_config import LoadConfig, LoadFormat from sglang.srt.configs.load_config import LoadConfig, LoadFormat
from sglang.srt.configs.model_config import ( from sglang.srt.configs.model_config import (
AttentionArch, AttentionArch,
@@ -255,8 +236,6 @@ UNBALANCED_MODEL_LOADING_TIMEOUT_S = 480 # leave more time for post data proces
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_UNSET: Any = object()
@dataclass @dataclass
class ModelRunnerOutput: class ModelRunnerOutput:
@@ -334,17 +313,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.req_to_token_pool = req_to_token_pool self.req_to_token_pool = req_to_token_pool
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
self.is_hybrid_swa = model_config.is_hybrid_swa self.is_hybrid_swa = model_config.is_hybrid_swa
self.is_hybrid_swa_compress = getattr( self.is_hybrid_swa_compress = model_config.is_hybrid_swa_compress
model_config, "is_hybrid_swa_compress", False
)
self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA
self.attention_chunk_size = model_config.attention_chunk_size self.attention_chunk_size = model_config.attention_chunk_size
rope_scaling = getattr(
model_config.hf_text_config, "rope_parameters", None
) or getattr(model_config.hf_text_config, "rope_scaling", {})
self.model_is_mrope = (
rope_scaling is not None and "mrope_section" in rope_scaling
)
self.enable_elastic_ep = server_args.elastic_ep_backend is not None self.enable_elastic_ep = server_args.elastic_ep_backend is not None
self.forward_pass_id = 0 self.forward_pass_id = 0
self.init_new_workspace = False self.init_new_workspace = False
@@ -535,8 +506,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
# For hisparse (must be set before initialize() so CUDA graph capture can see it) # For hisparse (must be set before initialize() so CUDA graph capture can see it)
self.hisparse_coordinator = None self.hisparse_coordinator = None
self._linear_attn_registry_cache: Any = _UNSET
# Load model weights and configure # Load model weights and configure
self.initialize() self.initialize()
self.check_quantized_moe_compatibility() self.check_quantized_moe_compatibility()
@@ -1796,69 +1765,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
return result return result
@property
def qwen3_next_config(self):
config = self.model_config.hf_config
if isinstance(config, Qwen3NextConfig):
return config
return None
@property
def hybrid_lightning_config(self):
config = self.model_config.hf_config
if isinstance(config, BailingHybridConfig):
return config
return None
@property
def hybrid_gdn_config(self):
config = self.model_config.hf_config.get_text_config()
if isinstance(
config,
Qwen3NextConfig
| Qwen3_5Config
| Qwen3_5MoeConfig
| InternS2PreviewConfig
| JetNemotronConfig
| JetVLMConfig,
):
return config
return None
@property
def mamba2_config(self):
config = self.model_config.hf_config
if isinstance(config, NemotronHConfig) and self.is_draft_worker:
# NemotronH MTP draft models have no Mamba layers (pattern like "*E")
# so they shouldn't use HybridLinearAttnBackend
pattern = getattr(config, "mtp_hybrid_override_pattern", None)
if pattern is not None and "M" not in pattern:
return None
if isinstance(
config,
FalconH1Config
| NemotronHConfig
| Lfm2Config
| Lfm2MoeConfig
| Lfm2VlConfig
| ZayaConfig,
):
return config
if isinstance(config, NemotronH_Nano_VL_V2_Config):
return config.llm_config
if isinstance(config, GraniteMoeHybridConfig):
has_mamba = any(
layer_type == "mamba"
for layer_type in getattr(config, "layer_types", [])
)
if not has_mamba:
return None
else:
return config
return None
@property @property
def effective_max_total_num_tokens(self): def effective_max_total_num_tokens(self):
"""Return the max token pool size considering hybrid swa settings.""" """Return the max token pool size considering hybrid swa settings."""
@@ -1867,38 +1773,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
else: else:
return self.max_total_num_tokens return self.max_total_num_tokens
@property
def kimi_linear_config(self):
config = self.model_config.hf_config
if isinstance(config, KimiLinearConfig):
return config
return None
def _get_linear_attn_registry_result(self):
if self._linear_attn_registry_cache is _UNSET:
self._linear_attn_registry_cache = get_linear_attn_config(
self.model_config.hf_config
)
return self._linear_attn_registry_cache
@property
def linear_attn_model_spec(self):
result = self._get_linear_attn_registry_result()
return result[0] if result else None
@property
def mambaish_config(self):
existing = (
self.mamba2_config
or self.hybrid_gdn_config
or self.kimi_linear_config
or self.hybrid_lightning_config
)
if existing:
return existing
result = self._get_linear_attn_registry_result()
return result[1] if result else None
def _record_kv_cache_dtype(self, resolved: str) -> None: def _record_kv_cache_dtype(self, resolved: str) -> None:
# Load-time resolution transition: the weight-resolved kv-cache dtype # Load-time resolution transition: the weight-resolved kv-cache dtype
# is declared into the flags tier; the dual-apply inside the helper # is declared into the flags tier; the dual-apply inside the helper
@@ -6,6 +6,10 @@ from typing import TYPE_CHECKING, Optional
import torch import torch
from sglang.srt.configs.hybrid_arch import (
hybrid_gdn_config,
mambaish_config,
)
from sglang.srt.configs.model_config import ( from sglang.srt.configs.model_config import (
get_dsa_index_head_dim, get_dsa_index_head_dim,
get_minimax_sparse_attention_config, get_minimax_sparse_attention_config,
@@ -113,7 +117,10 @@ class ModelRunnerKVCacheMixin:
) )
slack_gb = pre_model_load_memory * (1 - self.mem_fraction_static) slack_gb = pre_model_load_memory * (1 - self.mem_fraction_static)
if self.mambaish_config is not None and self.post_capture_kv_active: if (
mambaish_config(self.model_config) is not None
and self.post_capture_kv_active
):
# Mamba state is a fixed pre-capture allocation, so it can't ride the ~0 post-capture slack. # Mamba state is a fixed pre-capture allocation, so it can't ride the ~0 post-capture slack.
slack_gb = max( slack_gb = max(
slack_gb, slack_gb,
@@ -123,7 +130,7 @@ class ModelRunnerKVCacheMixin:
/ 1024, / 1024,
) )
rest_memory = available_gpu_memory - slack_gb rest_memory = available_gpu_memory - slack_gb
if self.mambaish_config is not None: if mambaish_config(self.model_config) is not None:
rest_memory = self.handle_max_mamba_cache(rest_memory) rest_memory = self.handle_max_mamba_cache(rest_memory)
# Loaded weights (target + draft) can exceed the static budget # Loaded weights (target + draft) can exceed the static budget
@@ -147,7 +154,7 @@ class ModelRunnerKVCacheMixin:
return int(rest_memory * (1 << 30)) # return in bytes return int(rest_memory * (1 << 30)) # return in bytes
def handle_max_mamba_cache(self: ModelRunner, total_rest_memory): def handle_max_mamba_cache(self: ModelRunner, total_rest_memory):
config = self.mambaish_config config = mambaish_config(self.model_config)
server_args = self.server_args server_args = self.server_args
assert config is not None assert config is not None
@@ -341,19 +348,21 @@ class ModelRunnerKVCacheMixin:
unsupported_pool_family = None unsupported_pool_family = None
if is_dsv4_model: if is_dsv4_model:
unsupported_pool_family = "DeepSeekV4TokenToKVPool" unsupported_pool_family = "DeepSeekV4TokenToKVPool"
elif current_platform.is_out_of_tree() and not self.mambaish_config: elif current_platform.is_out_of_tree() and not mambaish_config(
self.model_config
):
unsupported_pool_family = "out-of-tree platform KV pool" unsupported_pool_family = "out-of-tree platform KV pool"
elif ( elif self.server_args.attention_backend == "ascend" and not mambaish_config(
self.server_args.attention_backend == "ascend" and not self.mambaish_config self.model_config
): ):
unsupported_pool_family = "NPU/Ascend KV pool" unsupported_pool_family = "NPU/Ascend KV pool"
elif self.use_mla_backend and is_dsa_model: elif self.use_mla_backend and is_dsa_model:
unsupported_pool_family = "DSA/MLA KV pool" unsupported_pool_family = "DSA/MLA KV pool"
elif self.use_mla_backend and not self.mambaish_config: elif self.use_mla_backend and not mambaish_config(self.model_config):
unsupported_pool_family = "MLA KV pool" unsupported_pool_family = "MLA KV pool"
elif self.is_hybrid_swa: elif self.is_hybrid_swa:
unsupported_pool_family = "SWA KV pool" unsupported_pool_family = "SWA KV pool"
elif self.mambaish_config: elif mambaish_config(self.model_config):
unsupported_pool_family = "hybrid linear/Mamba KV pool" unsupported_pool_family = "hybrid linear/Mamba KV pool"
elif is_float4_e2m1fn_x2(self.kv_cache_dtype): elif is_float4_e2m1fn_x2(self.kv_cache_dtype):
unsupported_pool_family = "FP4 MHA KV pool" unsupported_pool_family = "FP4 MHA KV pool"
@@ -401,7 +410,7 @@ class ModelRunnerKVCacheMixin:
decode_max_bs, decode_max_bs,
running_requests, running_requests,
) )
if eager_decode_gap or self.mambaish_config is not None: if eager_decode_gap or mambaish_config(self.model_config) is not None:
headroom_gb = max( headroom_gb = max(
headroom_gb, headroom_gb,
self.server_args.mamba_pre_capture_reserve_mb( self.server_args.mamba_pre_capture_reserve_mb(
@@ -460,7 +469,7 @@ class ModelRunnerKVCacheMixin:
allocator.""" allocator."""
from sglang.srt.mem_cache.unified_memory_pool import init_unified_mamba_pools from sglang.srt.mem_cache.unified_memory_pool import init_unified_mamba_pools
config = self.mambaish_config config = mambaish_config(self.model_config)
assert config is not None assert config is not None
assert ( assert (
not self.use_mla_backend not self.use_mla_backend
@@ -616,7 +625,7 @@ class ModelRunnerKVCacheMixin:
and self.server_args.disaggregation_mode == "null" and self.server_args.disaggregation_mode == "null"
and self.req_to_token_pool is None and self.req_to_token_pool is None
): ):
if self.mambaish_config is not None: if mambaish_config(self.model_config) is not None:
self._init_unified_mamba_pools(max_num_reqs) self._init_unified_mamba_pools(max_num_reqs)
return return
if self.is_hybrid_swa and not is_deepseek_v4(self.model_config.hf_config): if self.is_hybrid_swa and not is_deepseek_v4(self.model_config.hf_config):
@@ -645,7 +654,7 @@ class ModelRunnerKVCacheMixin:
# Extra slots for pre-allocated requests # Extra slots for pre-allocated requests
pre_alloc_size = self.server_args.disaggregation_decode_extra_slots pre_alloc_size = self.server_args.disaggregation_decode_extra_slots
if config := self.mambaish_config: if config := mambaish_config(self.model_config):
self.req_to_token_pool = HybridMambaDecodeReqToTokenPool( self.req_to_token_pool = HybridMambaDecodeReqToTokenPool(
size=max_num_reqs, size=max_num_reqs,
max_context_len=self.model_config.context_len max_context_len=self.model_config.context_len
@@ -677,7 +686,7 @@ class ModelRunnerKVCacheMixin:
enable_memory_saver=self.server_args.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
pre_alloc_size=pre_alloc_size, pre_alloc_size=pre_alloc_size,
) )
elif config := self.mambaish_config: elif config := mambaish_config(self.model_config):
self.req_to_token_pool = HybridReqToTokenPool( self.req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs, size=max_num_reqs,
mamba_size=self.server_args.max_mamba_cache_size, mamba_size=self.server_args.max_mamba_cache_size,
@@ -821,7 +830,9 @@ class ModelRunnerKVCacheMixin:
self.server_args.max_speculative_num_draft_tokens or 0 self.server_args.max_speculative_num_draft_tokens or 0
), ),
) )
elif current_platform.is_out_of_tree() and not self.mambaish_config: elif current_platform.is_out_of_tree() and not mambaish_config(
self.model_config
):
if self.use_mla_backend and is_dsa_model: if self.use_mla_backend and is_dsa_model:
PoolCls = current_platform.get_dsa_kv_pool_cls() PoolCls = current_platform.get_dsa_kv_pool_cls()
self.token_to_kv_pool = PoolCls( self.token_to_kv_pool = PoolCls(
@@ -871,8 +882,8 @@ class ModelRunnerKVCacheMixin:
start_layer=self.start_layer, start_layer=self.start_layer,
end_layer=self.end_layer, end_layer=self.end_layer,
) )
elif ( elif self.server_args.attention_backend == "ascend" and not mambaish_config(
self.server_args.attention_backend == "ascend" and not self.mambaish_config self.model_config
): ):
if self.is_hybrid_swa: if self.is_hybrid_swa:
from sglang.srt.hardware_backend.npu.memory_pool_npu import ( from sglang.srt.hardware_backend.npu.memory_pool_npu import (
@@ -987,7 +998,7 @@ class ModelRunnerKVCacheMixin:
index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config), index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config),
**pool_kwargs, **pool_kwargs,
) )
elif self.use_mla_backend and not self.mambaish_config: elif self.use_mla_backend and not mambaish_config(self.model_config):
assert not is_dsa_model assert not is_dsa_model
if is_float4_e2m1fn_x2(self.kv_cache_dtype): if is_float4_e2m1fn_x2(self.kv_cache_dtype):
self.token_to_kv_pool = MLATokenToKVPoolFP4( self.token_to_kv_pool = MLATokenToKVPoolFP4(
@@ -1074,7 +1085,7 @@ class ModelRunnerKVCacheMixin:
start_layer=self.start_layer, start_layer=self.start_layer,
end_layer=self.end_layer, end_layer=self.end_layer,
) )
elif config := self.mambaish_config: elif config := mambaish_config(self.model_config):
extra_args = {} extra_args = {}
if self.use_mla_backend: if self.use_mla_backend:
extra_args = { extra_args = {
@@ -1178,7 +1189,7 @@ class ModelRunnerKVCacheMixin:
elif _is_npu and ( elif _is_npu and (
self.server_args.attention_backend == "ascend" self.server_args.attention_backend == "ascend"
or is_dsv4_model or is_dsv4_model
or self.hybrid_gdn_config is not None or hybrid_gdn_config(self.model_config) is not None
): ):
if self.is_hybrid_swa: if self.is_hybrid_swa:
# DSV4 on NPU: SWA allocator subclass that also drives the # DSV4 on NPU: SWA allocator subclass that also drives the
@@ -1359,7 +1370,7 @@ class ModelRunnerKVCacheMixin:
requested_per_worker = None requested_per_worker = None
max_num_reqs = min(estimated, token_capacity // 2) max_num_reqs = min(estimated, token_capacity // 2)
if self.mambaish_config is not None: if mambaish_config(self.model_config) is not None:
ratio = self._calculate_mamba_ratio() ratio = self._calculate_mamba_ratio()
max_num_reqs = min( max_num_reqs = min(
max_num_reqs, self.server_args.max_mamba_cache_size // ratio max_num_reqs, self.server_args.max_mamba_cache_size // ratio
@@ -19,6 +19,7 @@ from typing import TYPE_CHECKING, Optional
import torch import torch
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.configs.model_config import ( from sglang.srt.configs.model_config import (
get_dsa_index_head_dim, get_dsa_index_head_dim,
get_minimax_sparse_attention_config, get_minimax_sparse_attention_config,
@@ -124,7 +125,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
def __init__(self, mr: ModelRunner): def __init__(self, mr: ModelRunner):
# Determine effective number of layers for KV cache # Determine effective number of layers for KV cache
if mambaish := mr.mambaish_config: if mambaish := mambaish_config(mr.model_config):
effective_layer_ids = [ effective_layer_ids = [
i i
for i in mambaish.full_attention_layer_ids for i in mambaish.full_attention_layer_ids
@@ -9,6 +9,7 @@ from sglang.kernels.ops.speculative.dflash import (
_compute_dflash_accept_bonus_triton_unchecked, _compute_dflash_accept_bonus_triton_unchecked,
_prepare_dflash_draft_block_unchecked, _prepare_dflash_draft_block_unchecked,
) )
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.distributed import get_tp_group from sglang.srt.distributed import get_tp_group
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.scheduler import GenerationBatchResult
@@ -277,12 +278,11 @@ class DFlashWorkerV2(BaseSpecWorker):
def init_attention_backends(self): def init_attention_backends(self):
self._draft_worker.init_attention_backends() self._draft_worker.init_attention_backends()
self._need_mamba_verify_commit = ( self._need_mamba_verify_commit = mambaish_config(
self.model_runner.mambaish_config is not None self.model_runner.model_config
and hasattr( ) is not None and hasattr(
self.model_runner.attn_backend, self.model_runner.attn_backend,
"update_mamba_state_after_mtp_verify", "update_mamba_state_after_mtp_verify",
)
) )
def init_cuda_graphs(self): def init_cuda_graphs(self):
@@ -1555,7 +1555,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
) )
if ( if (
_is_npu _is_npu
and self._target_worker.model_runner.model_is_mrope and self._target_worker.model_runner.model_config.model_is_mrope
and batch.spec_info is not None and batch.spec_info is not None
and getattr(batch.spec_info, "positions", None) is not None and getattr(batch.spec_info, "positions", None) is not None
and not batch.forward_mode.is_idle() and not batch.forward_mode.is_idle()
+2 -1
View File
@@ -37,6 +37,7 @@ from sglang.kernels.ops.speculative.cache_locs import (
from sglang.kernels.ops.speculative.eagle import ( from sglang.kernels.ops.speculative.eagle import (
fill_accept_out_cache_loc_func as fill_accept_out_cache_loc_func, fill_accept_out_cache_loc_func as fill_accept_out_cache_loc_func,
) )
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.distributed.parallel_state import ( from sglang.srt.distributed.parallel_state import (
GroupCoordinator, GroupCoordinator,
patch_tensor_parallel_group, patch_tensor_parallel_group,
@@ -649,7 +650,7 @@ def commit_mamba_states_after_verify(
commit hook. commit hook.
""" """
model_runner = target_worker.model_runner model_runner = target_worker.model_runner
if model_runner.mambaish_config is None: if mambaish_config(model_runner.model_config) is None:
return return
attn_backend = model_runner.attn_backend attn_backend = model_runner.attn_backend
if not hasattr(attn_backend, "update_mamba_state_after_mtp_verify"): if not hasattr(attn_backend, "update_mamba_state_after_mtp_verify"):
@@ -94,6 +94,7 @@ def _make_model_runner(
mc.get_num_kv_heads = lambda tp_size: num_kv_heads mc.get_num_kv_heads = lambda tp_size: num_kv_heads
mc.get_swa_num_kv_heads = lambda tp_size: swa_num_kv_heads or num_kv_heads mc.get_swa_num_kv_heads = lambda tp_size: swa_num_kv_heads or num_kv_heads
mc.hf_config = SimpleNamespace(architectures=["LlamaForCausalLM"]) mc.hf_config = SimpleNamespace(architectures=["LlamaForCausalLM"])
mc.hf_config.get_text_config = lambda: mc.hf_config
mr.model_config = mc mr.model_config = mc
mr.kv_cache_dtype = "fake_bf16" mr.kv_cache_dtype = "fake_bf16"