Support Intern-S2-Mobius (#33691)
Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
co-authored by
Xinyuan Tong
parent
a25c330eb1
commit
55f02e6887
@@ -1292,6 +1292,7 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict:
|
|||||||
"Qwen3NextForCausalLM",
|
"Qwen3NextForCausalLM",
|
||||||
"Qwen3_5MoeForConditionalGeneration",
|
"Qwen3_5MoeForConditionalGeneration",
|
||||||
"InternS2PreviewForConditionalGeneration",
|
"InternS2PreviewForConditionalGeneration",
|
||||||
|
"InternS2MobiusForConditionalGeneration",
|
||||||
"Qwen3_5ForConditionalGeneration",
|
"Qwen3_5ForConditionalGeneration",
|
||||||
)
|
)
|
||||||
def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
|
def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||||
@@ -1322,6 +1323,14 @@ def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@_register_for("InternS2MobiusForConditionalGeneration")
|
||||||
|
def _interns2_mobius_baseline_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||||
|
"""Select the only MoE runner validated for the 2,560-expert baseline."""
|
||||||
|
if server_args.moe_runner_backend == "auto":
|
||||||
|
return {"moe_runner_backend": "triton_kernel"}
|
||||||
|
return {}
|
||||||
|
|
||||||
|
|
||||||
@_register_for("Qwen3VLForConditionalGeneration")
|
@_register_for("Qwen3VLForConditionalGeneration")
|
||||||
def _qwen3vl_overrides(server_args: Any, hf_config: Any) -> dict:
|
def _qwen3vl_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||||
|
|
||||||
@@ -1470,6 +1479,7 @@ _MAMBA_RADIX_CACHE_ARCHS = frozenset(
|
|||||||
"Qwen3NextForCausalLM",
|
"Qwen3NextForCausalLM",
|
||||||
"Qwen3_5MoeForConditionalGeneration",
|
"Qwen3_5MoeForConditionalGeneration",
|
||||||
"InternS2PreviewForConditionalGeneration",
|
"InternS2PreviewForConditionalGeneration",
|
||||||
|
"InternS2MobiusForConditionalGeneration",
|
||||||
"Qwen3_5ForConditionalGeneration",
|
"Qwen3_5ForConditionalGeneration",
|
||||||
"MiniCPMV4_6ForConditionalGeneration",
|
"MiniCPMV4_6ForConditionalGeneration",
|
||||||
"NemotronHForCausalLM",
|
"NemotronHForCausalLM",
|
||||||
|
|||||||
@@ -15,6 +15,11 @@ from sglang.srt.configs.inkling import (
|
|||||||
InklingModelConfig,
|
InklingModelConfig,
|
||||||
InklingVisionConfig,
|
InklingVisionConfig,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.configs.interns2_mobius import (
|
||||||
|
InternS2MobiusConfig,
|
||||||
|
InternS2MobiusTextConfig,
|
||||||
|
InternS2MobiusVisionConfig,
|
||||||
|
)
|
||||||
from sglang.srt.configs.interns2preview import InternS2PreviewConfig
|
from sglang.srt.configs.interns2preview import InternS2PreviewConfig
|
||||||
from sglang.srt.configs.janus_pro import MultiModalityConfig
|
from sglang.srt.configs.janus_pro import MultiModalityConfig
|
||||||
from sglang.srt.configs.jet_nemotron import JetNemotronConfig
|
from sglang.srt.configs.jet_nemotron import JetNemotronConfig
|
||||||
@@ -81,6 +86,9 @@ __all__ = [
|
|||||||
"Qwen3_5TextConfig",
|
"Qwen3_5TextConfig",
|
||||||
"Qwen3_5MoeTextConfig",
|
"Qwen3_5MoeTextConfig",
|
||||||
"InternS2PreviewConfig",
|
"InternS2PreviewConfig",
|
||||||
|
"InternS2MobiusConfig",
|
||||||
|
"InternS2MobiusTextConfig",
|
||||||
|
"InternS2MobiusVisionConfig",
|
||||||
"DotsVLMConfig",
|
"DotsVLMConfig",
|
||||||
"DotsOCRConfig",
|
"DotsOCRConfig",
|
||||||
"FalconH1Config",
|
"FalconH1Config",
|
||||||
|
|||||||
@@ -0,0 +1,54 @@
|
|||||||
|
from typing import ClassVar
|
||||||
|
|
||||||
|
from sglang.srt.configs.qwen3_5 import (
|
||||||
|
Qwen3_5MoeConfig,
|
||||||
|
Qwen3_5MoeTextConfig,
|
||||||
|
Qwen3_5MoeVisionConfig,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class InternS2MobiusVisionConfig(Qwen3_5MoeVisionConfig):
|
||||||
|
model_type = "interns2_mobius"
|
||||||
|
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
class InternS2MobiusTextConfig(Qwen3_5MoeTextConfig):
|
||||||
|
model_type = "interns2_mobius_text"
|
||||||
|
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
layer_types = kwargs.get("layer_types")
|
||||||
|
if "dtype" not in kwargs and "torch_dtype" not in kwargs:
|
||||||
|
kwargs["dtype"] = "bfloat16"
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
# The local Qwen3NextConfig currently consumes layer_types without
|
||||||
|
# retaining it. Mobius checkpoints provide the authoritative per-layer
|
||||||
|
# schedule, so keep it verbatim rather than regenerating it.
|
||||||
|
if layer_types is not None:
|
||||||
|
self.layer_types = layer_types
|
||||||
|
if not hasattr(self, "num_blocks"):
|
||||||
|
self.num_blocks = 4
|
||||||
|
|
||||||
|
|
||||||
|
class InternS2MobiusConfig(Qwen3_5MoeConfig):
|
||||||
|
model_type = "interns2_mobius"
|
||||||
|
sub_configs: ClassVar = {
|
||||||
|
"vision_config": InternS2MobiusVisionConfig,
|
||||||
|
"text_config": InternS2MobiusTextConfig,
|
||||||
|
}
|
||||||
|
|
||||||
|
def __init__(self, text_config=None, vision_config=None, **kwargs):
|
||||||
|
if "dtype" not in kwargs and "torch_dtype" not in kwargs:
|
||||||
|
kwargs["dtype"] = "bfloat16"
|
||||||
|
|
||||||
|
if isinstance(text_config, dict):
|
||||||
|
text_config = dict(text_config)
|
||||||
|
if "dtype" not in text_config and "torch_dtype" not in text_config:
|
||||||
|
text_config["dtype"] = "bfloat16"
|
||||||
|
|
||||||
|
super().__init__(
|
||||||
|
text_config=text_config,
|
||||||
|
vision_config=vision_config,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
@@ -678,7 +678,20 @@ class ModelConfig:
|
|||||||
"Qwen3_5ForCausalLM",
|
"Qwen3_5ForCausalLM",
|
||||||
"Qwen3_5MoeForCausalLM",
|
"Qwen3_5MoeForCausalLM",
|
||||||
"InternS2PreviewForConditionalGeneration",
|
"InternS2PreviewForConditionalGeneration",
|
||||||
|
"InternS2MobiusForConditionalGeneration",
|
||||||
]:
|
]:
|
||||||
|
if (
|
||||||
|
self.hf_config.architectures[0]
|
||||||
|
== "InternS2MobiusForConditionalGeneration"
|
||||||
|
):
|
||||||
|
# The target owns 2,560 experts through four shared physical
|
||||||
|
# banks, while its bundled MTP layer is an ordinary Qwen3.5
|
||||||
|
# MoE layer with the checkpoint-declared smaller expert set.
|
||||||
|
self.hf_text_config.model_type = "qwen3_5_moe_text"
|
||||||
|
self.hf_text_config.num_experts = self.hf_text_config.mtp_num_experts
|
||||||
|
self.hf_text_config.num_experts_per_tok = (
|
||||||
|
self.hf_text_config.mtp_num_experts_per_tok
|
||||||
|
)
|
||||||
self.hf_config.architectures[0] = "Qwen3_5ForCausalLMMTP"
|
self.hf_config.architectures[0] = "Qwen3_5ForCausalLMMTP"
|
||||||
self.hf_config.num_nextn_predict_layers = 1
|
self.hf_config.num_nextn_predict_layers = 1
|
||||||
|
|
||||||
@@ -1808,6 +1821,7 @@ multimodal_model_archs = [
|
|||||||
"Qwen3_5ForConditionalGeneration",
|
"Qwen3_5ForConditionalGeneration",
|
||||||
"Qwen3_5MoeForConditionalGeneration",
|
"Qwen3_5MoeForConditionalGeneration",
|
||||||
"InternS2PreviewForConditionalGeneration",
|
"InternS2PreviewForConditionalGeneration",
|
||||||
|
"InternS2MobiusForConditionalGeneration",
|
||||||
"Qwen3ASRForConditionalGeneration",
|
"Qwen3ASRForConditionalGeneration",
|
||||||
"Qwen3OmniMoeForConditionalGeneration",
|
"Qwen3OmniMoeForConditionalGeneration",
|
||||||
"KimiVLForConditionalGeneration",
|
"KimiVLForConditionalGeneration",
|
||||||
@@ -1858,6 +1872,7 @@ multimodal_piecewise_cuda_graph_supported_model_archs = [
|
|||||||
# Multimodal archs whose LM prefill is validated under breakable CUDA graph;
|
# Multimodal archs whose LM prefill is validated under breakable CUDA graph;
|
||||||
# embed-carrying batches are rejected at replay (can_run_graph) and run eager.
|
# embed-carrying batches are rejected at replay (can_run_graph) and run eager.
|
||||||
multimodal_breakable_cuda_graph_supported_model_archs = [
|
multimodal_breakable_cuda_graph_supported_model_archs = [
|
||||||
|
"InternS2MobiusForConditionalGeneration",
|
||||||
"Qwen3_5ForConditionalGeneration",
|
"Qwen3_5ForConditionalGeneration",
|
||||||
"Qwen3_5MoeForConditionalGeneration",
|
"Qwen3_5MoeForConditionalGeneration",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -74,6 +74,7 @@ def get_rope_index(
|
|||||||
model_type.startswith("qwen3_vl")
|
model_type.startswith("qwen3_vl")
|
||||||
or model_type.startswith("qwen3_vl_moe")
|
or model_type.startswith("qwen3_vl_moe")
|
||||||
or model_type.startswith("qwen3_5")
|
or model_type.startswith("qwen3_5")
|
||||||
|
or model_type.startswith("interns2_mobius")
|
||||||
) and video_grid_thw is not None:
|
) and video_grid_thw is not None:
|
||||||
video_grid_thw = torch.repeat_interleave(
|
video_grid_thw = torch.repeat_interleave(
|
||||||
video_grid_thw, video_grid_thw[:, 0], dim=0
|
video_grid_thw, video_grid_thw[:, 0], dim=0
|
||||||
@@ -160,6 +161,7 @@ def get_rope_index(
|
|||||||
"qwen3_5",
|
"qwen3_5",
|
||||||
"qwen3_5_moe",
|
"qwen3_5_moe",
|
||||||
"intern_s2_preview",
|
"intern_s2_preview",
|
||||||
|
"interns2_mobius",
|
||||||
):
|
):
|
||||||
t_index = (
|
t_index = (
|
||||||
torch.arange(llm_grid_t, device=position_ids.device)
|
torch.arange(llm_grid_t, device=position_ids.device)
|
||||||
|
|||||||
@@ -981,6 +981,12 @@ class LoRAManager:
|
|||||||
):
|
):
|
||||||
layer_id = get_layer_id(module_name)
|
layer_id = get_layer_id(module_name)
|
||||||
if layer_id is None:
|
if layer_id is None:
|
||||||
|
if module_name.startswith("model.meta_mlp."):
|
||||||
|
raise ValueError(
|
||||||
|
"LoRA on Intern-S2-Mobius model.meta_mlp routed banks "
|
||||||
|
"is not supported by the baseline; remove routed-expert "
|
||||||
|
"targets or use a future bank-aware LoRA implementation."
|
||||||
|
)
|
||||||
# FusedMoE submodules outside the decoder layer hierarchy
|
# FusedMoE submodules outside the decoder layer hierarchy
|
||||||
# (e.g. nested helpers under non-".layers." prefixes) have
|
# (e.g. nested helpers under non-".layers." prefixes) have
|
||||||
# no resolvable layer id; skip them so we don't index
|
# no resolvable layer id; skip them so we don't index
|
||||||
|
|||||||
@@ -0,0 +1,819 @@
|
|||||||
|
"""Inference-only Intern-S2-Mobius model."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from collections.abc import Iterable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from sglang.srt.configs.interns2_mobius import (
|
||||||
|
InternS2MobiusConfig,
|
||||||
|
InternS2MobiusTextConfig,
|
||||||
|
)
|
||||||
|
from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce
|
||||||
|
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
|
||||||
|
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
||||||
|
from sglang.srt.layers.layernorm import GemmaRMSNorm
|
||||||
|
from sglang.srt.layers.linear import (
|
||||||
|
QKVParallelLinear,
|
||||||
|
ReplicatedLinear,
|
||||||
|
RowParallelLinear,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.moe import (
|
||||||
|
should_skip_post_experts_all_reduce,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
|
||||||
|
from sglang.srt.layers.moe.topk import TopK
|
||||||
|
from sglang.srt.layers.moe.utils import (
|
||||||
|
RoutingMethodType,
|
||||||
|
filter_moe_weight_param_global_expert,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
|
from sglang.srt.layers.rotary_embedding import get_rope
|
||||||
|
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||||
|
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.qwen2_moe import Qwen2MoeMLP
|
||||||
|
from sglang.srt.models.qwen3_5 import (
|
||||||
|
QWEN3_5_KV_SCALE_MAPPER,
|
||||||
|
Qwen3_5AttentionDecoderLayer,
|
||||||
|
Qwen3_5ForCausalLM,
|
||||||
|
Qwen3_5ForConditionalGeneration,
|
||||||
|
Qwen3_5GatedDeltaNet,
|
||||||
|
_enable_qwen35_fused_ar_quant,
|
||||||
|
_linear_accepts_fp8_tuple,
|
||||||
|
)
|
||||||
|
from sglang.srt.runtime_context import get_forward, get_parallel, get_stream
|
||||||
|
from sglang.srt.utils import add_prefix, is_cuda, make_layers
|
||||||
|
from sglang.srt.utils.hf_transformers_utils import get_rope_config
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
_is_cuda = is_cuda()
|
||||||
|
|
||||||
|
_MOBIUS_PACKED_WEIGHT_MAPPING = (
|
||||||
|
("qkv_proj", "q_proj", "q"),
|
||||||
|
("qkv_proj", "k_proj", "k"),
|
||||||
|
("qkv_proj", "v_proj", "v"),
|
||||||
|
("gate_up_proj", "gate_proj", 0),
|
||||||
|
("gate_up_proj", "up_proj", 1),
|
||||||
|
("in_proj_qkvz.", "in_proj_qkv.", (0, 1, 2)),
|
||||||
|
("in_proj_qkvz.", "in_proj_z.", 3),
|
||||||
|
("in_proj_ba.", "in_proj_b.", 0),
|
||||||
|
("in_proj_ba.", "in_proj_a.", 1),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_intentional_mobius_skip(name: str, tie_word_embeddings: bool) -> bool:
|
||||||
|
return (
|
||||||
|
name.startswith("mtp.")
|
||||||
|
or name.endswith(
|
||||||
|
(
|
||||||
|
".rotary_emb.inv_freq",
|
||||||
|
".rotary_emb.cos_cached",
|
||||||
|
".rotary_emb.sin_cached",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
or (tie_word_embeddings and name == "lm_head.weight")
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_mobius_weight_name(name: str) -> str:
|
||||||
|
if name.startswith("model.language_model."):
|
||||||
|
name = "model." + name.removeprefix("model.language_model.")
|
||||||
|
elif name.startswith("model.visual."):
|
||||||
|
name = "visual." + name.removeprefix("model.visual.")
|
||||||
|
|
||||||
|
if ".self_attn." in name:
|
||||||
|
name = name.replace(".self_attn.", ".")
|
||||||
|
if name.startswith("visual."):
|
||||||
|
name = name.replace(".attn.qkv.", ".attn.qkv_proj.")
|
||||||
|
return name
|
||||||
|
|
||||||
|
|
||||||
|
def _load_fused_mobius_expert_weight(
|
||||||
|
*,
|
||||||
|
name: str,
|
||||||
|
loaded_weight: torch.Tensor,
|
||||||
|
params_dict: dict[str, nn.Parameter],
|
||||||
|
num_experts: int,
|
||||||
|
record_slot,
|
||||||
|
) -> None:
|
||||||
|
if name.endswith("experts.gate_up_proj"):
|
||||||
|
parameter_name = name.replace("experts.gate_up_proj", "experts.w13_weight")
|
||||||
|
if parameter_name not in params_dict:
|
||||||
|
raise KeyError(
|
||||||
|
f"Mobius fused gate/up destination is missing: {parameter_name}"
|
||||||
|
)
|
||||||
|
if loaded_weight.shape[0] != num_experts:
|
||||||
|
raise ValueError(
|
||||||
|
f"Expected {num_experts} experts in {name}, got {loaded_weight.shape[0]}"
|
||||||
|
)
|
||||||
|
gate_weights, up_weights = loaded_weight.chunk(2, dim=-2)
|
||||||
|
parameter = params_dict[parameter_name]
|
||||||
|
loader = parameter.weight_loader
|
||||||
|
for expert_id in range(num_experts):
|
||||||
|
for shard_id, expert_weight in (
|
||||||
|
("w1", gate_weights[expert_id]),
|
||||||
|
("w3", up_weights[expert_id]),
|
||||||
|
):
|
||||||
|
record_slot(parameter_name, shard_id, expert_id)
|
||||||
|
loader(
|
||||||
|
parameter,
|
||||||
|
expert_weight,
|
||||||
|
parameter_name,
|
||||||
|
shard_id,
|
||||||
|
expert_id,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if name.endswith("experts.down_proj"):
|
||||||
|
parameter_name = name.replace("experts.down_proj", "experts.w2_weight")
|
||||||
|
if parameter_name not in params_dict:
|
||||||
|
raise KeyError(
|
||||||
|
f"Mobius fused down destination is missing: {parameter_name}"
|
||||||
|
)
|
||||||
|
if loaded_weight.shape[0] != num_experts:
|
||||||
|
raise ValueError(
|
||||||
|
f"Expected {num_experts} experts in {name}, got {loaded_weight.shape[0]}"
|
||||||
|
)
|
||||||
|
parameter = params_dict[parameter_name]
|
||||||
|
loader = parameter.weight_loader
|
||||||
|
for expert_id in range(num_experts):
|
||||||
|
record_slot(parameter_name, "w2", expert_id)
|
||||||
|
loader(
|
||||||
|
parameter,
|
||||||
|
loaded_weight[expert_id],
|
||||||
|
parameter_name,
|
||||||
|
"w2",
|
||||||
|
expert_id,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
raise KeyError(f"Unexpected Mobius fused expert tensor: {name}")
|
||||||
|
|
||||||
|
|
||||||
|
def _expected_mobius_load_slots(
|
||||||
|
params_dict: dict[str, nn.Parameter], num_experts: int
|
||||||
|
) -> set[tuple[str, object, int | None]]:
|
||||||
|
expected = set()
|
||||||
|
seen_parameters = set()
|
||||||
|
for name, parameter in params_dict.items():
|
||||||
|
# RadixLinearAttention exposes A_log/dt_bias aliases that refer to the
|
||||||
|
# same Parameter already owned by Qwen3_5GatedDeltaNet. Require one
|
||||||
|
# canonical load, not one load per module alias.
|
||||||
|
parameter_id = id(parameter)
|
||||||
|
if parameter_id in seen_parameters:
|
||||||
|
continue
|
||||||
|
seen_parameters.add(parameter_id)
|
||||||
|
if ".meta_mlp." in name and name.endswith("experts.w13_weight"):
|
||||||
|
for expert_id in range(num_experts):
|
||||||
|
expected.add((name, "w1", expert_id))
|
||||||
|
expected.add((name, "w3", expert_id))
|
||||||
|
elif ".meta_mlp." in name and name.endswith("experts.w2_weight"):
|
||||||
|
for expert_id in range(num_experts):
|
||||||
|
expected.add((name, "w2", expert_id))
|
||||||
|
elif ".qkv_proj." in name and name.startswith("model.layers."):
|
||||||
|
for shard_id in ("q", "k", "v"):
|
||||||
|
expected.add((name, shard_id, None))
|
||||||
|
elif ".mlp.shared_expert.gate_up_proj." in name:
|
||||||
|
for shard_id in (0, 1):
|
||||||
|
expected.add((name, shard_id, None))
|
||||||
|
elif ".in_proj_qkvz." in name:
|
||||||
|
expected.add((name, (0, 1, 2), None))
|
||||||
|
expected.add((name, 3, None))
|
||||||
|
elif ".in_proj_ba." in name:
|
||||||
|
expected.add((name, 0, None))
|
||||||
|
expected.add((name, 1, None))
|
||||||
|
else:
|
||||||
|
expected.add((name, None, None))
|
||||||
|
return expected
|
||||||
|
|
||||||
|
|
||||||
|
def _load_mobius_weights_strict(
|
||||||
|
owner: nn.Module,
|
||||||
|
config: InternS2MobiusTextConfig,
|
||||||
|
weights: Iterable[tuple[str, torch.Tensor]],
|
||||||
|
) -> set[str]:
|
||||||
|
params_dict = dict(owner.named_parameters(remove_duplicate=False))
|
||||||
|
expected_slots = _expected_mobius_load_slots(params_dict, config.num_experts)
|
||||||
|
loaded_slots: set[tuple[str, object, int | None]] = set()
|
||||||
|
loaded_sources: set[str] = set()
|
||||||
|
|
||||||
|
def record_slot(name, shard_id=None, expert_id=None):
|
||||||
|
slot = (name, shard_id, expert_id)
|
||||||
|
if slot in loaded_slots:
|
||||||
|
raise ValueError(f"Mobius destination load is duplicated: {slot}")
|
||||||
|
loaded_slots.add(slot)
|
||||||
|
|
||||||
|
for source_name, loaded_weight in weights:
|
||||||
|
if _is_intentional_mobius_skip(source_name, config.tie_word_embeddings):
|
||||||
|
continue
|
||||||
|
if source_name in loaded_sources:
|
||||||
|
raise ValueError(f"Mobius checkpoint key is duplicated: {source_name}")
|
||||||
|
loaded_sources.add(source_name)
|
||||||
|
|
||||||
|
name = _normalize_mobius_weight_name(source_name)
|
||||||
|
if ".meta_mlp." in name and name.endswith(
|
||||||
|
("experts.gate_up_proj", "experts.down_proj")
|
||||||
|
):
|
||||||
|
_load_fused_mobius_expert_weight(
|
||||||
|
name=name,
|
||||||
|
loaded_weight=loaded_weight,
|
||||||
|
params_dict=params_dict,
|
||||||
|
num_experts=config.num_experts,
|
||||||
|
record_slot=record_slot,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
for parameter_name, weight_name, shard_id in _MOBIUS_PACKED_WEIGHT_MAPPING:
|
||||||
|
if weight_name not in name:
|
||||||
|
continue
|
||||||
|
# Vision qkv is already fused in the checkpoint.
|
||||||
|
if name.startswith("visual."):
|
||||||
|
continue
|
||||||
|
destination = name.replace(weight_name, parameter_name)
|
||||||
|
if destination not in params_dict:
|
||||||
|
raise KeyError(
|
||||||
|
f"Mobius packed destination is missing: {destination} "
|
||||||
|
f"(from {source_name})"
|
||||||
|
)
|
||||||
|
parameter = params_dict[destination]
|
||||||
|
loader = parameter.weight_loader
|
||||||
|
record_slot(destination, shard_id)
|
||||||
|
loader(parameter, loaded_weight, shard_id)
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
if name not in params_dict:
|
||||||
|
raise KeyError(
|
||||||
|
f"Mobius destination is missing: {name} (from {source_name})"
|
||||||
|
)
|
||||||
|
parameter = params_dict[name]
|
||||||
|
loader = getattr(parameter, "weight_loader", default_weight_loader)
|
||||||
|
record_slot(name)
|
||||||
|
loader(parameter, loaded_weight)
|
||||||
|
|
||||||
|
missing = expected_slots - loaded_slots
|
||||||
|
unexpected = loaded_slots - expected_slots
|
||||||
|
if missing or unexpected:
|
||||||
|
details = []
|
||||||
|
if missing:
|
||||||
|
details.append(f"missing destinations: {sorted(map(str, missing))[:20]}")
|
||||||
|
if unexpected:
|
||||||
|
details.append(
|
||||||
|
f"unexpected destinations: {sorted(map(str, unexpected))[:20]}"
|
||||||
|
)
|
||||||
|
raise ValueError("Mobius weight coverage failure; " + "; ".join(details))
|
||||||
|
return loaded_sources
|
||||||
|
|
||||||
|
|
||||||
|
class InternS2MobiusRoutedExpertBank(nn.Module):
|
||||||
|
"""One physical routed-expert bank, without a shared expert or reduction."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
bank_id: int,
|
||||||
|
config: InternS2MobiusTextConfig,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
prefix: str = "",
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.bank_id = bank_id
|
||||||
|
self.tp_size = get_parallel().tp_size
|
||||||
|
self.num_experts = config.num_experts
|
||||||
|
self.topk = TopK(
|
||||||
|
top_k=config.num_experts_per_tok,
|
||||||
|
renormalize=config.norm_topk_prob,
|
||||||
|
layer_id=bank_id,
|
||||||
|
)
|
||||||
|
self.experts = get_moe_impl_class(quant_config)(
|
||||||
|
layer_id=bank_id,
|
||||||
|
top_k=config.num_experts_per_tok,
|
||||||
|
num_experts=config.num_experts,
|
||||||
|
hidden_size=config.hidden_size,
|
||||||
|
intermediate_size=config.moe_intermediate_size,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=add_prefix("experts", prefix),
|
||||||
|
routing_method_type=RoutingMethodType.RenormalizeNaive,
|
||||||
|
num_fused_shared_experts=0,
|
||||||
|
# The layer-local shared path consumes the same post-attention input.
|
||||||
|
inplace=False,
|
||||||
|
)
|
||||||
|
self.gate = ReplicatedLinear(
|
||||||
|
config.hidden_size,
|
||||||
|
config.num_experts,
|
||||||
|
bias=False,
|
||||||
|
quant_config=None,
|
||||||
|
prefix=add_prefix("gate", prefix),
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_moe_weights(self):
|
||||||
|
return [
|
||||||
|
parameter.data
|
||||||
|
for name, parameter in self.experts.named_parameters()
|
||||||
|
if name != "correction_bias"
|
||||||
|
and filter_moe_weight_param_global_expert(
|
||||||
|
name, parameter, self.experts.num_local_experts
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
def forward_routed(
|
||||||
|
self,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
forward_batch: ForwardBatch | None = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Return an unreduced TP-partial routed result without mutating input."""
|
||||||
|
del forward_batch
|
||||||
|
original_shape = hidden_states.shape
|
||||||
|
hidden_states = hidden_states.reshape(-1, original_shape[-1])
|
||||||
|
|
||||||
|
if hidden_states.shape[0] == 0:
|
||||||
|
# Always enter the expert implementation so collective-capable
|
||||||
|
# implementations keep identical participation on idle ranks.
|
||||||
|
topk_output = self.topk.empty_topk_output(hidden_states.device)
|
||||||
|
else:
|
||||||
|
router_logits, _ = self.gate(hidden_states)
|
||||||
|
topk_output = self.topk(hidden_states, router_logits)
|
||||||
|
|
||||||
|
output = self.experts(hidden_states, topk_output)
|
||||||
|
return output.reshape(original_shape)
|
||||||
|
|
||||||
|
|
||||||
|
def _mobius_reduce_combined_output(combined: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Apply the one ordinary TP reduction unless the scoped runtime owns it."""
|
||||||
|
if get_parallel().tp_size > 1 and not should_skip_post_experts_all_reduce(
|
||||||
|
is_tp_path=True
|
||||||
|
):
|
||||||
|
return tensor_model_parallel_all_reduce(combined)
|
||||||
|
return combined
|
||||||
|
|
||||||
|
|
||||||
|
def _get_mobius_routed_bank(meta_mlp: nn.ModuleList, layer_id: int) -> nn.Module:
|
||||||
|
if not meta_mlp:
|
||||||
|
raise ValueError(
|
||||||
|
"Intern-S2-Mobius requires at least one physical routed-expert bank"
|
||||||
|
)
|
||||||
|
return meta_mlp[layer_id % len(meta_mlp)]
|
||||||
|
|
||||||
|
|
||||||
|
class _InternS2MobiusDecoderMixin:
|
||||||
|
def _init_mobius_mlp(
|
||||||
|
self,
|
||||||
|
config: InternS2MobiusTextConfig,
|
||||||
|
quant_config: QuantizationConfig | None,
|
||||||
|
layer_prefix: str,
|
||||||
|
) -> None:
|
||||||
|
self.mlp = InternS2MobiusLayerMlp(
|
||||||
|
config=config,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=add_prefix("mlp", layer_prefix),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _forward_mobius_mlp(
|
||||||
|
self,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
meta_mlp: nn.ModuleList,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
# Evaluate the layer-local branch first. Besides matching the intended
|
||||||
|
# equation, this prevents a broken/in-place routed backend from
|
||||||
|
# corrupting the input consumed by the shared expert or its gate.
|
||||||
|
if hidden_states.shape[0] == 0:
|
||||||
|
shared = torch.zeros_like(hidden_states)
|
||||||
|
else:
|
||||||
|
shared = self.mlp.shared_expert(hidden_states)
|
||||||
|
gate, _ = self.mlp.shared_expert_gate(hidden_states)
|
||||||
|
shared = torch.sigmoid(gate) * shared
|
||||||
|
|
||||||
|
routed = _get_mobius_routed_bank(meta_mlp, self.layer_id).forward_routed(
|
||||||
|
hidden_states, forward_batch
|
||||||
|
)
|
||||||
|
return _mobius_reduce_combined_output(routed + shared)
|
||||||
|
|
||||||
|
def _forward_after_attention(
|
||||||
|
self,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
residual: torch.Tensor | None,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
meta_mlp: nn.ModuleList,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||||
|
hidden_states, residual = self.layer_communicator.prepare_mlp(
|
||||||
|
hidden_states, residual, forward_batch
|
||||||
|
)
|
||||||
|
mlp_reduce_scatter = self.layer_communicator.should_use_reduce_scatter(
|
||||||
|
forward_batch
|
||||||
|
)
|
||||||
|
# Model-side all-reduce fusion is intentionally disabled for baseline.
|
||||||
|
with get_forward().scoped(
|
||||||
|
fuse_mlp_allreduce=False,
|
||||||
|
mlp_reduce_scatter=mlp_reduce_scatter,
|
||||||
|
):
|
||||||
|
hidden_states = self._forward_mobius_mlp(
|
||||||
|
hidden_states, forward_batch, meta_mlp
|
||||||
|
)
|
||||||
|
hidden_states, residual = self.layer_communicator.postprocess_layer(
|
||||||
|
hidden_states, residual, forward_batch
|
||||||
|
)
|
||||||
|
return hidden_states, residual
|
||||||
|
|
||||||
|
|
||||||
|
class InternS2MobiusLayerMlp(nn.Module):
|
||||||
|
"""Layer-local shared branch; routed banks live on the causal model."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: InternS2MobiusTextConfig,
|
||||||
|
quant_config: QuantizationConfig | None,
|
||||||
|
prefix: str,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.shared_expert = Qwen2MoeMLP(
|
||||||
|
hidden_size=config.hidden_size,
|
||||||
|
intermediate_size=config.shared_expert_intermediate_size,
|
||||||
|
hidden_act=config.hidden_act,
|
||||||
|
quant_config=quant_config,
|
||||||
|
reduce_results=False,
|
||||||
|
prefix=add_prefix("shared_expert", prefix),
|
||||||
|
)
|
||||||
|
self.shared_expert_gate = ReplicatedLinear(
|
||||||
|
config.hidden_size,
|
||||||
|
1,
|
||||||
|
bias=False,
|
||||||
|
quant_config=None,
|
||||||
|
prefix=add_prefix("shared_expert_gate", prefix),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class InternS2MobiusLinearDecoderLayer(_InternS2MobiusDecoderMixin, nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: InternS2MobiusTextConfig,
|
||||||
|
layer_id: int,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
prefix: str = "",
|
||||||
|
alt_stream: torch.cuda.Stream | None = None,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.config = config
|
||||||
|
self.layer_id = layer_id
|
||||||
|
self.linear_attn = Qwen3_5GatedDeltaNet(
|
||||||
|
config, layer_id, quant_config, alt_stream, prefix
|
||||||
|
)
|
||||||
|
|
||||||
|
layer_prefix = prefix.removesuffix(".linear_attn")
|
||||||
|
self._init_mobius_mlp(config, quant_config, layer_prefix)
|
||||||
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
|
layer_id=layer_id,
|
||||||
|
num_layers=config.num_hidden_layers,
|
||||||
|
is_layer_sparse=True,
|
||||||
|
is_previous_layer_sparse=True,
|
||||||
|
is_next_layer_sparse=True,
|
||||||
|
)
|
||||||
|
self.input_layernorm = GemmaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||||
|
self.post_attention_layernorm = GemmaRMSNorm(
|
||||||
|
config.hidden_size, eps=config.rms_norm_eps
|
||||||
|
)
|
||||||
|
enable_fused_ar_quant = (
|
||||||
|
_enable_qwen35_fused_ar_quant()
|
||||||
|
and _linear_accepts_fp8_tuple(self.linear_attn.in_proj_qkvz)
|
||||||
|
)
|
||||||
|
self.layer_communicator = LayerCommunicator(
|
||||||
|
layer_scatter_modes=self.layer_scatter_modes,
|
||||||
|
input_layernorm=self.input_layernorm,
|
||||||
|
post_attention_layernorm=self.post_attention_layernorm,
|
||||||
|
allow_reduce_scatter=True,
|
||||||
|
is_last_layer=(layer_id == config.num_hidden_layers - 1),
|
||||||
|
enable_fused_ar_quant=enable_fused_ar_quant,
|
||||||
|
fused_ar_quant_keep_bf16=enable_fused_ar_quant,
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
residual: torch.Tensor | None,
|
||||||
|
meta_mlp: nn.ModuleList,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
forward_batch = kwargs["forward_batch"]
|
||||||
|
hidden_states, residual = (
|
||||||
|
self.layer_communicator.prepare_attn_and_capture_last_layer_outputs(
|
||||||
|
hidden_states,
|
||||||
|
residual,
|
||||||
|
forward_batch,
|
||||||
|
captured_last_layer_outputs=kwargs.get("captured_last_layer_outputs"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if not forward_batch.forward_mode.is_idle():
|
||||||
|
hidden_states = self.linear_attn(hidden_states, forward_batch)
|
||||||
|
return self._forward_after_attention(
|
||||||
|
hidden_states, residual, forward_batch, meta_mlp
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class InternS2MobiusAttentionDecoderLayer(
|
||||||
|
_InternS2MobiusDecoderMixin, Qwen3_5AttentionDecoderLayer
|
||||||
|
):
|
||||||
|
"""Mobius-owned full-attention constructor reusing Qwen3.5 methods."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: InternS2MobiusTextConfig,
|
||||||
|
layer_id: int,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
prefix: str = "",
|
||||||
|
alt_stream: torch.cuda.Stream | None = None,
|
||||||
|
) -> None:
|
||||||
|
nn.Module.__init__(self)
|
||||||
|
self.config = config
|
||||||
|
self.hidden_size = config.hidden_size
|
||||||
|
self.attn_tp_rank = get_parallel().attn_tp_rank
|
||||||
|
self.attn_tp_size = get_parallel().attn_tp_size
|
||||||
|
self.total_num_heads = config.num_attention_heads
|
||||||
|
if self.total_num_heads % self.attn_tp_size != 0:
|
||||||
|
raise ValueError("num_attention_heads must be divisible by attention TP")
|
||||||
|
self.num_heads = self.total_num_heads // self.attn_tp_size
|
||||||
|
self.total_num_kv_heads = config.num_key_value_heads
|
||||||
|
if self.total_num_kv_heads >= self.attn_tp_size:
|
||||||
|
if self.total_num_kv_heads % self.attn_tp_size != 0:
|
||||||
|
raise ValueError(
|
||||||
|
"num_key_value_heads must be divisible by attention TP"
|
||||||
|
)
|
||||||
|
elif self.attn_tp_size % self.total_num_kv_heads != 0:
|
||||||
|
raise ValueError("attention TP must be divisible by num_key_value_heads")
|
||||||
|
self.num_kv_heads = max(1, self.total_num_kv_heads // self.attn_tp_size)
|
||||||
|
self.head_dim = config.head_dim or (self.hidden_size // self.num_heads)
|
||||||
|
self.q_size = self.num_heads * self.head_dim
|
||||||
|
self.kv_size = self.num_kv_heads * self.head_dim
|
||||||
|
self.scaling = self.head_dim**-0.5
|
||||||
|
self.max_position_embeddings = getattr(config, "max_position_embeddings", 8192)
|
||||||
|
self.rope_theta, rope_scaling = get_rope_config(config)
|
||||||
|
self.partial_rotary_factor = getattr(config, "partial_rotary_factor", 1.0)
|
||||||
|
self.layer_id = layer_id
|
||||||
|
if rope_scaling and not ("rope_type" in rope_scaling or "type" in rope_scaling):
|
||||||
|
rope_scaling = None
|
||||||
|
self.attn_output_gate = getattr(config, "attn_output_gate", True)
|
||||||
|
self.rotary_emb = get_rope(
|
||||||
|
head_size=self.head_dim,
|
||||||
|
rotary_dim=self.head_dim,
|
||||||
|
max_position=self.max_position_embeddings,
|
||||||
|
rope_scaling=rope_scaling,
|
||||||
|
base=self.rope_theta,
|
||||||
|
partial_rotary_factor=self.partial_rotary_factor,
|
||||||
|
is_neox_style=True,
|
||||||
|
dtype=torch.get_default_dtype(),
|
||||||
|
)
|
||||||
|
self.qkv_proj = QKVParallelLinear(
|
||||||
|
config.hidden_size,
|
||||||
|
self.head_dim,
|
||||||
|
self.total_num_heads * (1 + self.attn_output_gate),
|
||||||
|
self.total_num_kv_heads,
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
tp_rank=self.attn_tp_rank,
|
||||||
|
tp_size=self.attn_tp_size,
|
||||||
|
prefix=add_prefix("qkv_proj", prefix),
|
||||||
|
)
|
||||||
|
self.o_proj = RowParallelLinear(
|
||||||
|
self.total_num_heads * self.head_dim,
|
||||||
|
config.hidden_size,
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
reduce_results=False,
|
||||||
|
tp_rank=self.attn_tp_rank,
|
||||||
|
tp_size=self.attn_tp_size,
|
||||||
|
prefix=add_prefix("o_proj", prefix),
|
||||||
|
)
|
||||||
|
self.attn = RadixAttention(
|
||||||
|
self.num_heads,
|
||||||
|
self.head_dim,
|
||||||
|
self.scaling,
|
||||||
|
num_kv_heads=self.num_kv_heads,
|
||||||
|
layer_id=layer_id,
|
||||||
|
prefix=f"{prefix}.attn",
|
||||||
|
quant_config=quant_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
layer_prefix = prefix.removesuffix(".self_attn")
|
||||||
|
self._init_mobius_mlp(config, quant_config, layer_prefix)
|
||||||
|
self.layer_scatter_modes = LayerScatterModes.init_new(
|
||||||
|
layer_id=layer_id,
|
||||||
|
num_layers=config.num_hidden_layers,
|
||||||
|
is_layer_sparse=True,
|
||||||
|
is_previous_layer_sparse=True,
|
||||||
|
is_next_layer_sparse=True,
|
||||||
|
)
|
||||||
|
self.input_layernorm = GemmaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||||
|
self.post_attention_layernorm = GemmaRMSNorm(
|
||||||
|
config.hidden_size, eps=config.rms_norm_eps
|
||||||
|
)
|
||||||
|
self.q_norm = GemmaRMSNorm(self.head_dim, eps=config.rms_norm_eps)
|
||||||
|
self.k_norm = GemmaRMSNorm(self.head_dim, eps=config.rms_norm_eps)
|
||||||
|
enable_fused_ar_quant = (
|
||||||
|
_enable_qwen35_fused_ar_quant() and _linear_accepts_fp8_tuple(self.qkv_proj)
|
||||||
|
)
|
||||||
|
self.layer_communicator = LayerCommunicator(
|
||||||
|
layer_scatter_modes=self.layer_scatter_modes,
|
||||||
|
input_layernorm=self.input_layernorm,
|
||||||
|
post_attention_layernorm=self.post_attention_layernorm,
|
||||||
|
allow_reduce_scatter=True,
|
||||||
|
is_last_layer=(layer_id == config.num_hidden_layers - 1),
|
||||||
|
enable_fused_ar_quant=enable_fused_ar_quant,
|
||||||
|
fused_ar_quant_keep_bf16=False,
|
||||||
|
)
|
||||||
|
self.alt_stream = alt_stream
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
residual: torch.Tensor | None,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
meta_mlp: nn.ModuleList,
|
||||||
|
captured_last_layer_outputs: list[torch.Tensor] | None = None,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
del kwargs
|
||||||
|
hidden_states, residual = (
|
||||||
|
self.layer_communicator.prepare_attn_and_capture_last_layer_outputs(
|
||||||
|
hidden_states,
|
||||||
|
residual,
|
||||||
|
forward_batch,
|
||||||
|
captured_last_layer_outputs=captured_last_layer_outputs,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if not forward_batch.forward_mode.is_idle():
|
||||||
|
hidden_states = self.self_attention(
|
||||||
|
positions=positions,
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
)
|
||||||
|
return self._forward_after_attention(
|
||||||
|
hidden_states, residual, forward_batch, meta_mlp
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class InternS2MobiusForCausalLM(Qwen3_5ForCausalLM):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: InternS2MobiusTextConfig,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
prefix: str = "",
|
||||||
|
) -> None:
|
||||||
|
nn.Module.__init__(self)
|
||||||
|
self.config = config
|
||||||
|
self.hidden_size = config.hidden_size
|
||||||
|
self.pp_group = get_pp_group()
|
||||||
|
if self.pp_group.world_size != 1:
|
||||||
|
raise ValueError(
|
||||||
|
"Intern-S2-Mobius baseline does not support pipeline parallelism"
|
||||||
|
)
|
||||||
|
|
||||||
|
alt_stream = get_stream("alt") if _is_cuda else None
|
||||||
|
self.embed_tokens = VocabParallelEmbedding(
|
||||||
|
config.vocab_size,
|
||||||
|
config.hidden_size,
|
||||||
|
org_num_embeddings=config.vocab_size,
|
||||||
|
enable_tp=not is_dp_attention_enabled(),
|
||||||
|
)
|
||||||
|
|
||||||
|
bank_prefix = prefix.replace("model.language_model", "model")
|
||||||
|
self.meta_mlp = nn.ModuleList(
|
||||||
|
[
|
||||||
|
InternS2MobiusRoutedExpertBank(
|
||||||
|
bank_id=bank_id,
|
||||||
|
config=config,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=add_prefix(f"meta_mlp.{bank_id}", bank_prefix),
|
||||||
|
)
|
||||||
|
for bank_id in range(config.num_blocks)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
if len(self.meta_mlp) != config.num_blocks:
|
||||||
|
raise AssertionError(
|
||||||
|
"physical routed-expert bank count does not match num_blocks"
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_layer(idx: int, prefix: str):
|
||||||
|
checkpoint_type = config.layer_types[idx]
|
||||||
|
if checkpoint_type == "full_attention":
|
||||||
|
return InternS2MobiusAttentionDecoderLayer(
|
||||||
|
config=config,
|
||||||
|
layer_id=idx,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=add_prefix("self_attn", prefix),
|
||||||
|
alt_stream=alt_stream,
|
||||||
|
)
|
||||||
|
if checkpoint_type == "linear_attention":
|
||||||
|
return InternS2MobiusLinearDecoderLayer(
|
||||||
|
config=config,
|
||||||
|
layer_id=idx,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=add_prefix("linear_attn", prefix),
|
||||||
|
alt_stream=alt_stream,
|
||||||
|
)
|
||||||
|
raise ValueError(f"Unsupported Mobius layer type: {checkpoint_type}")
|
||||||
|
|
||||||
|
self.layers, self._start_layer, self._end_layer = make_layers(
|
||||||
|
config.num_hidden_layers,
|
||||||
|
get_layer,
|
||||||
|
pp_rank=0,
|
||||||
|
pp_size=1,
|
||||||
|
prefix=f"{prefix}.layers",
|
||||||
|
)
|
||||||
|
self.norm = GemmaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||||
|
self.layers_to_capture = []
|
||||||
|
|
||||||
|
def get_hidden_dim(self, module_name: str, layer_idx: int):
|
||||||
|
if module_name == "gate_up_proj":
|
||||||
|
return (
|
||||||
|
self.config.hidden_size,
|
||||||
|
self.config.shared_expert_intermediate_size * 2,
|
||||||
|
)
|
||||||
|
if module_name == "down_proj":
|
||||||
|
return self.config.shared_expert_intermediate_size, self.config.hidden_size
|
||||||
|
return super().get_hidden_dim(module_name, layer_idx)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids: torch.Tensor,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
input_embeds: torch.Tensor | None = None,
|
||||||
|
pp_proxy_tensors: PPProxyTensors | None = None,
|
||||||
|
input_deepstack_embeds: torch.Tensor | None = None,
|
||||||
|
) -> torch.Tensor | PPProxyTensors:
|
||||||
|
if pp_proxy_tensors is not None:
|
||||||
|
raise ValueError("Intern-S2-Mobius baseline does not support PP tensors")
|
||||||
|
hidden_states = (
|
||||||
|
self.embed_tokens(input_ids) if input_embeds is None else input_embeds
|
||||||
|
)
|
||||||
|
residual = None
|
||||||
|
aux_hidden_states = []
|
||||||
|
for layer_idx, layer in enumerate(self.layers):
|
||||||
|
hidden_states, residual = layer(
|
||||||
|
positions=positions,
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
residual=residual,
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
meta_mlp=self.meta_mlp,
|
||||||
|
captured_last_layer_outputs=(
|
||||||
|
aux_hidden_states
|
||||||
|
if getattr(layer, "_is_layer_to_capture", False)
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
input_deepstack_embeds is not None
|
||||||
|
and input_deepstack_embeds.numel() > 0
|
||||||
|
and layer_idx < 3
|
||||||
|
):
|
||||||
|
start = self.hidden_size * layer_idx
|
||||||
|
hidden_states.add_(
|
||||||
|
input_deepstack_embeds[:, start : start + self.hidden_size]
|
||||||
|
)
|
||||||
|
|
||||||
|
if hidden_states.shape[0] != 0:
|
||||||
|
if residual is None:
|
||||||
|
hidden_states = self.norm(hidden_states)
|
||||||
|
else:
|
||||||
|
hidden_states, _ = self.norm(hidden_states, residual)
|
||||||
|
return (
|
||||||
|
hidden_states
|
||||||
|
if not aux_hidden_states
|
||||||
|
else (hidden_states, aux_hidden_states)
|
||||||
|
)
|
||||||
|
|
||||||
|
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
|
||||||
|
raise ValueError(
|
||||||
|
"Load Intern-S2-Mobius through its conditional-generation wrapper "
|
||||||
|
"so vision, language, lm_head, and strict coverage are handled together"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class InternS2MobiusForConditionalGeneration(Qwen3_5ForConditionalGeneration):
|
||||||
|
packed_modules_mapping = InternS2MobiusForCausalLM.packed_modules_mapping
|
||||||
|
supported_lora_modules = InternS2MobiusForCausalLM.supported_lora_modules
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: InternS2MobiusConfig,
|
||||||
|
quant_config: QuantizationConfig | None = None,
|
||||||
|
prefix: str = "",
|
||||||
|
language_model_cls=InternS2MobiusForCausalLM,
|
||||||
|
) -> None:
|
||||||
|
super().__init__(config, quant_config, prefix, language_model_cls)
|
||||||
|
|
||||||
|
def should_apply_lora(self, module_name: str) -> bool:
|
||||||
|
# Meta banks require bank-aware adapter ownership and are excluded.
|
||||||
|
return module_name.startswith("model.layers.")
|
||||||
|
|
||||||
|
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
|
||||||
|
return _load_mobius_weights_strict(
|
||||||
|
self,
|
||||||
|
self.config,
|
||||||
|
QWEN3_5_KV_SCALE_MAPPER.apply(weights),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
EntryClass = [InternS2MobiusForConditionalGeneration]
|
||||||
@@ -17,6 +17,9 @@ from sglang.srt.managers.schedule_batch import (
|
|||||||
MultimodalDataItem,
|
MultimodalDataItem,
|
||||||
MultimodalProcessorOutput,
|
MultimodalProcessorOutput,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.models.interns2_mobius import (
|
||||||
|
InternS2MobiusForConditionalGeneration,
|
||||||
|
)
|
||||||
from sglang.srt.models.interns2preview import InternS2PreviewForConditionalGeneration
|
from sglang.srt.models.interns2preview import InternS2PreviewForConditionalGeneration
|
||||||
from sglang.srt.models.qwen2_5_vl import Qwen2_5_VLForConditionalGeneration
|
from sglang.srt.models.qwen2_5_vl import Qwen2_5_VLForConditionalGeneration
|
||||||
from sglang.srt.models.qwen2_vl import Qwen2VLForConditionalGeneration
|
from sglang.srt.models.qwen2_vl import Qwen2VLForConditionalGeneration
|
||||||
@@ -292,6 +295,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
|||||||
Qwen3_5MoeForConditionalGeneration,
|
Qwen3_5MoeForConditionalGeneration,
|
||||||
Qwen3_5ForCausalLMMTP,
|
Qwen3_5ForCausalLMMTP,
|
||||||
InternS2PreviewForConditionalGeneration,
|
InternS2PreviewForConditionalGeneration,
|
||||||
|
InternS2MobiusForConditionalGeneration,
|
||||||
Qwen3OmniMoeForConditionalGeneration,
|
Qwen3OmniMoeForConditionalGeneration,
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -305,6 +309,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
|||||||
"qwen3_5",
|
"qwen3_5",
|
||||||
"qwen3_5_moe",
|
"qwen3_5_moe",
|
||||||
"intern_s2_preview",
|
"intern_s2_preview",
|
||||||
|
"interns2_mobius",
|
||||||
):
|
):
|
||||||
# Two workers overlap CPU preprocessing without over-fragmenting
|
# Two workers overlap CPU preprocessing without over-fragmenting
|
||||||
# burst arrivals into smaller GPU prefill batches. Higher counts can
|
# burst arrivals into smaller GPU prefill batches. Higher counts can
|
||||||
@@ -514,6 +519,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
|||||||
"qwen3_5",
|
"qwen3_5",
|
||||||
"qwen3_5_moe",
|
"qwen3_5_moe",
|
||||||
"intern_s2_preview",
|
"intern_s2_preview",
|
||||||
|
"interns2_mobius",
|
||||||
):
|
):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -757,6 +763,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
|
|||||||
"qwen3_5",
|
"qwen3_5",
|
||||||
"qwen3_5_moe",
|
"qwen3_5_moe",
|
||||||
"intern_s2_preview",
|
"intern_s2_preview",
|
||||||
|
"interns2_mobius",
|
||||||
):
|
):
|
||||||
processor_kwargs.update(
|
processor_kwargs.update(
|
||||||
video_metadata=video_metadata,
|
video_metadata=video_metadata,
|
||||||
|
|||||||
@@ -5057,9 +5057,21 @@ class ServerArgs:
|
|||||||
self._resolved_overrides = []
|
self._resolved_overrides = []
|
||||||
return
|
return
|
||||||
|
|
||||||
hf_config = self.get_model_config().hf_config
|
model_config = self.get_model_config()
|
||||||
|
hf_config = model_config.hf_config
|
||||||
model_arch = hf_config.architectures[0]
|
model_arch = hf_config.architectures[0]
|
||||||
|
|
||||||
|
if model_arch == "InternS2MobiusForConditionalGeneration":
|
||||||
|
unsupported = []
|
||||||
|
if self.pp_size != 1:
|
||||||
|
unsupported.append("pipeline parallelism (--pp-size must be 1)")
|
||||||
|
if self.ep_size != 1:
|
||||||
|
unsupported.append("expert parallelism (--ep-size must be 1)")
|
||||||
|
if unsupported:
|
||||||
|
raise ValueError(
|
||||||
|
"Intern-S2-Mobius does not support: " + "; ".join(unsupported) + "."
|
||||||
|
)
|
||||||
|
|
||||||
if self.enable_dsa_cache_layer_split and not is_deepseek_dsa(hf_config):
|
if self.enable_dsa_cache_layer_split and not is_deepseek_dsa(hf_config):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--enable-dsa-cache-layer-split is only supported for DSA "
|
"--enable-dsa-cache-layer-split is only supported for DSA "
|
||||||
|
|||||||
@@ -36,6 +36,8 @@ from sglang.srt.configs import (
|
|||||||
InklingMMConfig,
|
InklingMMConfig,
|
||||||
InklingModelConfig,
|
InklingModelConfig,
|
||||||
InklingVisionConfig,
|
InklingVisionConfig,
|
||||||
|
InternS2MobiusConfig,
|
||||||
|
InternS2MobiusTextConfig,
|
||||||
InternS2PreviewConfig,
|
InternS2PreviewConfig,
|
||||||
JetNemotronConfig,
|
JetNemotronConfig,
|
||||||
JetVLMConfig,
|
JetVLMConfig,
|
||||||
@@ -116,6 +118,8 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
|
|||||||
Qwen3_5TextConfig,
|
Qwen3_5TextConfig,
|
||||||
Qwen3_5MoeTextConfig,
|
Qwen3_5MoeTextConfig,
|
||||||
InternS2PreviewConfig,
|
InternS2PreviewConfig,
|
||||||
|
InternS2MobiusConfig,
|
||||||
|
InternS2MobiusTextConfig,
|
||||||
JetNemotronConfig,
|
JetNemotronConfig,
|
||||||
JetVLMConfig,
|
JetVLMConfig,
|
||||||
KimiK25Config,
|
KimiK25Config,
|
||||||
|
|||||||
Reference in New Issue
Block a user