[Feature] Xiaomi MiMo-V2.5-Pro day0 support (#23808)
This commit is contained in:
@@ -337,9 +337,9 @@ class ModelConfig:
|
||||
|
||||
if is_draft_model and self.hf_config.architectures[0] == "MiMoForCausalLM":
|
||||
self.hf_config.architectures[0] = "MiMoMTP"
|
||||
if (
|
||||
is_draft_model
|
||||
and self.hf_config.architectures[0] == "MiMoV2FlashForCausalLM"
|
||||
if is_draft_model and self.hf_config.architectures[0] in (
|
||||
"MiMoV2ForCausalLM",
|
||||
"MiMoV2FlashForCausalLM",
|
||||
):
|
||||
self.hf_config.architectures[0] = "MiMoV2MTP"
|
||||
if is_draft_model and self.hf_config.architectures[0] == "Step3p5ForCausalLM":
|
||||
@@ -397,6 +397,7 @@ class ModelConfig:
|
||||
self.has_attention_sinks = self._detect_attention_sinks()
|
||||
|
||||
self.is_hybrid_swa_compress = self.hf_config.architectures[0] in [
|
||||
"MiMoV2ForCausalLM",
|
||||
"MiMoV2FlashForCausalLM",
|
||||
"MiMoV2MTP",
|
||||
"Gemma4ForCausalLM",
|
||||
@@ -416,7 +417,14 @@ class ModelConfig:
|
||||
return True
|
||||
|
||||
# MiMoV2 creates sinks only when the config flags are set.
|
||||
if any(a in archs for a in ("MiMoV2FlashForCausalLM", "MiMoV2MTP")):
|
||||
if any(
|
||||
a in archs
|
||||
for a in (
|
||||
"MiMoV2FlashForCausalLM",
|
||||
"MiMoV2ForCausalLM",
|
||||
"MiMoV2MTP",
|
||||
)
|
||||
):
|
||||
return getattr(
|
||||
self.hf_text_config, "add_swa_attention_sink_bias", False
|
||||
) or getattr(self.hf_text_config, "add_full_attention_sink_bias", False)
|
||||
@@ -1512,6 +1520,7 @@ def is_hybrid_swa_model(model_architectures: List[str]):
|
||||
hybrid_swa_archs = {
|
||||
"Llama4ForConditionalGeneration",
|
||||
"GptOssForCausalLM",
|
||||
"MiMoV2ForCausalLM",
|
||||
"MiMoV2FlashForCausalLM",
|
||||
"MiMoV2MTP",
|
||||
"Step3p5ForCausalLM",
|
||||
@@ -1542,7 +1551,13 @@ def get_hybrid_layer_ids(
|
||||
full_attention_layer_ids = [
|
||||
i for i, x in enumerate(layer_types) if x == "full_attention"
|
||||
]
|
||||
elif "MiMoV2FlashForCausalLM" in model_architectures:
|
||||
elif any(
|
||||
x in model_architectures
|
||||
for x in (
|
||||
"MiMoV2ForCausalLM",
|
||||
"MiMoV2FlashForCausalLM",
|
||||
)
|
||||
):
|
||||
hybrid_layer_pattern = getattr(hf_text_config, "hybrid_layer_pattern", None)
|
||||
swa_attention_layer_ids = [
|
||||
i for i in range(num_hidden_layers) if hybrid_layer_pattern[i] == 1
|
||||
|
||||
@@ -76,7 +76,7 @@ from sglang.srt.utils import (
|
||||
make_layers,
|
||||
)
|
||||
|
||||
MiMoV2FlashConfig = None
|
||||
MiMoV2Config = None
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -178,7 +178,7 @@ class MiMoV2MoE(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: MiMoV2FlashConfig,
|
||||
config: MiMoV2Config,
|
||||
layer_id: int,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
@@ -562,7 +562,7 @@ class MiMoV2Attention(nn.Module):
|
||||
class MiMoV2DecoderLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: MiMoV2FlashConfig,
|
||||
config: MiMoV2Config,
|
||||
layer_id: int = 0,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
@@ -582,7 +582,11 @@ class MiMoV2DecoderLayer(nn.Module):
|
||||
and rope_scaling.get("rope_type") == "default"
|
||||
):
|
||||
rope_scaling = None
|
||||
max_position_embeddings = getattr(config, "max_position_embeddings", 32768)
|
||||
max_position_embeddings = getattr(
|
||||
config,
|
||||
"context_len",
|
||||
getattr(config, "max_position_embeddings", 32768),
|
||||
)
|
||||
|
||||
if self.is_swa_layer():
|
||||
self.self_attn = MiMoV2Attention(
|
||||
@@ -792,7 +796,7 @@ class MiMoV2DecoderLayer(nn.Module):
|
||||
class MiMoV2Model(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: MiMoV2FlashConfig,
|
||||
config: MiMoV2Config,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
decoder_layer_type: type[nn.Module] = MiMoV2DecoderLayer,
|
||||
@@ -943,7 +947,7 @@ class MiMoV2Model(nn.Module):
|
||||
)
|
||||
|
||||
|
||||
class MiMoV2FlashForCausalLM(nn.Module):
|
||||
class MiMoV2ForCausalLM(nn.Module):
|
||||
# BitandBytes specific attributes
|
||||
default_bitsandbytes_target_modules = [
|
||||
".gate_proj.",
|
||||
@@ -965,7 +969,7 @@ class MiMoV2FlashForCausalLM(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: MiMoV2FlashConfig,
|
||||
config: MiMoV2Config,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
@@ -1099,6 +1103,16 @@ class MiMoV2FlashForCausalLM(nn.Module):
|
||||
if "mtp" in name:
|
||||
continue
|
||||
|
||||
# Support fused qkv_proj checkpoint (Pro format)
|
||||
if "qkv_proj" in name:
|
||||
if name in params_dict:
|
||||
tp_size = get_attention_tp_size()
|
||||
tp_rank = get_attention_tp_rank()
|
||||
param = params_dict[name]
|
||||
loaded_weight = loaded_weight.chunk(tp_size, dim=0)[tp_rank]
|
||||
default_weight_loader(param, loaded_weight)
|
||||
continue
|
||||
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
@@ -1173,4 +1187,9 @@ class MiMoV2FlashForCausalLM(nn.Module):
|
||||
)
|
||||
|
||||
|
||||
EntryClass = MiMoV2FlashForCausalLM
|
||||
# Keep the old Flash architecture name loadable while new configs use MiMoV2ForCausalLM.
|
||||
class MiMoV2FlashForCausalLM(MiMoV2ForCausalLM):
|
||||
pass
|
||||
|
||||
|
||||
EntryClass = [MiMoV2ForCausalLM, MiMoV2FlashForCausalLM]
|
||||
+21
-6
@@ -28,6 +28,7 @@ from sglang.srt.layers.communicator import (
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
get_attention_tp_rank,
|
||||
get_attention_tp_size,
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
@@ -39,15 +40,15 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.models.mimo_v2_flash import (
|
||||
from sglang.srt.models.mimo_v2 import (
|
||||
MiMoV2Attention,
|
||||
MiMoV2FlashForCausalLM,
|
||||
MiMoV2ForCausalLM,
|
||||
MiMoV2MLP,
|
||||
)
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
MiMoV2FlashConfig = None
|
||||
MiMoV2Config = None
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -55,7 +56,7 @@ logger = logging.getLogger(__name__)
|
||||
class MiMoV2MTPLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: MiMoV2FlashConfig,
|
||||
config: MiMoV2Config,
|
||||
layer_id: int = 0,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
@@ -71,7 +72,11 @@ class MiMoV2MTPLayer(nn.Module):
|
||||
and rope_scaling.get("rope_type") == "default"
|
||||
):
|
||||
rope_scaling = None
|
||||
max_position_embeddings = getattr(config, "max_position_embeddings", 32768)
|
||||
max_position_embeddings = getattr(
|
||||
config,
|
||||
"context_len",
|
||||
getattr(config, "max_position_embeddings", 32768),
|
||||
)
|
||||
|
||||
self.self_attn = MiMoV2Attention(
|
||||
hidden_size=self.hidden_size,
|
||||
@@ -228,7 +233,7 @@ class MiMoV2ModelNextN(nn.Module):
|
||||
return hidden_states, hidden_states_before_norm
|
||||
|
||||
|
||||
class MiMoV2MTP(MiMoV2FlashForCausalLM):
|
||||
class MiMoV2MTP(MiMoV2ForCausalLM):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -296,6 +301,16 @@ class MiMoV2MTP(MiMoV2FlashForCausalLM):
|
||||
continue
|
||||
name = self.map_model_name_to_mtp_param_name(name)
|
||||
|
||||
# Support fused qkv_proj checkpoint (Pro format)
|
||||
if "qkv_proj" in name:
|
||||
if name in params_dict:
|
||||
tp_size = get_attention_tp_size()
|
||||
tp_rank = get_attention_tp_rank()
|
||||
param = params_dict[name]
|
||||
loaded_weight = loaded_weight.chunk(tp_size, dim=0)[tp_rank]
|
||||
default_weight_loader(param, loaded_weight)
|
||||
continue
|
||||
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
|
||||
if weight_name not in name:
|
||||
@@ -1948,11 +1948,17 @@ class ServerArgs:
|
||||
self.ep_size == 1
|
||||
), "Triton kernel MoE is only supported when ep_size == 1"
|
||||
|
||||
elif "MiMoV2FlashForCausalLM" in model_arch:
|
||||
elif any(
|
||||
x in model_arch
|
||||
for x in (
|
||||
"MiMoV2ForCausalLM",
|
||||
"MiMoV2FlashForCausalLM",
|
||||
)
|
||||
):
|
||||
if self.speculative_algorithm == "EAGLE":
|
||||
self.enable_multi_layer_eagle = True
|
||||
logger.info(
|
||||
"Enable multi-layer EAGLE speculative decoding for MiMoV2FlashForCausalLM model."
|
||||
"Enable multi-layer EAGLE speculative decoding for MiMoV2 model."
|
||||
)
|
||||
if not envs.SGLANG_ENABLE_SPEC_V2.get():
|
||||
envs.SGLANG_ENABLE_SPEC_V2.set(True)
|
||||
@@ -1963,11 +1969,11 @@ class ServerArgs:
|
||||
if self.enable_hierarchical_cache:
|
||||
self.swa_full_tokens_ratio = 1.0
|
||||
logger.warning(
|
||||
"Reset swa_full_tokens_ratio to 1.0 for MiMoV2FlashForCausalLM model with hierarchical cache"
|
||||
"Reset swa_full_tokens_ratio to 1.0 for MiMoV2 model with hierarchical cache"
|
||||
)
|
||||
self.disable_hybrid_swa_memory = True
|
||||
logger.warning(
|
||||
"Disable hybrid SWA memory for MiMoV2FlashForCausalLM model with hierarchical cache"
|
||||
"Disable hybrid SWA memory for MiMoV2 model with hierarchical cache"
|
||||
)
|
||||
elif "Step3p5ForCausalLM" in model_arch:
|
||||
if self.speculative_algorithm == "EAGLE":
|
||||
@@ -7318,6 +7324,7 @@ def auto_choose_speculative_params(self: ServerArgs):
|
||||
"BailingMoeV2_5ForCausalLM",
|
||||
"MistralLarge3ForCausalLM",
|
||||
"PixtralForConditionalGeneration",
|
||||
"MiMoV2ForCausalLM",
|
||||
"MiMoV2FlashForCausalLM",
|
||||
]:
|
||||
return (3, 1, 4)
|
||||
|
||||
@@ -2838,6 +2838,7 @@ def is_fa3_default_architecture(hf_config):
|
||||
"GlmOcrForConditionalGeneration",
|
||||
"Step3VLForConditionalGeneration",
|
||||
"StepVLForConditionalGeneration",
|
||||
"MiMoV2ForCausalLM",
|
||||
"MiMoV2FlashForCausalLM",
|
||||
}
|
||||
return architectures[0] in default_archs
|
||||
|
||||
Reference in New Issue
Block a user