From f0bf652534c54fc71528fa9508330e8e980bc8cb Mon Sep 17 00:00:00 2001
From: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
Date: Thu, 17 Sep 2026 00:19:26 +0800
Subject: [PATCH] Add Ling-3.0-flash-VL model support (#38526)
---
.../kernels/ops/attention/rotary_triton.py | 4 +-
.../arg_groups/model_overrides/__init__.py | 1 +
.../model_overrides/bailing_moe_v3.py | 48 +++
python/sglang/srt/arg_groups/overrides.py | 2 +
.../batch_invariant_ops.py | 3 +
python/sglang/srt/configs/__init__.py | 5 +-
python/sglang/srt/configs/bailing_hybrid.py | 80 +++-
python/sglang/srt/configs/bailing_moe_v2.py | 172 ++++++++
python/sglang/srt/configs/hybrid_arch.py | 4 +
python/sglang/srt/configs/model_config.py | 41 +-
python/sglang/srt/configs/qwen3_vl.py | 2 +
python/sglang/srt/layers/moe/topk.py | 90 ++++-
python/sglang/srt/layers/multi_gate.py | 233 +++++++++++
.../layers/rotary_embedding/bailing_mrope.py | 366 ++++++++++++++++++
.../srt/layers/rotary_embedding/base.py | 13 +-
.../srt/layers/rotary_embedding/factory.py | 48 ++-
.../rotary_embedding/mrope_rope_index.py | 1 +
python/sglang/srt/managers/schedule_batch.py | 47 ++-
python/sglang/srt/managers/scheduler.py | 27 +-
.../srt/model_executor/forward_batch_info.py | 80 ++++
python/sglang/srt/models/bailing_mm.py | 270 +++++++++++++
python/sglang/srt/models/bailing_mm_v3.py | 325 ++++++++++++++++
python/sglang/srt/models/bailing_moe.py | 131 ++++++-
python/sglang/srt/models/bailing_moe_v3.py | 189 +++++++--
python/sglang/srt/models/qwen3_vl.py | 72 ++--
.../srt/multimodal/processors/bailing_mm.py | 261 +++++++++++++
.../sglang/srt/parser/template_detection.py | 17 +
.../srt/utils/hf_transformers/common.py | 4 +
.../unit/eplb/test_waterfill_eplb.py | 10 +-
.../unit/layers/test_bailing_mrope_shape.py | 212 ++++++++++
.../test_bailing_modality_metadata.py | 197 ++++++++++
.../unit/models/test_bailing_vl_loader.py | 277 +++++++++++++
.../test_shared_experts_fusion_gates.py | 63 ++-
.../multimodal/test_bailing_mm_processor.py | 34 ++
.../multimodal/test_gpu_feature_transport.py | 9 +-
.../unit/parser/test_template_manager.py | 49 ++-
test/registered/unit/test_model_overrides.py | 68 ++++
37 files changed, 3332 insertions(+), 123 deletions(-)
create mode 100644 python/sglang/srt/arg_groups/model_overrides/bailing_moe_v3.py
create mode 100644 python/sglang/srt/configs/bailing_moe_v2.py
create mode 100644 python/sglang/srt/layers/multi_gate.py
create mode 100644 python/sglang/srt/layers/rotary_embedding/bailing_mrope.py
create mode 100644 python/sglang/srt/models/bailing_mm.py
create mode 100644 python/sglang/srt/models/bailing_mm_v3.py
create mode 100644 python/sglang/srt/multimodal/processors/bailing_mm.py
create mode 100644 test/registered/unit/layers/test_bailing_mrope_shape.py
create mode 100644 test/registered/unit/managers/test_bailing_modality_metadata.py
create mode 100644 test/registered/unit/models/test_bailing_vl_loader.py
create mode 100644 test/registered/unit/multimodal/test_bailing_mm_processor.py
diff --git a/python/sglang/kernels/ops/attention/rotary_triton.py b/python/sglang/kernels/ops/attention/rotary_triton.py
index 0a8dc2c33..3e066bb07 100644
--- a/python/sglang/kernels/ops/attention/rotary_triton.py
+++ b/python/sglang/kernels/ops/attention/rotary_triton.py
@@ -245,8 +245,8 @@ def triton_ernie45_rope_fused_inplace(
section_h, section_w, section_t = mrope_section
assert section_h == section_w, "Ernie4.5 layout assumes section_h == section_w"
assert section_h + section_w + section_t == rd // 2
- if cos_sin_cache.dtype != q.dtype or cos_sin_cache.device != q.device:
- cos_sin_cache = cos_sin_cache.to(device=q.device, dtype=q.dtype)
+ if cos_sin_cache.device != q.device:
+ cos_sin_cache = cos_sin_cache.to(device=q.device)
pad_n_qh = triton.next_power_of_2(n_qh)
pad_n_kh = triton.next_power_of_2(n_kh)
pad_hd = triton.next_power_of_2(head_size)
diff --git a/python/sglang/srt/arg_groups/model_overrides/__init__.py b/python/sglang/srt/arg_groups/model_overrides/__init__.py
index 2d6240e9f..871b71012 100644
--- a/python/sglang/srt/arg_groups/model_overrides/__init__.py
+++ b/python/sglang/srt/arg_groups/model_overrides/__init__.py
@@ -8,6 +8,7 @@ nobody would own that value, and which module supplied it would come down to
the order of the imports below. Keep each field owned by one family module.
"""
+from sglang.srt.arg_groups.model_overrides import bailing_moe_v3 # noqa: F401
from sglang.srt.arg_groups.model_overrides import cohere2_moe # noqa: F401
from sglang.srt.arg_groups.model_overrides import deepseek_v2 # noqa: F401
from sglang.srt.arg_groups.model_overrides import deepseek_v4 # noqa: F401
diff --git a/python/sglang/srt/arg_groups/model_overrides/bailing_moe_v3.py b/python/sglang/srt/arg_groups/model_overrides/bailing_moe_v3.py
new file mode 100644
index 000000000..817b02480
--- /dev/null
+++ b/python/sglang/srt/arg_groups/model_overrides/bailing_moe_v3.py
@@ -0,0 +1,48 @@
+"""Config-time override declarations for bailing_moe_v3.
+
+Architectures: BailingMoeV3ForCausalLM,
+BailingMoeV3VLForConditionalGeneration.
+"""
+
+import logging
+from typing import Any
+
+from sglang.srt.arg_groups.model_override_base import (
+ _register_for,
+ model_config_of,
+ resolving_view,
+)
+from sglang.srt.runtime_context import get_platform
+
+logger = logging.getLogger(__name__)
+
+
+@_register_for(
+ "BailingMoeV3ForCausalLM",
+ "BailingMoeV3VLForConditionalGeneration",
+)
+def _bailing_moe_v3_overrides(server_args: Any, hf_config: Any) -> dict:
+ cfg = resolving_view(server_args)
+ if (
+ cfg.moe_runner_backend != "auto"
+ or cfg.device != "cuda"
+ or cfg.moe_a2a_backend != "none"
+ or get_platform().is_hip
+ ):
+ return {}
+ if not (
+ get_platform().is_sm90 or get_platform().is_sm100 or get_platform().is_sm120
+ ):
+ return {}
+
+ model_config = model_config_of(server_args)
+ if model_config.quantization != "fp8" or not model_config.is_fp4_experts:
+ return {}
+
+ model_arch = hf_config.architectures[0]
+ logger.info(
+ "Bailing V3 mixed FP8/MXFP4 checkpoint: "
+ "moe_runner_backend=flashinfer_mxfp4 for %s.",
+ model_arch,
+ )
+ return {"moe_runner_backend": "flashinfer_mxfp4"}
diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py
index 5813f8fc5..e01cc31ac 100644
--- a/python/sglang/srt/arg_groups/overrides.py
+++ b/python/sglang/srt/arg_groups/overrides.py
@@ -435,6 +435,7 @@ _MAMBA_RADIX_CACHE_ARCHS = frozenset(
"KimiK3ForConditionalGeneration",
"BailingMoeV2_5ForCausalLM",
"BailingMoeV3ForCausalLM",
+ "BailingMoeV3VLForConditionalGeneration",
"Qwen3NextForCausalLM",
"Qwen3_5MoeForConditionalGeneration",
"InternS2PreviewForConditionalGeneration",
@@ -476,6 +477,7 @@ _MAMBA_EXTRA_BUFFER_ARCHS = frozenset(
"MiniCPMV4_6ForConditionalGeneration",
"BailingMoeV2_5ForCausalLM",
"BailingMoeV3ForCausalLM",
+ "BailingMoeV3VLForConditionalGeneration",
"FalconH1ForCausalLM",
"GraniteMoeHybridForCausalLM",
"Glm5NextForConditionalGeneration",
diff --git a/python/sglang/srt/batch_invariant_ops/batch_invariant_ops.py b/python/sglang/srt/batch_invariant_ops/batch_invariant_ops.py
index 509bd4156..882f3535d 100644
--- a/python/sglang/srt/batch_invariant_ops/batch_invariant_ops.py
+++ b/python/sglang/srt/batch_invariant_ops/batch_invariant_ops.py
@@ -296,6 +296,8 @@ def matmul_persistent(
# DeepGEMM has minimum dimension requirements for TMA descriptors
MIN_DEEPGEMM_DIM = 16
+ element_size = a.element_size()
+ deepgemm_tma_aligned = (N * element_size) % 16 == 0 and (K * element_size) % 16 == 0
if (
_ENABLE_MM_DEEPGEMM
@@ -305,6 +307,7 @@ def matmul_persistent(
and a.is_contiguous()
and b.transpose(0, 1).is_contiguous()
and N >= MIN_DEEPGEMM_DIM
+ and deepgemm_tma_aligned
):
if _ENABLE_MM_COMPARISON_TEST:
out_triton = _matmul_persistent_triton(
diff --git a/python/sglang/srt/configs/__init__.py b/python/sglang/srt/configs/__init__.py
index 9e5c28d15..476983b04 100644
--- a/python/sglang/srt/configs/__init__.py
+++ b/python/sglang/srt/configs/__init__.py
@@ -1,5 +1,6 @@
from sglang.srt.configs.afmoe import AfmoeConfig
-from sglang.srt.configs.bailing_hybrid import BailingHybridConfig
+from sglang.srt.configs.bailing_hybrid import BailingHybridConfig, BailingMoeV3VLConfig
+from sglang.srt.configs.bailing_moe_v2 import BailingMM2Config
from sglang.srt.configs.chatglm import ChatGLMConfig
from sglang.srt.configs.cohere2_moe import Cohere2MoeConfig
from sglang.srt.configs.cosmos3 import (
@@ -86,6 +87,8 @@ from sglang.srt.configs.zaya import ZayaConfig
__all__ = [
"AfmoeConfig",
"BailingHybridConfig",
+ "BailingMM2Config",
+ "BailingMoeV3VLConfig",
"ExaoneConfig",
"ChatGLMConfig",
"Cosmos3Config",
diff --git a/python/sglang/srt/configs/bailing_hybrid.py b/python/sglang/srt/configs/bailing_hybrid.py
index 567c30963..6c24da731 100644
--- a/python/sglang/srt/configs/bailing_hybrid.py
+++ b/python/sglang/srt/configs/bailing_hybrid.py
@@ -26,11 +26,20 @@ from sglang.srt.configs.mamba_utils import (
Mamba2CacheParams,
Mamba2StateShape,
)
+from sglang.srt.configs.qwen3_vl import Qwen3VLMoeVisionConfig
from sglang.srt.runtime_context import get_parallel
logger = logging.get_logger(__name__)
+def is_bailing_multi_gate_enabled(config: PretrainedConfig) -> bool:
+ """Select MultiRouter only when the checkpoint config declares it."""
+ return (
+ bool(getattr(config, "multi_gate", False))
+ or getattr(config, "router_type", "topN") == "MultiRouter"
+ )
+
+
class HybridLayerType(enum.Enum):
full_attention = "attention"
linear_attention = "linear_attention"
@@ -76,7 +85,7 @@ class BailingHybridConfig(PretrainedConfig):
use_qk_norm=True,
num_nextn_predict_layers=0,
mtp_loss_scaling_factor=0,
- moe_router_enable_expert_bias=True,
+ moe_router_enable_expert_bias=False,
routed_scaling_factor=1.0,
layer_group_size=1,
group_norm_size=1,
@@ -225,3 +234,72 @@ class BailingHybridConfig(PretrainedConfig):
)
return Mamba2CacheParams(shape=shape, layers=self.linear_layer_ids)
+
+
+class BailingMoeV3VLConfig(PretrainedConfig):
+ model_type = "bailing_moe_v3_vl"
+
+ def __init__(
+ self,
+ text_config=None,
+ vision_config=None,
+ image_token_id=157157,
+ video_token_id=156909,
+ vision_start_token_id=157158,
+ vision_end_token_id=157159,
+ tie_word_embeddings=False,
+ mrope_section=None,
+ **kwargs,
+ ):
+ if isinstance(vision_config, dict):
+ vision_config = dict(vision_config)
+ # The public Bailing checkpoint omits deepstack entirely. Do not
+ # inherit Qwen3-VL's architecture-specific deepstack defaults.
+ vision_config.setdefault("deepstack_visual_indexes", [])
+ vision_config = Qwen3VLMoeVisionConfig(**vision_config)
+ elif vision_config is None:
+ vision_config = Qwen3VLMoeVisionConfig(deepstack_visual_indexes=[])
+
+ if isinstance(text_config, dict):
+ text_config = BailingHybridConfig(**text_config)
+ elif text_config is None:
+ text_config = BailingHybridConfig()
+
+ self.vision_config = vision_config
+ self.text_config = text_config
+ self.image_token_id = image_token_id
+ self.video_token_id = video_token_id
+ self.vision_start_token_id = vision_start_token_id
+ self.vision_end_token_id = vision_end_token_id
+
+ if mrope_section is None:
+ mrope_section = text_config.rope_parameters.get(
+ "mrope_section", [8, 12, 12]
+ )
+ self.mrope_section = mrope_section
+ text_config.rope_parameters.update(
+ rope_type="default",
+ mrope_section=mrope_section,
+ video_rope=True,
+ )
+
+ if self.text_config.architectures is None:
+ self.text_config.architectures = ["BailingMoeV3ForCausalLM"]
+
+ super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
+
+ @property
+ def rope_scaling(self):
+ """The language model's effective RoPE parameters (v5 backcompat alias)."""
+ return self.text_config.rope_parameters
+
+ @rope_scaling.setter
+ def rope_scaling(self, value):
+ # A top-level rope_scaling override (e.g. --json-model-override-args
+ # '{"rope_scaling": ...}') targets the language model's rope. Merge it
+ # into the text config's rope_parameters so the mrope_section and
+ # video_rope markers injected above survive the override.
+ if isinstance(value, dict) and hasattr(self, "text_config"):
+ self.text_config.rope_parameters.update(value)
+ else:
+ PretrainedConfig.rope_scaling.fset(self, value)
diff --git a/python/sglang/srt/configs/bailing_moe_v2.py b/python/sglang/srt/configs/bailing_moe_v2.py
new file mode 100644
index 000000000..1602e140d
--- /dev/null
+++ b/python/sglang/srt/configs/bailing_moe_v2.py
@@ -0,0 +1,172 @@
+# Copyright 2023-2026 SGLang Team
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+from transformers.configuration_utils import PretrainedConfig
+
+from sglang.srt.configs.qwen3_vl import Qwen3VLMoeVisionConfig
+
+
+class BailingMoeV2Config(PretrainedConfig):
+ model_type = "bailing_moe_v2"
+ ignore_keys_at_rope_validation = {"mrope_section", "use_video_rope"}
+
+ def __init__(
+ self,
+ vocab_size=30592,
+ hidden_size=1024,
+ intermediate_size=None,
+ num_hidden_layers=24,
+ num_attention_heads=16,
+ num_key_value_heads=0,
+ hidden_act="silu",
+ use_qkv_bias=False,
+ use_qk_norm=False,
+ use_bias=True,
+ rms_norm_eps=1e-5,
+ norm_head=False,
+ tie_word_embeddings=False,
+ embedding_dropout=0.1,
+ attention_dropout=0.1,
+ output_dropout=0.1,
+ initializer_range=0.02,
+ max_position_embeddings=16384,
+ rope_theta=10000.0,
+ use_cache=True,
+ use_sliding_window=False,
+ sliding_window=81920,
+ max_window_layers=28,
+ rope_scaling=None,
+ pad_token_id=126081,
+ num_experts=16,
+ num_shared_experts=0,
+ num_experts_per_tok=2,
+ n_group=8,
+ topk_group=4,
+ routed_scaling_factor=2.5,
+ moe_intermediate_size=None,
+ first_k_dense_replace=0,
+ head_dim=None,
+ output_router_logits=False,
+ partial_rotary_factor=0.5,
+ router_type="topN",
+ norm_topk_prob=True,
+ moe_router_enable_expert_bias=False,
+ _attn_implementation="flash_attention_2",
+ use_interleaved_frame_timestamp=True,
+ **kwargs,
+ ):
+ self.num_hidden_layers = num_hidden_layers
+ self.vocab_size = vocab_size
+ self.hidden_size = hidden_size
+ self.intermediate_size = intermediate_size
+ self.num_attention_heads = num_attention_heads
+ self.num_key_value_heads = num_key_value_heads
+ self.hidden_act = hidden_act
+ self.use_qkv_bias = use_qkv_bias
+ self.use_qk_norm = use_qk_norm
+ self.use_bias = use_bias
+ self.norm_head = norm_head
+ self.rms_norm_eps = rms_norm_eps
+ self.embedding_dropout = embedding_dropout
+ self.attention_dropout = attention_dropout
+ self.output_dropout = output_dropout
+ self.initializer_range = initializer_range
+ self.max_position_embeddings = max_position_embeddings
+ self.rope_theta = rope_theta
+ self.use_cache = use_cache
+ self.use_sliding_window = use_sliding_window
+ self.sliding_window = sliding_window
+ self.max_window_layers = max_window_layers
+ self.head_dim = head_dim or hidden_size // num_attention_heads
+ self.rope_scaling = rope_scaling
+ self.num_experts = num_experts
+ self.num_shared_experts = num_shared_experts
+ self.num_experts_per_tok = num_experts_per_tok
+ self.n_group = n_group
+ self.topk_group = topk_group
+ self.moe_intermediate_size = moe_intermediate_size
+ self.first_k_dense_replace = first_k_dense_replace
+ self.output_router_logits = output_router_logits
+ self.routed_scaling_factor = routed_scaling_factor
+ self.partial_rotary_factor = partial_rotary_factor
+ self.router_type = router_type
+ self.norm_topk_prob = norm_topk_prob
+ self.moe_router_enable_expert_bias = moe_router_enable_expert_bias
+ self.use_interleaved_frame_timestamp = use_interleaved_frame_timestamp
+ super().__init__(
+ pad_token_id=pad_token_id,
+ tie_word_embeddings=tie_word_embeddings,
+ **kwargs,
+ )
+ self._attn_implementation = _attn_implementation
+
+
+class WhisperEncoderConfig(PretrainedConfig):
+ def __init__(
+ self,
+ whisper_encoder_config: dict | None = None,
+ ds_kernel_size=3,
+ ds_stride=2,
+ **kwargs,
+ ):
+ self.whisper_encoder_config = whisper_encoder_config
+ self.ds_kernel_size = ds_kernel_size
+ self.ds_stride = ds_stride
+ super().__init__(**kwargs)
+
+
+class BailingMM2Config(PretrainedConfig):
+ model_type = "bailingmm_moe_v2_lite"
+
+ def __init__(
+ self,
+ mlp_depth=1,
+ llm_config=None,
+ vision_config=None,
+ audio_config=None,
+ mrope_section=None,
+ **kwargs,
+ ):
+ if isinstance(audio_config, dict):
+ audio_config = WhisperEncoderConfig(**audio_config)
+ elif audio_config is not None and not isinstance(
+ audio_config, WhisperEncoderConfig
+ ):
+ raise TypeError(
+ "audio_config must be a dict, WhisperEncoderConfig, or None; "
+ f"got {type(audio_config).__name__}"
+ )
+ self.audio_config = audio_config
+
+ if isinstance(vision_config, dict):
+ vision_config = Qwen3VLMoeVisionConfig(**vision_config)
+ elif vision_config is None:
+ vision_config = Qwen3VLMoeVisionConfig()
+ self.vision_config = vision_config
+
+ if isinstance(llm_config, dict):
+ llm_config = BailingMoeV2Config(**llm_config)
+ elif llm_config is None:
+ llm_config = BailingMoeV2Config()
+ self.llm_config = llm_config
+ self.mlp_depth = mlp_depth
+
+ if mrope_section is None:
+ mrope_section = llm_config.rope_parameters.get("mrope_section", [8, 12, 12])
+ self.mrope_section = mrope_section
+ llm_config.rope_parameters.update(
+ rope_type="default",
+ mrope_section=mrope_section,
+ video_rope=True,
+ )
+ super().__init__(**kwargs)
diff --git a/python/sglang/srt/configs/hybrid_arch.py b/python/sglang/srt/configs/hybrid_arch.py
index e821792c0..5b364d0b8 100644
--- a/python/sglang/srt/configs/hybrid_arch.py
+++ b/python/sglang/srt/configs/hybrid_arch.py
@@ -26,6 +26,7 @@ from sglang.srt.configs import (
Qwen3NextConfig,
ZayaConfig,
)
+from sglang.srt.utils.hf_transformers.common import get_hf_text_config
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
@@ -113,6 +114,9 @@ def kimi_linear_config(model_config: ModelConfig):
return config
if isinstance(config, BailingHybridConfig) and config.use_kda:
return config
+ text_config = get_hf_text_config(config)
+ if isinstance(text_config, BailingHybridConfig) and text_config.use_kda:
+ return text_config
text_config = getattr(config, "text_config", None)
if isinstance(text_config, KimiLinearConfig):
return text_config
diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py
index e7ce9277a..bff067b17 100644
--- a/python/sglang/srt/configs/model_config.py
+++ b/python/sglang/srt/configs/model_config.py
@@ -26,6 +26,7 @@ import torch
from transformers import PretrainedConfig
from sglang.srt.arg_groups.overrides import resolving_view
+from sglang.srt.configs.bailing_hybrid import is_bailing_multi_gate_enabled
from sglang.srt.configs.embedding_model_spec import resolve_embedding_model_spec
from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_config
from sglang.srt.environ import envs
@@ -51,6 +52,14 @@ MIMO_V2_MODEL_ARCHS = (
)
MIMO_V2_MULTIMODAL_ARCHS = ("MiMoV2ForCausalLM",)
+BAILING_MULTI_GATE_MM_ARCHS = frozenset(
+ {
+ "BailingMMNativeForConditionalGeneration",
+ "BailingMM2NativeForConditionalGeneration",
+ "BailingMoeV3VLForConditionalGeneration",
+ }
+)
+
SWA_SINK_ARCHS = frozenset(
{
"GptOssForCausalLM",
@@ -66,6 +75,17 @@ def _quant_config_to_dict(quant_config):
return quant_config
+def requires_mm_token_modalities(
+ model_architectures: Optional[List[str]], hf_text_config: PretrainedConfig
+) -> bool:
+ """Whether a Bailing multimodal wrapper uses modality-specific routers."""
+ return bool(
+ model_architectures
+ and any(arch in BAILING_MULTI_GATE_MM_ARCHS for arch in model_architectures)
+ and is_bailing_multi_gate_enabled(hf_text_config)
+ )
+
+
def unwrap_modelopt_quantization_config(quant_config: dict) -> dict:
quantization = quant_config.get("quantization", quant_config)
if not isinstance(quantization, dict):
@@ -450,6 +470,9 @@ class ModelConfig:
)
)
self.hf_text_config = get_hf_text_config(self.hf_config)
+ self.requires_mm_token_modalities = requires_mm_token_modalities(
+ self.hf_config.architectures, self.hf_text_config
+ )
self.is_embedding_gemma = is_embedding_gemma(self.hf_text_config)
self.embedding_model_spec = resolve_embedding_model_spec(
self.hf_config.architectures,
@@ -1208,15 +1231,20 @@ class ModelConfig:
self.qk_rope_head_dim = self.hf_text_config.qk_rope_head_dim
self.v_head_dim = self.hf_config.v_head_dim
self._init_mla_scaling(self.hf_config.rope_scaling)
- elif "BailingMoeV3ForCausalLM" in self.hf_config.architectures:
+ elif (
+ "BailingMoeV3ForCausalLM" in self.hf_config.architectures
+ or "BailingMoeV3VLForConditionalGeneration" in self.hf_config.architectures
+ ):
self.head_dim = 128
self.attention_arch = AttentionArch.MLA
- self.kv_lora_rank = self.hf_config.kv_lora_rank
+ self.kv_lora_rank = self.hf_text_config.kv_lora_rank
self.qk_rope_head_dim = (
- 0 if self.hf_config.use_mla_nope else self.hf_config.qk_rope_head_dim
+ 0
+ if getattr(self.hf_text_config, "use_mla_nope", False)
+ else self.hf_text_config.qk_rope_head_dim
)
- self.v_head_dim = self.hf_config.v_head_dim
- self.qk_nope_head_dim = self.hf_config.qk_nope_head_dim
+ self.v_head_dim = self.hf_text_config.v_head_dim
+ self.qk_nope_head_dim = self.hf_text_config.qk_nope_head_dim
self.scaling = 1 / math.sqrt(self.qk_nope_head_dim + self.qk_rope_head_dim)
elif (
"SarvamMLAForCausalLM" in self.hf_config.architectures
@@ -2176,6 +2204,9 @@ multimodal_model_archs = [
"StepVLForConditionalGeneration",
"Step3p7ForConditionalGeneration",
"KimiK25ForConditionalGeneration",
+ "BailingMMNativeForConditionalGeneration",
+ "BailingMM2NativeForConditionalGeneration",
+ "BailingMoeV3VLForConditionalGeneration",
]
piecewise_cuda_graph_disabled_model_archs = [
diff --git a/python/sglang/srt/configs/qwen3_vl.py b/python/sglang/srt/configs/qwen3_vl.py
index 85068b5a6..21cb968ce 100644
--- a/python/sglang/srt/configs/qwen3_vl.py
+++ b/python/sglang/srt/configs/qwen3_vl.py
@@ -477,6 +477,7 @@ class Qwen3VLMoeVisionConfig(PretrainedConfig):
num_position_embeddings=2304,
deepstack_visual_indexes=[8, 16, 24],
initializer_range=0.02,
+ disable_merger_proj=False,
**kwargs,
):
super().__init__(**kwargs)
@@ -494,6 +495,7 @@ class Qwen3VLMoeVisionConfig(PretrainedConfig):
self.num_position_embeddings = num_position_embeddings
self.initializer_range = initializer_range
self.deepstack_visual_indexes = deepstack_visual_indexes
+ self.disable_merger_proj = disable_merger_proj
class Qwen3VLMoeConfig(PretrainedConfig):
diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py
index 36b8aa92d..acb633454 100644
--- a/python/sglang/srt/layers/moe/topk.py
+++ b/python/sglang/srt/layers/moe/topk.py
@@ -616,6 +616,7 @@ class TopK(BaseFusedOp):
*,
num_token_non_padded: Optional[torch.Tensor] = None,
expert_location_dispatch_info: Optional[ExpertLocationDispatchInfo] = None,
+ dynamic_expert_bias: Optional[torch.Tensor] = None,
) -> TopKOutput:
self.topk_config.torch_native = True
topk_output = select_experts(
@@ -625,6 +626,7 @@ class TopK(BaseFusedOp):
topk_config=self.topk_config,
num_token_non_padded=num_token_non_padded,
expert_location_dispatch_info=expert_location_dispatch_info,
+ dynamic_expert_bias=dynamic_expert_bias,
)
return self._apply_waterfill(topk_output, hidden_states.shape[0])
@@ -635,8 +637,11 @@ class TopK(BaseFusedOp):
*,
num_token_non_padded: Optional[torch.Tensor] = None,
expert_location_dispatch_info: Optional[ExpertLocationDispatchInfo] = None,
+ dynamic_expert_bias: Optional[torch.Tensor] = None,
) -> TopKOutput:
- if self.topk_config.output_format is not None:
+ if dynamic_expert_bias is not None:
+ output_format = TopKOutputFormat.STANDARD
+ elif self.topk_config.output_format is not None:
output_format = self.topk_config.output_format
elif get_moe_runner_backend().is_triton_kernels():
output_format = TopKOutputFormat.TRITON_KERNEL
@@ -699,6 +704,7 @@ class TopK(BaseFusedOp):
topk_config=self.topk_config,
num_token_non_padded=num_token_non_padded,
expert_location_dispatch_info=expert_location_dispatch_info,
+ dynamic_expert_bias=dynamic_expert_bias,
)
return self._apply_waterfill(topk_output, hidden_states.shape[0])
@@ -709,6 +715,7 @@ class TopK(BaseFusedOp):
*,
num_token_non_padded: Optional[torch.Tensor] = None,
expert_location_dispatch_info: Optional[ExpertLocationDispatchInfo] = None,
+ dynamic_expert_bias: Optional[torch.Tensor] = None,
) -> TopKOutput:
topk_output = select_experts(
hidden_states=hidden_states,
@@ -717,6 +724,7 @@ class TopK(BaseFusedOp):
topk_config=self.topk_config,
num_token_non_padded=num_token_non_padded,
expert_location_dispatch_info=expert_location_dispatch_info,
+ dynamic_expert_bias=dynamic_expert_bias,
)
return self._apply_waterfill(topk_output, hidden_states.shape[0])
@@ -727,7 +735,20 @@ class TopK(BaseFusedOp):
*,
num_token_non_padded: Optional[torch.Tensor] = None,
expert_location_dispatch_info: Optional[ExpertLocationDispatchInfo] = None,
+ dynamic_expert_bias: Optional[torch.Tensor] = None,
) -> TopKOutput:
+ if dynamic_expert_bias is not None:
+ self.topk_config.torch_native = False
+ return select_experts(
+ hidden_states=hidden_states,
+ layer_id=self.layer_id,
+ router_logits=router_logits,
+ topk_config=self.topk_config,
+ num_token_non_padded=num_token_non_padded,
+ expert_location_dispatch_info=expert_location_dispatch_info,
+ dynamic_expert_bias=dynamic_expert_bias,
+ )
+
from sglang.srt.hardware_backend.npu.moe.topk import fused_topk_npu
return fused_topk_npu(
@@ -1313,7 +1334,9 @@ def biased_topk_impl(
num_token = scores.shape[0]
num_experts = scores.shape[1]
- scores_for_choice = scores.view(num_token, -1) + correction_bias.unsqueeze(0)
+ scores_for_choice = scores.view(num_token, -1) + correction_bias.view(
+ -1, num_experts
+ )
_, topk_ids = torch.topk(
scores_for_choice,
k=topk,
@@ -1473,7 +1496,9 @@ def biased_grouped_topk_impl(
scores = gating_output.sigmoid()
num_token = scores.shape[0]
num_experts = scores.shape[1]
- scores_for_choice = scores.view(num_token, -1) + correction_bias.unsqueeze(0)
+ scores_for_choice = scores.view(num_token, -1) + correction_bias.view(
+ -1, num_experts
+ )
group_scores = (
scores_for_choice.view(num_token, num_expert_group, -1)
.topk(2, dim=-1)[0]
@@ -1619,13 +1644,18 @@ def biased_grouped_topk_gpu(
experts_per_group = (
num_experts // num_expert_group if num_expert_group else num_experts
)
+ dynamic_bias = correction_bias.ndim == 2
# topk for routed experts only (shared experts are appended separately below)
topk_routed = topk - num_fused_shared_experts
+ # The JIT router accepts one shared bias vector, not per-token bias rows.
if (
- (_is_cuda and num_expert_group and num_expert_group > 1)
- # ROCm also admits single-group routing; CUDA's condition is unchanged.
- or (_is_hip and num_expert_group)
+ not dynamic_bias
+ and (
+ (_is_cuda and num_expert_group and num_expert_group > 1)
+ # ROCm also admits single-group routing; CUDA's condition is unchanged.
+ or (_is_hip and num_expert_group)
+ )
) and envs.SGLANG_OPT_USE_JIT_KERNEL_GROUPED_TOPK.get():
# Opt-in: unified Triton router for DeepSeek-V3 grouped routing. Bit-exact
# with the flashinfer/AOT paths on DeepSeek-V3.2 e2e (validated); handles any
@@ -1671,6 +1701,7 @@ def biased_grouped_topk_gpu(
if num_expert_group > 1
else num_experts <= 384
)
+ and not dynamic_bias
):
# Pre-allocate output tensors (flashinfer mutates them in-place)
topk_weights = torch.empty(
@@ -1713,7 +1744,7 @@ def biased_grouped_topk_gpu(
return topk_weights, topk_ids
- elif _is_cuda and num_expert_group > 1:
+ elif _is_cuda and num_expert_group > 1 and not dynamic_bias:
# CUDA grouped fallback (flashinfer unavailable / constraints unmet): the
# unified Triton router replaces the retired AOT moe_fused_gate kernel. It
# handles any experts-per-group (no MAX_VPT=32 cap) and any num_experts.
@@ -1811,6 +1842,7 @@ def biased_grouped_topk_gpu(
# needs experts<=512 + topk<=8.
_jit_gate_ok = (
_is_cuda
+ and not dynamic_bias
and num_expert_group == 1
and (topk_group is None or topk_group == 1)
and (
@@ -2160,19 +2192,16 @@ def _post_process_topk_ids(
topk_ids, expert_location_dispatch_info, log2phy_prob
)
_mask_topk_ids_padded_region(topk_ids, num_token_non_padded)
- elif use_per_rank_shared_slots:
- # Shared experts appended as extra columns in topk_ids: their value
- # would be out-of-bounds for the logical-to-physical dispatch table,
- # so split, dispatch the routed cols, recombine.
+ elif num_fused_shared_experts > 0:
+ # Shared IDs are outside EPLB's routed-expert table for both global
+ # and per-rank layouts, so remap only routed columns.
shared_cols = topk_ids[:, -num_fused_shared_experts:]
routed_cols = topk_ids[:, :-num_fused_shared_experts]
routed_cols = _biased_grouped_topk_postprocess(
routed_cols, expert_location_dispatch_info, num_token_non_padded
)
topk_ids = torch.cat([routed_cols, shared_cols], dim=-1)
- # ExpertDistributionRecorder tracks EPLB physical routed experts.
- # Per-rank shared-slot remap later adds shared slots to the topk ID
- # space, so keep the routed physical ids separately for statistics.
+ # ExpertDistributionRecorder tracks only EPLB physical routed experts.
recorder_topk_ids = routed_cols
else:
topk_ids = _biased_grouped_topk_postprocess(
@@ -2323,6 +2352,7 @@ def select_experts(
layer_id: Optional[int] = None,
num_token_non_padded: Optional[torch.Tensor] = None,
expert_location_dispatch_info: Optional[ExpertLocationDispatchInfo] = None,
+ dynamic_expert_bias: Optional[torch.Tensor] = None,
) -> StandardTopKOutput:
top_k = topk_config.top_k
use_grouped_topk = topk_config.use_grouped_topk
@@ -2331,7 +2361,11 @@ def select_experts(
renormalize = topk_config.renormalize
num_fused_shared_experts = topk_config.num_fused_shared_experts
custom_routing_function = topk_config.custom_routing_function
- correction_bias = topk_config.correction_bias
+ correction_bias = (
+ topk_config.correction_bias
+ if dynamic_expert_bias is None
+ else dynamic_expert_bias
+ )
torch_native = topk_config.torch_native
routed_scaling_factor = topk_config.routed_scaling_factor
apply_routed_scaling_factor_on_output = (
@@ -2352,7 +2386,12 @@ def select_experts(
info=expert_location_dispatch_info,
)
- if _use_aiter and use_grouped_topk and correction_bias is not None:
+ if (
+ _use_aiter
+ and use_grouped_topk
+ and correction_bias is not None
+ and dynamic_expert_bias is None
+ ):
correction_bias = topk_config.correction_bias_for_dtype(router_logits.dtype)
# DeepSeek V2/V3/R1 series models use grouped_top_k
@@ -2370,7 +2409,24 @@ def select_experts(
if has_per_rank_fused_shared_slots(num_fused_shared_experts)
else num_fused_shared_experts
)
- if use_grouped_topk:
+ if dynamic_expert_bias is not None:
+ if scoring_func != "sigmoid":
+ raise ValueError(
+ "Per-token expert bias is only supported with sigmoid routing"
+ )
+ topk_weights, topk_ids = biased_grouped_topk_impl(
+ hidden_states=hidden_states,
+ gating_output=router_logits,
+ correction_bias=correction_bias,
+ topk=top_k,
+ renormalize=renormalize,
+ num_expert_group=num_expert_group or 1,
+ topk_group=topk_group or 1,
+ num_fused_shared_experts=num_fused_shared_experts,
+ routed_scaling_factor=routed_scaling_factor,
+ apply_routed_scaling_factor_on_output=apply_routed_scaling_factor_on_output,
+ )
+ elif use_grouped_topk:
assert topk_group is not None
assert num_expert_group is not None
if correction_bias is None:
diff --git a/python/sglang/srt/layers/multi_gate.py b/python/sglang/srt/layers/multi_gate.py
new file mode 100644
index 000000000..8036fef11
--- /dev/null
+++ b/python/sglang/srt/layers/multi_gate.py
@@ -0,0 +1,233 @@
+from typing import Optional, Tuple
+
+import torch
+import triton
+import triton.language as tl
+
+from sglang.srt.managers.schedule_batch import Modality
+
+MULTI_GATE_BLOCK_M = 64
+TEXT_MODALITY = 0
+VISION_MODALITY = 1
+AUDIO_MODALITY = 2
+
+
+@triton.jit
+def multi_gate_kernel(
+ hidden_states_ptr,
+ router_logits_ptr,
+ expert_bias_ptr,
+ text_gate_ptr,
+ image_gate_ptr,
+ audio_gate_ptr,
+ text_bias_ptr,
+ image_bias_ptr,
+ audio_bias_ptr,
+ token_indices_ptr,
+ modality_ids_ptr,
+ num_valid_tokens: tl.constexpr,
+ compute_type: tl.constexpr,
+ stride_am: tl.constexpr,
+ stride_ak: tl.constexpr,
+ stride_bk: tl.constexpr,
+ stride_bn: tl.constexpr,
+ stride_cm: tl.constexpr,
+ stride_cn: tl.constexpr,
+ M: tl.constexpr,
+ N: tl.constexpr,
+ K: tl.constexpr,
+ BLOCK_SIZE_M: tl.constexpr,
+ BLOCK_SIZE_N: tl.constexpr,
+ BLOCK_SIZE_K: tl.constexpr,
+ GROUP_SIZE_M: tl.constexpr,
+):
+ pid = tl.program_id(axis=0)
+ num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
+ num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
+ num_pid_in_group = GROUP_SIZE_M * num_pid_n
+ group_id = pid // num_pid_in_group
+ first_pid_m = group_id * GROUP_SIZE_M
+ group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
+ pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
+ pid_n = (pid % num_pid_in_group) // group_size_m
+
+ offs = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
+ token_indices = tl.load(token_indices_ptr + offs)
+ token_mask = token_indices < num_valid_tokens
+ modality_id = tl.load(modality_ids_ptr + pid_m).to(tl.int64)
+
+ if modality_id == VISION_MODALITY:
+ gate_ptr = image_gate_ptr
+ bias_ptr = image_bias_ptr
+ elif modality_id == AUDIO_MODALITY:
+ gate_ptr = audio_gate_ptr
+ bias_ptr = audio_bias_ptr
+ else:
+ gate_ptr = text_gate_ptr
+ bias_ptr = text_bias_ptr
+
+ offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)) % N
+ offs_k = tl.arange(0, BLOCK_SIZE_K)
+ hidden_ptrs = hidden_states_ptr + (
+ token_indices[:, None] * stride_am + offs_k[None, :] * stride_ak
+ )
+ gate_ptrs = gate_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn
+
+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
+ for k_start in range(0, K, BLOCK_SIZE_K):
+ hidden = tl.load(
+ hidden_ptrs,
+ mask=token_mask[:, None] & (offs_k[None, :] < K - k_start),
+ other=0.0,
+ ).to(compute_type)
+ gate = tl.load(gate_ptrs, mask=offs_k[:, None] < K - k_start, other=0.0)
+ accumulator += tl.dot(hidden, gate)
+ hidden_ptrs += BLOCK_SIZE_K * stride_ak
+ gate_ptrs += BLOCK_SIZE_K * stride_bk
+
+ output_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
+ output_offsets = stride_cm * token_indices[:, None] + stride_cn * output_n[None, :]
+ output_mask = token_mask[:, None] & (output_n[None, :] < N)
+ tl.store(router_logits_ptr + output_offsets, accumulator, mask=output_mask)
+
+ bias = tl.load(bias_ptr + output_n[None, :], mask=output_n[None, :] < N, other=0.0)
+ tl.store(expert_bias_ptr + output_offsets, bias, mask=output_mask)
+
+
+def _kernel_config(num_tokens: int) -> dict:
+ configs = {
+ 1024: (64, 128, 64, 64, 4, 3),
+ 2048: (64, 32, 128, 1, 8, 3),
+ 4096: (64, 64, 128, 32, 4, 3),
+ 8192: (64, 32, 128, 64, 8, 3),
+ }
+ key = min(configs, key=lambda candidate: abs(candidate - num_tokens))
+ block_m, block_n, block_k, group_m, num_warps, num_stages = configs[key]
+ return {
+ "BLOCK_SIZE_M": block_m,
+ "BLOCK_SIZE_N": block_n,
+ "BLOCK_SIZE_K": block_k,
+ "GROUP_SIZE_M": group_m,
+ "num_warps": num_warps,
+ "num_stages": num_stages,
+ }
+
+
+def create_multi_gate_mm_indices(
+ token_modalities: torch.Tensor,
+) -> Tuple[torch.Tensor, torch.Tensor]:
+ """Group token indices into padded blocks tagged by modality."""
+ if token_modalities.ndim != 1:
+ raise ValueError(
+ f"token_modalities must be one-dimensional, got {token_modalities.shape=}"
+ )
+ modality_indices = [
+ (token_modalities == 0).nonzero(as_tuple=False).squeeze(-1),
+ (
+ (token_modalities == Modality.IMAGE.value)
+ | (token_modalities == Modality.VIDEO.value)
+ )
+ .nonzero(as_tuple=False)
+ .squeeze(-1),
+ (token_modalities == Modality.AUDIO.value).nonzero(as_tuple=False).squeeze(-1),
+ ]
+ block_counts = [
+ (indices.shape[0] + MULTI_GATE_BLOCK_M - 1) // MULTI_GATE_BLOCK_M
+ for indices in modality_indices
+ ]
+ total_tokens = token_modalities.shape[0]
+ total_blocks = sum(block_counts)
+ token_indices = torch.full(
+ (total_blocks * MULTI_GATE_BLOCK_M,),
+ total_tokens,
+ dtype=torch.int32,
+ device=token_modalities.device,
+ )
+ modality_ids = torch.empty(
+ total_blocks, dtype=torch.int32, device=token_modalities.device
+ )
+ block_offset = 0
+ for modality, (indices, block_count) in enumerate(
+ zip(modality_indices, block_counts)
+ ):
+ token_offset = block_offset * MULTI_GATE_BLOCK_M
+ token_indices[token_offset : token_offset + indices.shape[0]] = indices
+ modality_ids[block_offset : block_offset + block_count] = modality
+ block_offset += block_count
+ return token_indices, modality_ids
+
+
+@torch.compiler.disable
+def multi_gate_triton_kernel(
+ hidden_states: torch.Tensor,
+ multi_gate_indices: Tuple[torch.Tensor, torch.Tensor],
+ text_weight: torch.Tensor,
+ image_weight: torch.Tensor,
+ audio_weight: torch.Tensor,
+ text_bias: torch.Tensor,
+ image_bias: torch.Tensor,
+ audio_bias: torch.Tensor,
+ config: Optional[dict] = None,
+) -> Tuple[torch.Tensor, torch.Tensor]:
+ token_indices, modality_ids = multi_gate_indices
+ num_tokens = hidden_states.shape[0]
+ weights = (text_weight, image_weight, audio_weight)
+ biases = (text_bias, image_bias, audio_bias)
+ num_experts = text_weight.shape[0]
+ if any(weight.shape != text_weight.shape for weight in weights[1:]):
+ raise ValueError("All modality gate weights must have the same shape")
+ if any(bias is None or bias.shape != (num_experts,) for bias in biases):
+ raise ValueError(
+ "Multi-gate routing requires one expert-bias vector per modality"
+ )
+
+ transposed_weights = tuple(weight.transpose(0, 1) for weight in weights)
+ router_logits = torch.empty(
+ (num_tokens, num_experts),
+ dtype=hidden_states.dtype,
+ device=hidden_states.device,
+ )
+ dynamic_expert_bias = torch.empty(
+ (num_tokens, num_experts), dtype=torch.float32, device=hidden_states.device
+ )
+ if text_weight.dtype == torch.bfloat16:
+ compute_type = tl.bfloat16
+ elif text_weight.dtype == torch.float16:
+ compute_type = tl.float16
+ elif text_weight.dtype == torch.float32:
+ compute_type = tl.float32
+ else:
+ raise ValueError(f"Unsupported multi-gate dtype: {text_weight.dtype}")
+
+ num_grouped_tokens = token_indices.shape[0]
+ config = config or _kernel_config(num_grouped_tokens)
+ if config["BLOCK_SIZE_M"] != MULTI_GATE_BLOCK_M:
+ raise ValueError(
+ f"Multi-gate BLOCK_SIZE_M must be {MULTI_GATE_BLOCK_M}, got {config['BLOCK_SIZE_M']}"
+ )
+ grid = lambda meta: (
+ triton.cdiv(num_grouped_tokens, meta["BLOCK_SIZE_M"])
+ * triton.cdiv(num_experts, meta["BLOCK_SIZE_N"]),
+ )
+ multi_gate_kernel[grid](
+ hidden_states,
+ router_logits,
+ dynamic_expert_bias,
+ *transposed_weights,
+ *biases,
+ token_indices,
+ modality_ids,
+ num_valid_tokens=num_tokens,
+ compute_type=compute_type,
+ stride_am=hidden_states.stride(0),
+ stride_ak=hidden_states.stride(1),
+ stride_bk=transposed_weights[0].stride(0),
+ stride_bn=transposed_weights[0].stride(1),
+ stride_cm=router_logits.stride(0),
+ stride_cn=router_logits.stride(1),
+ M=num_grouped_tokens,
+ N=num_experts,
+ K=hidden_states.shape[-1],
+ **config,
+ )
+ return router_logits, dynamic_expert_bias
diff --git a/python/sglang/srt/layers/rotary_embedding/bailing_mrope.py b/python/sglang/srt/layers/rotary_embedding/bailing_mrope.py
new file mode 100644
index 000000000..d875f0899
--- /dev/null
+++ b/python/sglang/srt/layers/rotary_embedding/bailing_mrope.py
@@ -0,0 +1,366 @@
+from typing import List, Optional, Tuple, Union
+
+import torch
+from transformers.configuration_utils import PretrainedConfig
+
+from sglang.kernels.ops.attention.rotary_triton import (
+ triton_ernie45_rope_fused_inplace,
+)
+from sglang.srt.environ import envs
+from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding
+from sglang.srt.layers.rotary_embedding.yarn import (
+ yarn_find_correction_range,
+ yarn_get_mscale_simple,
+ yarn_linear_ramp_mask,
+)
+
+
+class BailingMRotaryEmbedding(RotaryEmbedding):
+ """Bailing multimodal RoPE with centered height and width positions."""
+
+ def __init__(
+ self,
+ head_size: int,
+ rotary_dim: int,
+ max_position_embeddings: int,
+ base: int,
+ is_neox_style: bool,
+ dtype: torch.dtype,
+ mrope_section: Optional[List[int]] = None,
+ video_rope: bool = False,
+ scaling_factor: float = 1.0,
+ original_max_position_embeddings: Optional[int] = None,
+ extrapolation_factor: float = 1,
+ attn_factor: float = 1,
+ beta_fast: int = 32,
+ beta_slow: int = 1,
+ truncate: bool = True,
+ ) -> None:
+ self.scaling_factor = scaling_factor
+ self.extrapolation_factor = extrapolation_factor
+ self.attn_factor = attn_factor
+ self.beta_fast = beta_fast
+ self.beta_slow = beta_slow
+ self.truncate = truncate
+ self.original_max_position_embeddings = (
+ original_max_position_embeddings or max_position_embeddings
+ )
+ self.mscale = (
+ float(yarn_get_mscale_simple(scaling_factor) * attn_factor)
+ if scaling_factor > 1
+ else 1.0
+ )
+ # Bailing positions are bounded by the checkpoint context on both sides:
+ # text/time grow positive while centered height/width can be negative.
+ # YaRN only stretches the positive side; the negative side holds small
+ # centered media coordinates and stays at the checkpoint bound.
+ position_start = -max_position_embeddings if video_rope else 0
+ cache_length = (max_position_embeddings if video_rope else 0) + int(
+ self.original_max_position_embeddings * scaling_factor
+ )
+ super().__init__(
+ head_size,
+ rotary_dim,
+ cache_length,
+ base,
+ is_neox_style,
+ dtype,
+ position_start=position_start,
+ )
+
+ if mrope_section is not None:
+ if sum(mrope_section) != rotary_dim // 2:
+ raise ValueError(
+ "mrope_section must sum to rotary_dim // 2; "
+ f"got {mrope_section=} and {rotary_dim=}"
+ )
+ # The checkpoint stores [time, height, width], while the shared
+ # Ernie4.5 kernel consumes [height, width, time].
+ mrope_section = [mrope_section[1], mrope_section[2], mrope_section[0]]
+ self.mrope_section = mrope_section
+
+ def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor:
+ if self.scaling_factor <= 1:
+ return super()._compute_inv_freq(base)
+ # YaRN blend, same construction as YaRNScalingMRotaryEmbedding.
+ pos_freqs = self.base ** (
+ torch.arange(0, self.rotary_dim, 2, dtype=torch.float) / self.rotary_dim
+ )
+ inv_freq_extrapolation = 1.0 / pos_freqs
+ inv_freq_interpolation = 1.0 / (self.scaling_factor * pos_freqs)
+ low, high = yarn_find_correction_range(
+ self.beta_fast,
+ self.beta_slow,
+ self.rotary_dim,
+ self.base,
+ self.original_max_position_embeddings,
+ self.truncate,
+ )
+ inv_freq_mask = (
+ 1
+ - yarn_linear_ramp_mask(low, high, self.rotary_dim // 2, dtype=torch.float)
+ ) * self.extrapolation_factor
+ return (
+ inv_freq_interpolation * (1 - inv_freq_mask)
+ + inv_freq_extrapolation * inv_freq_mask
+ )
+
+ def _compute_cos_sin_cache(self) -> torch.Tensor:
+ cache = super()._compute_cos_sin_cache()
+ if self.mscale != 1.0:
+ cache = cache * self.mscale
+ return cache
+
+ def _ensure_cos_sin_cache_length(self, needed_max_pos: int):
+ if self.mscale == 1.0:
+ return super()._ensure_cos_sin_cache_length(needed_max_pos)
+ cur_len = int(self.cos_sin_cache.shape[0])
+ if needed_max_pos < cur_len:
+ return
+ # The base incremental path skips mscale, so rebuild the cache in one
+ # shot to keep every row on the same scale.
+ align = envs.SGLANG_ROPE_CACHE_ALIGN.get()
+ self.max_position_embeddings = ((needed_max_pos + align) // align) * align
+ self.cos_sin_cache = self._compute_cos_sin_cache().to(
+ device=self.cos_sin_cache.device, dtype=self.cos_sin_cache.dtype
+ )
+
+ def forward(
+ self,
+ positions: torch.Tensor,
+ query: torch.Tensor,
+ key: torch.Tensor,
+ fused_set_kv_buffer_arg=None,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ if positions.ndim not in (1, 2):
+ raise ValueError(
+ f"Bailing mRoPE expects 1D or 2D positions, got {positions.shape=}"
+ )
+ positions = positions - self.position_start
+ if positions.ndim == 2:
+ if fused_set_kv_buffer_arg is not None:
+ raise ValueError(
+ "fused_set_kv_buffer_arg is not supported for Bailing mRoPE"
+ )
+ if self.mrope_section is None:
+ raise ValueError("mrope_section is required for 2D Bailing positions")
+
+ query_shape = query.shape
+ key_shape = key.shape
+ if query.ndim == 3:
+ query = query.reshape(query_shape[0], -1)
+ key = key.reshape(key_shape[0], -1)
+ triton_ernie45_rope_fused_inplace(
+ query,
+ key,
+ self.cos_sin_cache,
+ positions,
+ self.mrope_section,
+ self.head_size,
+ self.rotary_dim,
+ self.is_neox_style,
+ )
+ if query_shape != query.shape:
+ query = query.view(query_shape)
+ key = key.view(key_shape)
+ return query, key
+ return RotaryEmbedding.forward(self, positions, query, key)
+
+ @staticmethod
+ def _text_config(hf_config: PretrainedConfig) -> PretrainedConfig:
+ text_config = getattr(hf_config, "text_config", None)
+ if text_config is None:
+ text_config = getattr(hf_config, "llm_config", None)
+ if text_config is None:
+ raise ValueError("Bailing VL config must define text_config or llm_config")
+ return text_config
+
+ @staticmethod
+ def _validate_position_bounds(
+ positions: torch.Tensor, text_config: PretrainedConfig
+ ) -> None:
+ if positions.numel() == 0:
+ return
+ bound = text_config.max_position_embeddings
+ positive_bound = bound
+ rope_parameters = getattr(text_config, "rope_parameters", None) or {}
+ rope_type = rope_parameters.get("rope_type") or rope_parameters.get("type")
+ if rope_type in ("yarn", "deepseek_yarn"):
+ factor = float(rope_parameters.get("factor", 1.0))
+ original = rope_parameters.get("original_max_position_embeddings", bound)
+ positive_bound = max(bound, int(original * factor))
+ min_position = int(positions.min().item())
+ max_position = int(positions.max().item())
+ if min_position < -bound or max_position >= positive_bound:
+ raise ValueError(
+ "Bailing mRoPE position exceeds the checkpoint bounds: "
+ f"min={min_position}, max={max_position}, "
+ f"allowed=[{-bound}, {positive_bound})"
+ )
+
+ @classmethod
+ def bailing_3drope_get_input_positions_tensor(
+ cls,
+ input_ids: torch.Tensor,
+ hf_config: PretrainedConfig,
+ image_grid_thw: Union[List[List[int]], torch.Tensor, None],
+ video_grid_thw: Union[List[List[int]], torch.Tensor, None],
+ second_per_grid_ts: Optional[List[float]] = None,
+ context_len: int = 0,
+ seq_len: Optional[int] = None,
+ audio_feature_lengths: Optional[torch.Tensor] = None,
+ use_audio_in_video: bool = False,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ del context_len, seq_len, audio_feature_lengths, use_audio_in_video
+ scale_factor = 2.0
+ second_per_grid_ts = None
+ spatial_merge_size = hf_config.vision_config.spatial_merge_size
+ text_config = cls._text_config(hf_config)
+ image_patch_id = text_config.image_patch_token
+ video_patch_id = text_config.video_patch_token
+ image_start_token_id = text_config.image_start_token
+ video_start_token_id = text_config.video_start_token
+ use_interleaved_frame_timestamp = getattr(
+ text_config, "use_interleaved_frame_timestamp", False
+ )
+
+ if image_grid_thw is None and video_grid_thw is None:
+ position_ids = (
+ torch.arange(input_ids.numel(), device=input_ids.device)
+ .view(1, 1, -1)
+ .expand(3, 1, -1)
+ )
+ cls._validate_position_bounds(position_ids, text_config)
+ position_delta = torch.zeros(
+ [1, 1], device=input_ids.device, dtype=input_ids.dtype
+ )
+ return position_ids, position_delta
+
+ if video_grid_thw is not None and use_interleaved_frame_timestamp:
+ video_grid_thw = torch.as_tensor(video_grid_thw).clone()
+ video_grid_thw = torch.repeat_interleave(
+ video_grid_thw, video_grid_thw[:, 0], dim=0
+ )
+ video_grid_thw[:, 0] = 1
+
+ image_count = 0
+ video_count = 0
+ if image_grid_thw is not None:
+ starts = torch.argwhere(input_ids == image_start_token_id).squeeze(1)
+ starts = starts[starts + 1 < input_ids.numel()]
+ if starts.numel() > 0:
+ image_count = int((input_ids[starts + 1] == image_patch_id).sum())
+ if video_grid_thw is not None:
+ start_token = (
+ image_start_token_id
+ if use_interleaved_frame_timestamp
+ else video_start_token_id
+ )
+ starts = torch.argwhere(input_ids == start_token).squeeze(1)
+ starts = starts[starts + 1 < input_ids.numel()]
+ if starts.numel() > 0:
+ video_count = int((input_ids[starts + 1] == video_patch_id).sum())
+
+ input_tokens = input_ids.tolist()
+ position_chunks = []
+ start = 0
+ image_index = video_index = 0
+ remaining_images = image_count
+ remaining_videos = video_count
+ device = input_ids.device
+
+ for _ in range(image_count + video_count):
+ image_start = (
+ input_tokens.index(image_patch_id, start)
+ if image_patch_id in input_tokens[start:] and remaining_images > 0
+ else len(input_tokens) + 1
+ )
+ video_start = (
+ input_tokens.index(video_patch_id, start)
+ if video_patch_id in input_tokens[start:] and remaining_videos > 0
+ else len(input_tokens) + 1
+ )
+ if image_start < video_start:
+ t, h, w = torch.as_tensor(image_grid_thw[image_index]).tolist()
+ seconds_per_grid = 0.0
+ image_index += 1
+ remaining_images -= 1
+ media_start = image_start
+ else:
+ t, h, w = torch.as_tensor(video_grid_thw[video_index]).tolist()
+ seconds_per_grid = (
+ second_per_grid_ts[video_index]
+ if second_per_grid_ts is not None
+ else 1.0
+ )
+ video_index += 1
+ remaining_videos -= 1
+ media_start = video_start
+
+ grid_t = int(t)
+ grid_h = int(h) // spatial_merge_size
+ grid_w = int(w) // spatial_merge_size
+ text_len = media_start - start
+ position_start = (
+ int(position_chunks[-1][0].max().item()) + 1 if position_chunks else 0
+ )
+ position_chunks.append(
+ torch.arange(text_len, device=device).view(1, -1).expand(3, -1)
+ + position_start
+ )
+
+ time_index = (
+ torch.arange(grid_t, device=device)
+ .view(-1, 1)
+ .expand(-1, grid_h * grid_w)
+ .flatten()
+ )
+ height_index = (
+ torch.arange(grid_h, device=device)
+ .view(1, -1, 1)
+ .expand(grid_t, -1, grid_w)
+ .flatten()
+ - (grid_h - 1) // 2
+ )
+ width_index = (
+ torch.arange(grid_w, device=device)
+ .view(1, 1, -1)
+ .expand(grid_t, grid_h, -1)
+ .flatten()
+ - (grid_w - 1) // 2
+ )
+ if second_per_grid_ts is not None:
+ time_index = time_index * seconds_per_grid * scale_factor
+ else:
+ time_index = time_index * scale_factor
+ time_index = time_index + text_len + position_start
+ position_chunks.append(
+ torch.stack(
+ [time_index, height_index + time_index, width_index + time_index]
+ )
+ )
+ start = media_start + grid_t * grid_h * grid_w
+
+ if start < len(input_tokens):
+ position_start = (
+ int(position_chunks[-1][0].max().item()) + 1 if position_chunks else 0
+ )
+ text_len = len(input_tokens) - start
+ position_chunks.append(
+ torch.arange(text_len, device=device).view(1, -1).expand(3, -1)
+ + position_start
+ )
+
+ positions = torch.cat(position_chunks, dim=1).reshape(3, -1)
+ if positions.shape[1] != input_ids.numel():
+ raise ValueError(
+ "Bailing mRoPE media grids do not match the prompt token spans: "
+ f"positions={positions.shape[1]}, tokens={input_ids.numel()}"
+ )
+ cls._validate_position_bounds(positions, text_config)
+ position_delta = (
+ (positions[0].max() + 1 - input_ids.numel())
+ .reshape(1, 1)
+ .to(dtype=input_ids.dtype)
+ )
+ return positions.unsqueeze(1).to(dtype=input_ids.dtype), position_delta
diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py
index a8bda5656..486a9d606 100644
--- a/python/sglang/srt/layers/rotary_embedding/base.py
+++ b/python/sglang/srt/layers/rotary_embedding/base.py
@@ -94,6 +94,7 @@ class RotaryEmbedding(BaseFusedOp):
base: int,
is_neox_style: bool,
dtype: torch.dtype,
+ position_start: int = 0,
) -> None:
super().__init__()
self.head_size = head_size
@@ -102,6 +103,7 @@ class RotaryEmbedding(BaseFusedOp):
self.base = base
self.is_neox_style = is_neox_style
self.dtype = dtype
+ self.position_start = position_start
self._force_native = (
publish_role() is not None
and get_exec().deterministic.rl_on_policy_target is not None
@@ -181,7 +183,11 @@ class RotaryEmbedding(BaseFusedOp):
def _compute_cos_sin_cache(self) -> torch.Tensor:
"""Compute the cos and sin cache."""
inv_freq = self._compute_inv_freq(self.base)
- t = torch.arange(self.max_position_embeddings, dtype=torch.float)
+ t = torch.arange(
+ self.position_start,
+ self.position_start + self.max_position_embeddings,
+ dtype=torch.float,
+ )
freqs = torch.einsum("i,j -> ij", t, inv_freq)
cos = freqs.cos()
@@ -205,8 +211,9 @@ class RotaryEmbedding(BaseFusedOp):
inv_freq = self._compute_inv_freq(self.base).to(device=device)
# Incremental computation for new positions only
- start = cur_len
- t_new = torch.arange(start, new_len, dtype=inv_freq.dtype, device=device)
+ start = self.position_start + cur_len
+ end = self.position_start + new_len
+ t_new = torch.arange(start, end, dtype=inv_freq.dtype, device=device)
if t_new.numel() == 0:
return
diff --git a/python/sglang/srt/layers/rotary_embedding/factory.py b/python/sglang/srt/layers/rotary_embedding/factory.py
index 32a9f0ecc..56de48197 100644
--- a/python/sglang/srt/layers/rotary_embedding/factory.py
+++ b/python/sglang/srt/layers/rotary_embedding/factory.py
@@ -8,6 +8,7 @@ from typing import Any, Dict, Optional, Tuple
import torch
+from sglang.srt.layers.rotary_embedding.bailing_mrope import BailingMRotaryEmbedding
from sglang.srt.layers.rotary_embedding.base import (
LinearScalingRotaryEmbedding,
RotaryEmbedding,
@@ -52,6 +53,21 @@ def _get_rope_param(rope_scaling, key, default, scaling_type):
return default
+def _bailing_yarn_kwargs(rope_scaling: Dict[str, Any], max_position: int) -> Dict:
+ """YaRN overrides for BailingMRotaryEmbedding; factor=1.0 is a no-op."""
+ return {
+ "scaling_factor": rope_scaling.get("factor", 1.0),
+ "original_max_position_embeddings": rope_scaling.get(
+ "original_max_position_embeddings", max_position
+ ),
+ "extrapolation_factor": rope_scaling.get("extrapolation_factor", 1),
+ "attn_factor": rope_scaling.get("attn_factor", 1),
+ "beta_fast": rope_scaling.get("beta_fast", 32),
+ "beta_slow": rope_scaling.get("beta_slow", 1),
+ "truncate": rope_scaling.get("truncate", True),
+ }
+
+
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
@@ -209,7 +225,21 @@ def get_rope(
original_max_position,
)
elif scaling_type == "default":
- if "mrope_section" in rope_scaling:
+ if "mrope_section" in rope_scaling and rope_scaling.get(
+ "video_rope", False
+ ):
+ rotary_emb = BailingMRotaryEmbedding(
+ head_size,
+ rotary_dim,
+ max_position,
+ base,
+ is_neox_style,
+ dtype,
+ mrope_section=rope_scaling["mrope_section"],
+ video_rope=True,
+ **_bailing_yarn_kwargs(rope_scaling, max_position),
+ )
+ elif "mrope_section" in rope_scaling:
rotary_emb = MRotaryEmbedding(
head_size,
rotary_dim,
@@ -300,7 +330,21 @@ def get_rope(
)
}
extra_kwargs["truncate"] = rope_scaling.get("truncate", True)
- if "mrope_section" in rope_scaling:
+ if "mrope_section" in rope_scaling and rope_scaling.get(
+ "video_rope", False
+ ):
+ rotary_emb = BailingMRotaryEmbedding(
+ head_size,
+ rotary_dim,
+ max_position,
+ base,
+ is_neox_style,
+ dtype,
+ mrope_section=rope_scaling["mrope_section"],
+ video_rope=True,
+ **_bailing_yarn_kwargs(rope_scaling, max_position),
+ )
+ elif "mrope_section" in rope_scaling:
rotary_emb = YaRNScalingMRotaryEmbedding(
head_size,
rotary_dim,
diff --git a/python/sglang/srt/layers/rotary_embedding/mrope_rope_index.py b/python/sglang/srt/layers/rotary_embedding/mrope_rope_index.py
index f0e17f3a2..a922f7b3c 100644
--- a/python/sglang/srt/layers/rotary_embedding/mrope_rope_index.py
+++ b/python/sglang/srt/layers/rotary_embedding/mrope_rope_index.py
@@ -78,6 +78,7 @@ def get_rope_index(
or model_type.startswith("interns2_mobius")
or model_type.startswith("cosmos3_omni")
or model_type.startswith("cosmos3_edge")
+ or model_type == "bailing_moe_v3_vl"
) and video_grid_thw is not None:
video_grid_thw = torch.repeat_interleave(
video_grid_thw, video_grid_thw[:, 0], dim=0
diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py
index 6c4e7b9e8..5297d9a84 100755
--- a/python/sglang/srt/managers/schedule_batch.py
+++ b/python/sglang/srt/managers/schedule_batch.py
@@ -649,6 +649,35 @@ class MultimodalProcessorOutput(
padded_input_ids[start : end + 1] = [item.pad_value] * (end - start + 1)
return padded_input_ids
+ @staticmethod
+ def build_token_modalities(
+ input_ids, mm_items: List[MultimodalDataItem]
+ ) -> Optional[List[int]]:
+ """Build the pre-padding token modality map from item offsets."""
+ if input_ids is None or not mm_items:
+ return None
+ if isinstance(input_ids, torch.Tensor):
+ num_tokens = input_ids.numel()
+ else:
+ num_tokens = len(flatten_nested_list(input_ids))
+ token_modalities = [0] * num_tokens
+ for item in mm_items:
+ if not item.offsets:
+ continue
+ modality = item.modality.value
+ for start, end in item.offsets:
+ if start < 0 or end < start or end >= num_tokens:
+ raise ValueError(
+ "Invalid multimodal token offsets: "
+ f"offset=({start}, {end}), num_tokens={num_tokens}"
+ )
+ if any(token_modalities[index] for index in range(start, end + 1)):
+ raise ValueError(
+ f"Overlapping multimodal token offsets at ({start}, {end})"
+ )
+ token_modalities[start : end + 1] = [modality] * (end - start + 1)
+ return token_modalities
+
@dataclasses.dataclass
class MultimodalInputs:
@@ -659,6 +688,7 @@ class MultimodalInputs:
padded_input_ids: Optional[List[int]] = None
image_pad_len: Optional[list] = None
num_image_tokens: Optional[int] = None
+ token_modalities: Optional[List[int]] = None
# image
im_token_id: Optional[int] = None
@@ -701,7 +731,9 @@ class MultimodalInputs:
item.feature = None
@staticmethod
- def from_processor_output(obj: MultimodalProcessorOutput):
+ def from_processor_output(
+ obj: MultimodalProcessorOutput, *, requires_mm_token_modalities: bool = False
+ ):
mm_items = obj.mm_items
assert isinstance(mm_items, list)
mm_items = [item for item in mm_items if item.is_valid()]
@@ -742,6 +774,12 @@ class MultimodalInputs:
if isinstance(item.feature, torch.Tensor):
item.feature = try_add_to_buffer(item.feature)
+ token_modalities = (
+ MultimodalProcessorOutput.build_token_modalities(obj.input_ids, mm_items)
+ if requires_mm_token_modalities
+ else None
+ )
+
for item in mm_items:
item.set_pad_value()
@@ -753,6 +791,7 @@ class MultimodalInputs:
mm_inputs = MultimodalInputs(
mm_items=mm_items,
padded_input_ids=obj.padded_input_ids,
+ token_modalities=token_modalities,
)
optional_args = [
"mrope_positions",
@@ -820,6 +859,12 @@ class MultimodalInputs:
if self_arg is not None:
setattr(self, arg, self_arg + getattr(other, arg))
+ if other.token_modalities is not None:
+ if self.token_modalities is None:
+ self.token_modalities = list(other.token_modalities)
+ else:
+ self.token_modalities += other.token_modalities
+
mrope_positions = self.mrope_positions
if mrope_positions is not None:
if other.mrope_positions is None:
diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py
index cecc87100..87d070392 100644
--- a/python/sglang/srt/managers/scheduler.py
+++ b/python/sglang/srt/managers/scheduler.py
@@ -978,11 +978,10 @@ class Scheduler(
initialize_mamba_selective_state_update_backend(self.server_args)
def init_moe_gemm_config(self):
- config_to_check = self.model_config.hf_config
- if hasattr(self.model_config.hf_config, "text_config"):
- config_to_check = self.model_config.hf_config.text_config
- elif hasattr(self.model_config, "hf_text_config"):
- config_to_check = self.model_config.hf_text_config
+ # Use the language config already normalized by ModelConfig. Multimodal
+ # wrappers expose it under different attributes (for example,
+ # ``text_config`` or ``llm_config``).
+ config_to_check = self.model_config.hf_text_config
# Different MoE architectures expose the per-token expert count under
# different attribute names (e.g. Gemma4 uses ``top_k_experts``,
@@ -2167,7 +2166,8 @@ class Scheduler(
tokenized_req.mm_inputs, MultimodalInputs
):
tokenized_req.mm_inputs = MultimodalInputs.from_processor_output(
- tokenized_req.mm_inputs
+ tokenized_req.mm_inputs,
+ requires_mm_token_modalities=self.model_config.requires_mm_token_modalities,
)
except Exception as error:
local_error = f"{type(error).__name__}: {error}"
@@ -2602,7 +2602,10 @@ class Scheduler(
if self.dp_tp_group.rank_in_group == 0:
try:
result = _MultimodalInputBroadcast(
- inputs=MultimodalInputs.from_processor_output(raw_mm_inputs)
+ inputs=MultimodalInputs.from_processor_output(
+ raw_mm_inputs,
+ requires_mm_token_modalities=self.model_config.requires_mm_token_modalities,
+ )
)
except Exception as error:
result = _MultimodalInputBroadcast(
@@ -2633,7 +2636,10 @@ class Scheduler(
result = obj_list[0]
else:
result = _MultimodalInputBroadcast(
- inputs=MultimodalInputs.from_processor_output(raw_mm_inputs)
+ inputs=MultimodalInputs.from_processor_output(
+ raw_mm_inputs,
+ requires_mm_token_modalities=self.model_config.requires_mm_token_modalities,
+ )
)
if result.error is not None:
@@ -2648,7 +2654,10 @@ class Scheduler(
if get_mm().enable_broadcast_mm_inputs_process:
return self._process_and_broadcast_mm_inputs(mm_inputs)
- return MultimodalInputs.from_processor_output(mm_inputs)
+ return MultimodalInputs.from_processor_output(
+ mm_inputs,
+ requires_mm_token_modalities=self.model_config.requires_mm_token_modalities,
+ )
@staticmethod
def _try_apply_padded_mm_input_ids(recv_req, req, image_inputs) -> bool:
diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py
index 6c0009840..b3d7dc097 100644
--- a/python/sglang/srt/model_executor/forward_batch_info.py
+++ b/python/sglang/srt/model_executor/forward_batch_info.py
@@ -69,6 +69,7 @@ from sglang.srt.utils import (
from sglang.srt.utils.common import ceil_align, is_pin_memory_available
if TYPE_CHECKING:
+ from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.layers.cp.base import BaseContextParallelMetadata
from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
@@ -85,6 +86,74 @@ _is_npu = is_npu()
_is_cpu = is_cpu()
+def _build_forward_token_modalities(
+ mm_inputs: Optional[List[MultimodalInputs]],
+ extend_prefix_lens: Optional[List[int]],
+ extend_seq_lens: Optional[List[int]],
+ num_tokens: int,
+ device: torch.device,
+) -> Optional[torch.Tensor]:
+ if not mm_inputs or extend_prefix_lens is None or extend_seq_lens is None:
+ return None
+ if not (len(mm_inputs) == len(extend_prefix_lens) == len(extend_seq_lens)):
+ raise ValueError(
+ "Multimodal metadata batch dimensions do not match: "
+ f"mm_inputs={len(mm_inputs)}, prefixes={len(extend_prefix_lens)}, "
+ f"extend_lens={len(extend_seq_lens)}"
+ )
+
+ modalities = []
+ has_multimodal_tokens = False
+ for mm_input, prefix_len, extend_len in zip(
+ mm_inputs, extend_prefix_lens, extend_seq_lens
+ ):
+ if mm_input is None or mm_input.token_modalities is None:
+ modalities.extend([0] * extend_len)
+ continue
+ end = prefix_len + extend_len
+ request_modalities = mm_input.token_modalities[prefix_len:end]
+ if len(request_modalities) != extend_len:
+ raise ValueError(
+ "Multimodal token metadata is shorter than the active forward span: "
+ f"prefix_len={prefix_len}, extend_len={extend_len}, "
+ f"metadata_len={len(mm_input.token_modalities)}"
+ )
+ has_multimodal_tokens |= any(request_modalities)
+ modalities.extend(request_modalities)
+
+ if len(modalities) != num_tokens:
+ raise ValueError(
+ "Multimodal token metadata does not match the forward batch: "
+ f"metadata_tokens={len(modalities)}, forward_tokens={num_tokens}"
+ )
+ if not has_multimodal_tokens:
+ return None
+ return torch.tensor(
+ modalities,
+ dtype=torch.int8,
+ pin_memory=is_pin_memory_available(device),
+ ).to(device, non_blocking=True)
+
+
+def _maybe_build_forward_token_modalities(
+ model_config: ModelConfig,
+ mm_inputs: Optional[List[MultimodalInputs]],
+ extend_prefix_lens: Optional[List[int]],
+ extend_seq_lens: Optional[List[int]],
+ num_tokens: int,
+ device: torch.device,
+) -> Optional[torch.Tensor]:
+ if not model_config.requires_mm_token_modalities:
+ return None
+ return _build_forward_token_modalities(
+ mm_inputs,
+ extend_prefix_lens,
+ extend_seq_lens,
+ num_tokens,
+ device,
+ )
+
+
def _elastic_should_preserve_local_token_counts(
*,
model_runner: ModelRunner,
@@ -483,6 +552,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# For multimodal
mm_inputs: Optional[List[MultimodalInputs]] = None
+ mm_token_modalities: Optional[torch.Tensor] = None
+ multi_gate_indices: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
# Encoder-decoder host fields
encoder_cached: Optional[List[bool]] = None
@@ -869,6 +940,15 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
device = model_runner.device
+ ret.mm_token_modalities = _maybe_build_forward_token_modalities(
+ model_runner.model_config,
+ ret.mm_inputs,
+ extend_prefix_lens if isinstance(extend_prefix_lens, list) else None,
+ extend_seq_lens if isinstance(extend_seq_lens, list) else None,
+ len(batch.input_ids) if batch.input_ids is not None else 0,
+ device,
+ )
+
model_runner.kv_index_translator.rebind_write_loc(ret)
if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get():
diff --git a/python/sglang/srt/models/bailing_mm.py b/python/sglang/srt/models/bailing_mm.py
new file mode 100644
index 000000000..59533f104
--- /dev/null
+++ b/python/sglang/srt/models/bailing_mm.py
@@ -0,0 +1,270 @@
+# Copyright 2023 Antgroup and The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""Legacy Bailing multimodal wrappers for image and video inference."""
+
+import logging
+from typing import Iterable, List, Optional, Set, Tuple
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from transformers import PretrainedConfig
+
+from sglang.srt.distributed import get_pp_group
+from sglang.srt.layers.quantization.base_config import QuantizationConfig
+from sglang.srt.layers.utils import PPMissingLayer
+from sglang.srt.managers.mm_utils import (
+ MultiModalityDataPaddingPatternMultimodalTokens,
+ general_mm_embed_routine,
+)
+from sglang.srt.managers.schedule_batch import (
+ MultimodalDataItem,
+ MultimodalInputs,
+)
+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.bailing_moe import BailingMoeV2ForCausalLM
+from sglang.srt.models.qwen2_5_vl import Qwen2_5_VisionTransformer
+from sglang.srt.multimodal.mm_utils import materialize_multimodal_features
+from sglang.srt.runtime_context import get_mm
+from sglang.srt.utils import add_prefix
+
+logger = logging.getLogger(__name__)
+
+
+class BailingMMNativeForConditionalGeneration(nn.Module):
+ """Bailing MoE V2 wrapper with optional image/video encoding."""
+
+ def __init__(
+ self,
+ config: PretrainedConfig,
+ quant_config: Optional[QuantizationConfig] = None,
+ prefix: str = "",
+ ) -> None:
+ super().__init__()
+ self.pp_group = get_pp_group()
+ self.config = config
+ self.quant_config = quant_config
+ self.use_data_parallel = get_mm().mm_enable_dp_encoder
+ text_config = config.llm_config
+ self.model = BailingMoeV2ForCausalLM(
+ text_config,
+ quant_config,
+ prefix=add_prefix("model", prefix),
+ )
+
+ if getattr(config, "audio_config", None) is not None:
+ raise ValueError(
+ "Audio is not supported by the Bailing SGLang port; "
+ "use an image/video-only checkpoint"
+ )
+
+ self._build_mm_encoders = self.pp_group.is_first_rank
+ self.vision = None
+ self.linear_proj = None
+ if config.vision_config is not None:
+ if self._build_mm_encoders:
+ vision_config = config.vision_config
+ architectures = getattr(vision_config, "architectures", None) or []
+ arch = architectures[0] if architectures else vision_config.model_type
+ if arch in {
+ "Qwen3MoeVisionTransformer",
+ "Qwen3_VisionTransformer",
+ "qwen3_vl_moe",
+ }:
+ from sglang.srt.models.qwen3_vl import Qwen3VLMoeVisionModel
+
+ vision_cls = Qwen3VLMoeVisionModel
+ elif arch == "Qwen2_5_VisionTransformer":
+ vision_cls = Qwen2_5_VisionTransformer
+ else:
+ raise ValueError(f"Unsupported Bailing vision architecture: {arch}")
+ self.vision = vision_cls(
+ vision_config,
+ quant_config=quant_config,
+ prefix=add_prefix("vision", prefix),
+ use_data_parallel=self.use_data_parallel,
+ )
+ projection_layers = [
+ nn.Linear(
+ vision_config.out_hidden_size,
+ self.model.config.hidden_size,
+ )
+ ]
+ for _ in range(1, config.mlp_depth):
+ projection_layers.extend(
+ [
+ nn.GELU(),
+ nn.Linear(
+ self.model.config.hidden_size,
+ self.model.config.hidden_size,
+ ),
+ ]
+ )
+ self.linear_proj = nn.Sequential(*projection_layers)
+ else:
+ self.vision = PPMissingLayer()
+ self.linear_proj = PPMissingLayer()
+
+ self.is_mrope_enabled = "mrope_section" in config.llm_config.rope_parameters
+ self.pattern = MultiModalityDataPaddingPatternMultimodalTokens()
+
+ def get_input_embeddings(self):
+ return self.model.get_input_embeddings()
+
+ def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
+ return self.pattern.pad_input_tokens(input_ids, mm_inputs)
+
+ def _get_vision_feature(
+ self, items: List[MultimodalDataItem], grid_thw: torch.Tensor
+ ) -> torch.Tensor:
+ if self.vision is None or not self._build_mm_encoders:
+ raise RuntimeError("Vision encoder is unavailable on this PP stage")
+ pixel_values = materialize_multimodal_features(
+ [item.feature for item in items],
+ device=self.vision.device,
+ dtype=self.vision.dtype,
+ )
+ if self.use_data_parallel:
+ from sglang.srt.multimodal.mm_utils import (
+ run_dp_sharded_mrope_vision_model,
+ )
+
+ vision_embeds = run_dp_sharded_mrope_vision_model(
+ self.vision,
+ pixel_values,
+ grid_thw.tolist(),
+ rope_type="rope_3d",
+ )
+ else:
+ vision_embeds = self.vision(pixel_values, grid_thw=grid_thw)
+ deepstack_indexes = getattr(
+ self.config.vision_config, "deepstack_visual_indexes", []
+ )
+ if deepstack_indexes:
+ expected_dim = (len(deepstack_indexes) + 1) * (
+ self.config.vision_config.out_hidden_size
+ )
+ if vision_embeds.shape[-1] != expected_dim:
+ raise ValueError(
+ "Unexpected Bailing vision embedding width: "
+ f"expected={expected_dim}, got={vision_embeds.shape[-1]}"
+ )
+ vision_embeds = vision_embeds[
+ ..., : self.config.vision_config.out_hidden_size
+ ]
+ return F.normalize(self.linear_proj(vision_embeds).float(), dim=-1)
+
+ def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
+ image_grid_thw = torch.concat([item.image_grid_thw for item in items], dim=0)
+ return self._get_vision_feature(items, image_grid_thw)
+
+ def get_video_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
+ video_grid_thw = torch.concat([item.video_grid_thw for item in items], dim=0)
+ return self._get_vision_feature(items, video_grid_thw)
+
+ def get_audio_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
+ raise ValueError("Audio inputs are not supported by the Bailing SGLang port")
+
+ @torch.no_grad()
+ def forward(
+ self,
+ input_ids: torch.Tensor,
+ positions: torch.Tensor,
+ forward_batch: ForwardBatch,
+ input_embeds: Optional[torch.Tensor] = None,
+ get_embedding: bool = False,
+ pp_proxy_tensors: Optional[PPProxyTensors] = None,
+ ) -> torch.Tensor:
+ del input_embeds, get_embedding
+ if self.is_mrope_enabled:
+ positions = forward_batch.mrope_positions
+ return general_mm_embed_routine(
+ input_ids=input_ids,
+ forward_batch=forward_batch,
+ language_model=self.model,
+ multimodal_model=self,
+ positions=positions,
+ pp_proxy_tensors=pp_proxy_tensors,
+ )
+
+ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> Set[str]:
+ params_dict = dict(self.named_parameters(remove_duplicate=False))
+ buffers_dict = dict(self.named_buffers())
+ loaded_non_text: Set[str] = set()
+ unexpected_non_text = []
+
+ def dispatch_text_weights():
+ for name, loaded_weight in weights:
+ is_multimodal = name.startswith(
+ ("model.visual.", "model.vision.", "model.linear_proj.")
+ )
+ if name.startswith("model.") and not is_multimodal:
+ yield name[len("model.") :], loaded_weight
+ continue
+ if name.startswith(("model.visual.", "model.vision.")):
+ _, _, suffix = name.partition(".")
+ _, _, suffix = suffix.partition(".")
+ name = f"vision.{suffix}"
+ elif name.startswith("model.linear_proj."):
+ name = name[len("model.") :]
+ mapped_name = name.replace("attn.qkv.", "attn.qkv_proj.")
+ target = params_dict.get(mapped_name)
+ if target is not None:
+ weight_loader = getattr(
+ target, "weight_loader", default_weight_loader
+ )
+ weight_loader(target, loaded_weight)
+ loaded_non_text.add(mapped_name)
+ elif mapped_name in buffers_dict:
+ buffers_dict[mapped_name].copy_(loaded_weight)
+ loaded_non_text.add(mapped_name)
+ else:
+ unexpected_non_text.append(name)
+
+ loaded_text = self.model.load_weights(dispatch_text_weights())
+ required_non_text = {
+ name
+ for name in params_dict
+ if self._build_mm_encoders
+ and (name.startswith("vision.") or name.startswith("linear_proj."))
+ }
+ missing_non_text = required_non_text - loaded_non_text
+ if missing_non_text:
+ raise RuntimeError(
+ "Missing required Bailing multimodal weights: "
+ f"{sorted(missing_non_text)[:20]}"
+ )
+ if unexpected_non_text:
+ logger.warning(
+ "Skipped %d Bailing checkpoint tensors; examples: %s",
+ len(unexpected_non_text),
+ unexpected_non_text[:10],
+ )
+ logger.info(
+ "Loaded Bailing weights: %d text tensors and %d multimodal tensors",
+ len(loaded_text),
+ len(loaded_non_text),
+ )
+ return set(loaded_text) | loaded_non_text
+
+
+class BailingMM2NativeForConditionalGeneration(BailingMMNativeForConditionalGeneration):
+ pass
+
+
+EntryClass = [
+ BailingMMNativeForConditionalGeneration,
+ BailingMM2NativeForConditionalGeneration,
+]
diff --git a/python/sglang/srt/models/bailing_mm_v3.py b/python/sglang/srt/models/bailing_mm_v3.py
new file mode 100644
index 000000000..dd604514f
--- /dev/null
+++ b/python/sglang/srt/models/bailing_mm_v3.py
@@ -0,0 +1,325 @@
+# Copyright 2023 Antgroup and The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""SGLang implementation of Bailing/Ling 3 VL image and video inference."""
+
+import logging
+from typing import Iterable, List, Optional, Set, Tuple
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from transformers import PretrainedConfig
+
+from sglang.srt.configs.bailing_hybrid import is_bailing_multi_gate_enabled
+from sglang.srt.distributed import get_pp_group
+from sglang.srt.layers.quantization.base_config import QuantizationConfig
+from sglang.srt.layers.utils import PPMissingLayer
+from sglang.srt.managers.mm_utils import (
+ MultiModalityDataPaddingPatternMultimodalTokens,
+ general_mm_embed_routine,
+)
+from sglang.srt.managers.schedule_batch import (
+ MultimodalDataItem,
+ MultimodalInputs,
+)
+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.bailing_moe_v3 import (
+ BailingMoeV3ForCausalLM,
+)
+from sglang.srt.models.qwen3_vl import Qwen3VLMoeVisionModel
+from sglang.srt.multimodal.mm_utils import materialize_multimodal_features
+from sglang.srt.runtime_context import get_mm
+from sglang.srt.utils import add_prefix
+
+logger = logging.getLogger(__name__)
+
+_PUBLIC_VISION_PARAMETER_PREFIXES = (
+ "visual.patch_embed.",
+ "visual.pos_embed.",
+ "visual.blocks.",
+ "visual.merger.",
+)
+_NON_TEXT_PREFIX_MAPPING = (
+ ("model.visual.", "visual."),
+ ("linear_proj.", "linear_proj."),
+ # Compatibility with the private training-checkpoint wrapper.
+ ("model.linear_proj.", "linear_proj."),
+)
+
+
+class BailingMoeV3VLForConditionalGeneration(nn.Module):
+ """Bailing MoE V3 language model with Qwen3 vision encoding."""
+
+ @staticmethod
+ def shared_experts_fusion_disable_reason(hf_config, quant_config):
+ return BailingMoeV3ForCausalLM.shared_experts_fusion_disable_reason(
+ hf_config.text_config, quant_config
+ )
+
+ def __init__(
+ self,
+ config: PretrainedConfig,
+ quant_config: Optional[QuantizationConfig] = None,
+ prefix: str = "",
+ ) -> None:
+ super().__init__()
+ self.pp_group = get_pp_group()
+ self.config = config
+ self.quant_config = quant_config
+ self.norm_query_embeds = getattr(config, "norm_query_embeds", False)
+ self.use_data_parallel = get_mm().mm_enable_dp_encoder
+
+ text_config = config.text_config
+ self.multi_gate_enabled = is_bailing_multi_gate_enabled(text_config)
+ self.model = BailingMoeV3ForCausalLM(
+ config=text_config,
+ quant_config=quant_config,
+ prefix=add_prefix("model", prefix),
+ )
+
+ if config.vision_config is None:
+ raise ValueError("BailingMoeV3VL requires vision_config")
+ self._build_mm_encoders = self.pp_group.is_first_rank
+ if self._build_mm_encoders:
+ self.visual = Qwen3VLMoeVisionModel(
+ config.vision_config,
+ quant_config=quant_config,
+ prefix=add_prefix("visual", prefix),
+ use_data_parallel=self.use_data_parallel,
+ )
+ else:
+ self.visual = PPMissingLayer()
+
+ self.disable_merger_proj = getattr(
+ config.vision_config, "disable_merger_proj", False
+ )
+ self.deepstack_visual_indexes = tuple(
+ getattr(config.vision_config, "deepstack_visual_indexes", None) or ()
+ )
+ self.vision_out_dim = (
+ config.vision_config.hidden_size
+ * config.vision_config.spatial_merge_size**2
+ if self.disable_merger_proj
+ else config.vision_config.out_hidden_size
+ )
+ if self._build_mm_encoders:
+ self.linear_proj = nn.Sequential(
+ nn.Linear(self.vision_out_dim, text_config.hidden_size),
+ nn.GELU(),
+ nn.Linear(text_config.hidden_size, text_config.hidden_size),
+ )
+ else:
+ self.linear_proj = PPMissingLayer()
+
+ self.is_mrope_enabled = "mrope_section" in text_config.rope_parameters
+ self.pattern = MultiModalityDataPaddingPatternMultimodalTokens()
+
+ def get_input_embeddings(self):
+ return self.model.get_input_embeddings()
+
+ def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
+ return self.pattern.pad_input_tokens(input_ids, mm_inputs)
+
+ def _materialize_items(self, items: List[MultimodalDataItem]) -> torch.Tensor:
+ return materialize_multimodal_features(
+ [item.feature for item in items],
+ device=self.visual.device,
+ dtype=self.visual.dtype,
+ )
+
+ def _get_vision_feature(
+ self, items: List[MultimodalDataItem], grid_thw: torch.Tensor
+ ) -> torch.Tensor:
+ pixel_values = self._materialize_items(items)
+ if self.use_data_parallel:
+ from sglang.srt.multimodal.mm_utils import (
+ run_dp_sharded_mrope_vision_model,
+ )
+
+ vision_embeds = run_dp_sharded_mrope_vision_model(
+ self.visual,
+ pixel_values,
+ grid_thw.tolist(),
+ rope_type="rope_3d",
+ )
+ else:
+ vision_embeds = self.visual(pixel_values, grid_thw=grid_thw)
+
+ if self.deepstack_visual_indexes:
+ expected_dim = (
+ len(self.deepstack_visual_indexes) + 1
+ ) * self.vision_out_dim
+ if vision_embeds.shape[-1] != expected_dim:
+ raise ValueError(
+ "Unexpected Bailing vision embedding width: "
+ f"expected={expected_dim}, got={vision_embeds.shape[-1]}"
+ )
+ vision_embeds = vision_embeds[..., : self.vision_out_dim]
+
+ vision_embeds = self.linear_proj(vision_embeds)
+ if self.norm_query_embeds:
+ vision_embeds = F.normalize(vision_embeds, dim=-1)
+ return vision_embeds
+
+ def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
+ if not self._build_mm_encoders:
+ raise RuntimeError("Vision encoder is only available on the first PP stage")
+ image_grid_thw = torch.concat([item.image_grid_thw for item in items], dim=0)
+ return self._get_vision_feature(items, image_grid_thw)
+
+ def get_video_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
+ if not self._build_mm_encoders:
+ raise RuntimeError("Vision encoder is only available on the first PP stage")
+ video_grid_thw = torch.concat([item.video_grid_thw for item in items], dim=0)
+ return self._get_vision_feature(items, video_grid_thw)
+
+ @torch.no_grad()
+ def forward(
+ self,
+ input_ids: torch.Tensor,
+ positions: torch.Tensor,
+ forward_batch: ForwardBatch,
+ input_embeds: Optional[torch.Tensor] = None,
+ get_embedding: bool = False,
+ pp_proxy_tensors: Optional[PPProxyTensors] = None,
+ ) -> torch.Tensor:
+ del input_embeds, get_embedding
+ if self.is_mrope_enabled:
+ positions = forward_batch.mrope_positions
+ return general_mm_embed_routine(
+ input_ids=input_ids,
+ forward_batch=forward_batch,
+ language_model=self.model,
+ multimodal_model=self,
+ positions=positions,
+ pp_proxy_tensors=pp_proxy_tensors,
+ )
+
+ @classmethod
+ def get_model_config_for_expert_location(cls, config):
+ return BailingMoeV3ForCausalLM.get_model_config_for_expert_location(
+ config.text_config
+ )
+
+ def _load_non_text_weight(
+ self,
+ name: str,
+ loaded_weight: torch.Tensor,
+ params_dict: dict,
+ ) -> Optional[str]:
+ # Public layout: model.visual.* for the encoder and top-level
+ # linear_proj.* for the bridge. Keep the private nested bridge alias.
+ for checkpoint_prefix, parameter_prefix in _NON_TEXT_PREFIX_MAPPING:
+ if name.startswith(checkpoint_prefix):
+ name = parameter_prefix + name[len(checkpoint_prefix) :]
+ break
+ else:
+ return None
+ if name.startswith("visual."):
+ name = name.replace("attn.qkv.", "attn.qkv_proj.")
+ if name not in params_dict:
+ return None
+ param = params_dict[name]
+ weight_loader = getattr(param, "weight_loader", default_weight_loader)
+ weight_loader(param, loaded_weight)
+ return name
+
+ def _required_non_text_weights(self, params_dict: dict) -> Set[str]:
+ required_prefixes = list(_PUBLIC_VISION_PARAMETER_PREFIXES)
+ if self.deepstack_visual_indexes:
+ required_prefixes.append("visual.deepstack_merger_list.")
+ required_prefixes.append("linear_proj.")
+ return {
+ name
+ for name in params_dict
+ if self._build_mm_encoders and name.startswith(tuple(required_prefixes))
+ }
+
+ def _required_router_weights(self) -> Set[str]:
+ gate_names = (
+ ("gate", "image_gate", "audio_gate")
+ if self.multi_gate_enabled
+ else ("gate",)
+ )
+ suffixes = tuple(
+ f".mlp.{gate_name}.{parameter_name}"
+ for gate_name in gate_names
+ for parameter_name in ("weight", "expert_bias")
+ )
+ return {
+ name for name, _ in self.model.named_parameters() if name.endswith(suffixes)
+ }
+
+ @staticmethod
+ def _map_text_weight_name(name: str) -> Optional[str]:
+ if name == "lm_head.weight":
+ return name
+ if not name.startswith("model.") or name.startswith(
+ ("model.visual.", "model.linear_proj.")
+ ):
+ return None
+ # The public checkpoint is already model.layers/model.norm/
+ # model.word_embeddings; retain the private model.model alias.
+ return name.replace("model.model.", "model.", 1)
+
+ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> Set[str]:
+ params_dict = dict(self.named_parameters(remove_duplicate=False))
+ loaded_non_text: Set[str] = set()
+ unexpected_non_text = []
+
+ def dispatch_text_weights():
+ for name, loaded_weight in weights:
+ text_name = self._map_text_weight_name(name)
+ if text_name is not None:
+ yield text_name, loaded_weight
+ continue
+ loaded_name = self._load_non_text_weight(
+ name, loaded_weight, params_dict
+ )
+ if loaded_name is None:
+ unexpected_non_text.append(name)
+ else:
+ loaded_non_text.add(loaded_name)
+
+ loaded_text = self.model.load_weights(dispatch_text_weights())
+ required_non_text = self._required_non_text_weights(params_dict)
+ missing_non_text = required_non_text - loaded_non_text
+ if missing_non_text:
+ raise RuntimeError(
+ f"Missing required Bailing VL weights: {sorted(missing_non_text)[:20]}"
+ )
+
+ required_router_weights = self._required_router_weights()
+ missing_router_weights = required_router_weights - loaded_text
+ if missing_router_weights:
+ raise RuntimeError(
+ "Missing required Bailing VL router weights: "
+ f"{sorted(missing_router_weights)[:20]}"
+ )
+ if unexpected_non_text:
+ logger.warning(
+ "Skipped %d non-text Bailing checkpoint tensors; examples: %s",
+ len(unexpected_non_text),
+ unexpected_non_text[:10],
+ )
+ logger.info(
+ "Loaded Bailing VL weights: %d text tensors and %d vision/projection tensors",
+ len(loaded_text),
+ len(loaded_non_text),
+ )
+ return set(loaded_text) | loaded_non_text
+
+
+EntryClass = [BailingMoeV3VLForConditionalGeneration]
diff --git a/python/sglang/srt/models/bailing_moe.py b/python/sglang/srt/models/bailing_moe.py
index 96a564127..7868411d1 100644
--- a/python/sglang/srt/models/bailing_moe.py
+++ b/python/sglang/srt/models/bailing_moe.py
@@ -61,6 +61,10 @@ from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.layers.moe.token_dispatcher import DeepEPDispatcher
from sglang.srt.layers.moe.topk import TopK
from sglang.srt.layers.moe.utils import filter_moe_weight_param_global_expert
+from sglang.srt.layers.multi_gate import (
+ create_multi_gate_mm_indices,
+ multi_gate_triton_kernel,
+)
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
@@ -187,6 +191,9 @@ class BailingMoESparseMoeBlock(nn.Module):
self.num_shared_experts = config.num_shared_experts
self.routed_scaling_factor = getattr(config, "routed_scaling_factor", 1.0)
self.score_function = getattr(config, "score_function", None)
+ self.multi_gate = getattr(config, "multi_gate", False) or (
+ getattr(config, "router_type", "topN") == "MultiRouter"
+ )
if config.hidden_act != "silu":
raise ValueError(
@@ -229,6 +236,43 @@ class BailingMoESparseMoeBlock(nn.Module):
self.gate.expert_bias.data if self.gate.expert_bias is not None else None
)
+ self.image_gate = None
+ self.audio_gate = None
+ self.image_correction_bias = None
+ self.audio_correction_bias = None
+ if self.multi_gate:
+ self.image_gate = BailingMoEGate(
+ config=config,
+ params_dtype=self.router_dtype,
+ prefix=add_prefix("image_gate", prefix),
+ )
+ self.audio_gate = BailingMoEGate(
+ config=config,
+ params_dtype=self.router_dtype,
+ prefix=add_prefix("audio_gate", prefix),
+ )
+ self.image_correction_bias = (
+ self.image_gate.expert_bias.data
+ if self.image_gate.expert_bias is not None
+ else None
+ )
+ self.audio_correction_bias = (
+ self.audio_gate.expert_bias.data
+ if self.audio_gate.expert_bias is not None
+ else None
+ )
+ if any(
+ bias is None
+ for bias in (
+ self.correction_bias,
+ self.image_correction_bias,
+ self.audio_correction_bias,
+ )
+ ):
+ raise ValueError(
+ "Bailing MultiRouter requires expert bias for text, image, and audio gates"
+ )
+
if self.score_function is not None:
assert (
self.score_function == "softmax" and self.correction_bias is None
@@ -246,6 +290,8 @@ class BailingMoESparseMoeBlock(nn.Module):
# num_fused_shared_experts=self.num_fused_shared_experts,
topk_group=self.topk_group,
correction_bias=self.correction_bias,
+ scoring_func=self.score_function
+ or ("sigmoid" if self.correction_bias is not None else "softmax"),
routed_scaling_factor=self.routed_scaling_factor,
)
@@ -303,7 +349,7 @@ class BailingMoESparseMoeBlock(nn.Module):
forward_batch: Optional[ForwardBatch] = None,
) -> torch.Tensor:
if not get_moe_a2a_backend().is_deepep():
- return self.forward_normal(hidden_states)
+ return self.forward_normal(hidden_states, forward_batch)
else:
return self.forward_deepep(hidden_states, forward_batch)
@@ -323,22 +369,64 @@ class BailingMoESparseMoeBlock(nn.Module):
shared_output = self.shared_experts(hidden_states)
return shared_output
- def _forward_router_experts(self, hidden_states: torch.Tensor):
- # router_logits: (num_tokens, n_experts)
- router_logits = self.gate(hidden_states)
- topk_output = self.topk(hidden_states, router_logits)
+ def _forward_gate(
+ self,
+ hidden_states: torch.Tensor,
+ forward_batch: Optional[ForwardBatch],
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
+ if (
+ self.multi_gate
+ and forward_batch is not None
+ and forward_batch.mm_token_modalities is not None
+ ):
+ if forward_batch.mm_token_modalities.shape[0] != hidden_states.shape[0]:
+ raise ValueError(
+ "Bailing modality metadata must align with MoE tokens: "
+ f"modalities={forward_batch.mm_token_modalities.shape[0]}, "
+ f"hidden_states={hidden_states.shape[0]}"
+ )
+ if forward_batch.multi_gate_indices is None:
+ forward_batch.multi_gate_indices = create_multi_gate_mm_indices(
+ forward_batch.mm_token_modalities
+ )
+ return multi_gate_triton_kernel(
+ hidden_states,
+ forward_batch.multi_gate_indices,
+ self.gate.weight,
+ self.image_gate.weight,
+ self.audio_gate.weight,
+ self.correction_bias,
+ self.image_correction_bias,
+ self.audio_correction_bias,
+ )
+ return self.gate(hidden_states), None
+
+ def _forward_router_experts(
+ self,
+ hidden_states: torch.Tensor,
+ forward_batch: Optional[ForwardBatch],
+ ):
+ router_logits, dynamic_expert_bias = self._forward_gate(
+ hidden_states, forward_batch
+ )
+ topk_output = self.topk(
+ hidden_states,
+ router_logits,
+ dynamic_expert_bias=dynamic_expert_bias,
+ )
return self.experts(hidden_states, topk_output)
def forward_normal_dual_stream(
self,
hidden_states: torch.Tensor,
+ forward_batch: Optional[ForwardBatch],
) -> torch.Tensor:
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
shared_output = self._forward_shared_experts(hidden_states.clone())
with torch.cuda.stream(self.alt_stream):
- router_output = self._forward_router_experts(hidden_states)
+ router_output = self._forward_router_experts(hidden_states, forward_batch)
current_stream.wait_stream(self.alt_stream)
return router_output, shared_output
@@ -346,6 +434,7 @@ class BailingMoESparseMoeBlock(nn.Module):
def forward_normal(
self,
hidden_states: torch.Tensor,
+ forward_batch: Optional[ForwardBatch] = None,
) -> torch.Tensor:
num_tokens, hidden_size = hidden_states.shape
hidden_states = hidden_states.view(-1, hidden_size)
@@ -356,11 +445,13 @@ class BailingMoESparseMoeBlock(nn.Module):
and get_is_capture_mode()
):
final_hidden_states, shared_output = self.forward_normal_dual_stream(
- hidden_states
+ hidden_states, forward_batch
)
else:
shared_output = self._forward_shared_experts(hidden_states)
- final_hidden_states = self._forward_router_experts(hidden_states)
+ final_hidden_states = self._forward_router_experts(
+ hidden_states, forward_batch
+ )
if self.num_shared_experts > 0:
final_hidden_states = final_hidden_states + shared_output
@@ -377,13 +468,16 @@ class BailingMoESparseMoeBlock(nn.Module):
shared_output = None
forward_mode = forward_batch.forward_mode
if is_non_idle_and_non_empty(forward_mode, hidden_states):
- router_logits = self.gate(hidden_states)
+ router_logits, dynamic_expert_bias = self._forward_gate(
+ hidden_states, forward_batch
+ )
if self.num_shared_experts > 0:
shared_output = self.shared_experts(hidden_states)
topk_output = self.topk(
hidden_states,
router_logits,
+ dynamic_expert_bias=dynamic_expert_bias,
num_token_non_padded=forward_batch.num_token_non_padded,
expert_location_dispatch_info=ExpertLocationDispatchInfo.init_new(
layer_id=self.layer_id,
@@ -482,6 +576,7 @@ class BailingMoEAttention(nn.Module):
base=config.rope_parameters["rope_theta"],
rope_scaling=config.rope_parameters,
)
+ self.is_mrope_enabled = "mrope_section" in config.rope_parameters
self.attn = RadixAttention(
self.num_heads,
@@ -516,6 +611,7 @@ class BailingMoEAttention(nn.Module):
can_fuse_set_kv = (
self.head_dim == self.rotary_emb.rotary_dim
and enable_fused_set_kv_buffer(forward_batch)
+ and not self.is_mrope_enabled
)
q, k = self.rotary_emb(
positions,
@@ -698,7 +794,10 @@ class BailingMoEModel(nn.Module):
self.config = config
self.vocab_size = config.vocab_size
self.embed_dim = config.hidden_size
- if self.pp_group.is_first_rank:
+ keep_word_embeddings = self.pp_group.is_first_rank or (
+ config.tie_word_embeddings and self.pp_group.is_last_rank
+ )
+ if keep_word_embeddings:
self.word_embeddings = VocabParallelEmbedding(
self.vocab_size,
self.embed_dim,
@@ -824,6 +923,9 @@ class BailingMoEForCausalLM(nn.Module):
self.capture_aux_hidden_states = False
+ def get_input_embeddings(self):
+ return self.model.word_embeddings
+
@property
def start_layer(self):
return self.model.start_layer
@@ -892,6 +994,10 @@ class BailingMoEForCausalLM(nn.Module):
("gate_up_proj", "gate_proj", 0),
("gate_up_proj", "up_proj", 1),
]
+ replaced_params_mapping = {
+ "key_layernorm": "k_norm",
+ "query_layernorm": "q_norm",
+ }
if is_nextn:
nextn_layer_prefix = f"model.layers.{nextn_layer_id}"
@@ -928,6 +1034,11 @@ class BailingMoEForCausalLM(nn.Module):
loaded_weight = F.normalize(loaded_weight, dim=0, p=2, eps=1e-7)
+ for param_name, weight_name in replaced_params_mapping.items():
+ if weight_name in name:
+ name = name.replace(weight_name, param_name)
+ break
+
if is_nextn:
if not name.startswith(nextn_layer_prefix):
continue
diff --git a/python/sglang/srt/models/bailing_moe_v3.py b/python/sglang/srt/models/bailing_moe_v3.py
index bac6a39d4..a2dfac58d 100644
--- a/python/sglang/srt/models/bailing_moe_v3.py
+++ b/python/sglang/srt/models/bailing_moe_v3.py
@@ -18,6 +18,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import (
is_fp8_fnuz,
)
from sglang.srt.configs import KimiLinearConfig
+from sglang.srt.configs.bailing_hybrid import is_bailing_multi_gate_enabled
from sglang.srt.distributed import (
get_pp_group,
moe_expert_parallel_all_reduce,
@@ -50,6 +51,10 @@ from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.layers.moe.topk import TopK
from sglang.srt.layers.moe.utils import is_shared_experts_fusion_disabled
+from sglang.srt.layers.multi_gate import (
+ create_multi_gate_mm_indices,
+ multi_gate_triton_kernel,
+)
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.quantization.fp8_utils import (
block_quant_dequant,
@@ -98,17 +103,23 @@ from sglang.srt.utils import (
_is_fp8_fnuz = is_fp8_fnuz()
-if _is_cuda:
- from sglang.kernels.ops.quantization.awq_dequantize import awq_dequantize
-elif _is_cpu and _is_cpu_amx_available:
- pass
-elif _is_hip:
- from sglang.kernels.ops.quantization.awq_triton import (
- awq_dequantize_triton as awq_dequantize,
- )
-elif not (_is_cpu and _is_cpu_amx_available):
- from vllm._custom_ops import awq_dequantize
+def _get_awq_dequantize():
+ """Load the platform AWQ kernel only when an AWQ tensor needs it."""
+ if _is_cuda:
+ from sglang.kernels.ops.quantization.awq_dequantize import awq_dequantize
+
+ return awq_dequantize
+ if _is_hip:
+ from sglang.kernels.ops.quantization.awq_triton import awq_dequantize_triton
+
+ return awq_dequantize_triton
+ if not (_is_cpu and _is_cpu_amx_available):
+ from vllm._custom_ops import awq_dequantize
+
+ return awq_dequantize
+ return None
+
_is_flashinfer_available = is_flashinfer_available()
_is_sm100_supported = is_cuda() and get_platform().is_sm100
@@ -182,6 +193,17 @@ class DsV3MLA(DeepseekV2AttentionMLA):
else:
self.g_proj = None
+ if "mrope_section" in rope_scaling and rope_scaling.get("video_rope", False):
+ rope_scaling["rope_type"] = "default"
+ self.rotary_emb = get_rope(
+ qk_rope_head_dim,
+ rotary_dim=qk_rope_head_dim,
+ max_position=max_position_embeddings,
+ base=rope_theta,
+ rope_scaling=rope_scaling,
+ is_neox_style=not getattr(config, "rope_interleave", True),
+ )
+
def forward(
self,
positions: torch.Tensor,
@@ -410,6 +432,16 @@ class BailingMoEGate(nn.Module):
return logits
+def _get_bailing_num_shared_experts(config: PretrainedConfig) -> int:
+ shared_intermediate_size = getattr(
+ config, "moe_shared_expert_intermediate_size", None
+ )
+ num_shared_experts = getattr(config, "num_shared_experts", 0) or 0
+ if shared_intermediate_size is None:
+ return num_shared_experts
+ return max(1, num_shared_experts) if shared_intermediate_size > 0 else 0
+
+
class BailingMoE(nn.Module):
@staticmethod
def _get_swiglu_limit(limit_list, layer_num):
@@ -441,9 +473,10 @@ class BailingMoE(nn.Module):
self.norm_expert_prob = getattr(config, "norm_topk_prob", False)
self.hidden_size = config.hidden_size
self.intermediate_size = config.moe_intermediate_size
- self.num_shared_experts = getattr(config, "num_shared_experts", 0)
+ self.num_shared_experts = _get_bailing_num_shared_experts(config)
self.routed_scaling_factor = getattr(config, "routed_scaling_factor", 1.0)
self.score_function = getattr(config, "score_function", None)
+ self.multi_gate = is_bailing_multi_gate_enabled(config)
self.num_fused_shared_experts = num_fused_shared_experts
@@ -489,6 +522,43 @@ class BailingMoE(nn.Module):
self.gate.expert_bias.data if self.gate.expert_bias is not None else None
)
+ self.image_gate = None
+ self.audio_gate = None
+ self.image_correction_bias = None
+ self.audio_correction_bias = None
+ if self.multi_gate:
+ self.image_gate = BailingMoEGate(
+ config=config,
+ params_dtype=self.router_dtype,
+ prefix=add_prefix("image_gate", prefix),
+ )
+ self.audio_gate = BailingMoEGate(
+ config=config,
+ params_dtype=self.router_dtype,
+ prefix=add_prefix("audio_gate", prefix),
+ )
+ self.image_correction_bias = (
+ self.image_gate.expert_bias.data
+ if self.image_gate.expert_bias is not None
+ else None
+ )
+ self.audio_correction_bias = (
+ self.audio_gate.expert_bias.data
+ if self.audio_gate.expert_bias is not None
+ else None
+ )
+ if any(
+ bias is None
+ for bias in (
+ self.correction_bias,
+ self.image_correction_bias,
+ self.audio_correction_bias,
+ )
+ ):
+ raise ValueError(
+ "Bailing MultiRouter requires expert bias for text, image, and audio gates"
+ )
+
if self.score_function is not None:
assert (
self.score_function == "softmax" and self.correction_bias is None
@@ -530,6 +600,8 @@ class BailingMoE(nn.Module):
num_expert_group=self.num_expert_group,
topk_group=self.topk_group,
correction_bias=self.correction_bias,
+ scoring_func=self.score_function
+ or ("sigmoid" if self.correction_bias is not None else "softmax"),
routed_scaling_factor=self.routed_scaling_factor,
apply_routed_scaling_factor_on_output=(
self.experts.should_fuse_routed_scaling_factor_in_topk
@@ -650,11 +722,12 @@ class BailingMoE(nn.Module):
) -> torch.Tensor:
if self._enable_a2a_moe:
return self.forward_deepep(hidden_states, forward_batch)
- return self.forward_normal(hidden_states)
+ return self.forward_normal(hidden_states, forward_batch)
def forward_normal(
self,
hidden_states: torch.Tensor,
+ forward_batch: Optional[ForwardBatch] = None,
) -> torch.Tensor:
num_tokens, hidden_size = hidden_states.shape
hidden_states = hidden_states.view(-1, hidden_size)
@@ -670,11 +743,13 @@ class BailingMoE(nn.Module):
and get_is_capture_mode()
):
final_hidden_states, shared_output = self.forward_normal_dual_stream(
- hidden_states
+ hidden_states, forward_batch
)
else:
shared_output = self._forward_shared_experts(hidden_states)
- final_hidden_states = self._forward_router_experts(hidden_states)
+ final_hidden_states = self._forward_router_experts(
+ hidden_states, forward_batch
+ )
if shared_output is not None:
final_hidden_states = final_hidden_states + shared_output
@@ -699,14 +774,57 @@ class BailingMoE(nn.Module):
return None
return self.shared_experts(hidden_states)
- def _forward_router_experts(self, hidden_states: torch.Tensor) -> torch.Tensor:
- router_logits = self.gate(hidden_states)
- topk_output = self.topk(hidden_states, router_logits)
+ def _forward_gate(
+ self,
+ hidden_states: torch.Tensor,
+ forward_batch: Optional[ForwardBatch],
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
+ if (
+ self.multi_gate
+ and forward_batch is not None
+ and forward_batch.mm_token_modalities is not None
+ ):
+ if forward_batch.mm_token_modalities.shape[0] != hidden_states.shape[0]:
+ raise ValueError(
+ "Bailing modality metadata must align with MoE tokens: "
+ f"modalities={forward_batch.mm_token_modalities.shape[0]}, "
+ f"hidden_states={hidden_states.shape[0]}"
+ )
+ if forward_batch.multi_gate_indices is None:
+ forward_batch.multi_gate_indices = create_multi_gate_mm_indices(
+ forward_batch.mm_token_modalities
+ )
+ return multi_gate_triton_kernel(
+ hidden_states,
+ forward_batch.multi_gate_indices,
+ self.gate.weight,
+ self.image_gate.weight,
+ self.audio_gate.weight,
+ self.correction_bias,
+ self.image_correction_bias,
+ self.audio_correction_bias,
+ )
+ return self.gate(hidden_states), None
+
+ def _forward_router_experts(
+ self,
+ hidden_states: torch.Tensor,
+ forward_batch: Optional[ForwardBatch],
+ ) -> torch.Tensor:
+ router_logits, dynamic_expert_bias = self._forward_gate(
+ hidden_states, forward_batch
+ )
+ topk_output = self.topk(
+ hidden_states,
+ router_logits,
+ dynamic_expert_bias=dynamic_expert_bias,
+ )
return self.experts(hidden_states, topk_output)
def forward_normal_dual_stream(
self,
hidden_states: torch.Tensor,
+ forward_batch: Optional[ForwardBatch],
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
@@ -714,7 +832,9 @@ class BailingMoE(nn.Module):
shared_output = self._forward_shared_experts(hidden_states.clone())
with torch.cuda.stream(self.alt_stream):
- final_hidden_states = self._forward_router_experts(hidden_states)
+ final_hidden_states = self._forward_router_experts(
+ hidden_states, forward_batch
+ )
current_stream.wait_stream(self.alt_stream)
return final_hidden_states, shared_output
@@ -733,10 +853,13 @@ class BailingMoE(nn.Module):
topk_output = self.topk.empty_topk_output(hidden_states.device)
else:
shared_output = self._forward_shared_experts(hidden_states)
- router_logits = self.gate(hidden_states)
+ router_logits, dynamic_expert_bias = self._forward_gate(
+ hidden_states, forward_batch
+ )
topk_output = self.topk(
hidden_states,
router_logits,
+ dynamic_expert_bias=dynamic_expert_bias,
num_token_non_padded=forward_batch.num_token_non_padded,
expert_location_dispatch_info=ExpertLocationDispatchInfo.init_new(
layer_id=self.layer_id,
@@ -952,7 +1075,7 @@ class BailingMoELinearDecoderLayer(nn.Module):
kv_lora_rank=config.kv_lora_rank,
rope_theta=config.rope_parameters.get("rope_theta", 600000),
rope_scaling=config.rope_parameters,
- max_position_embeddings=262144,
+ max_position_embeddings=config.max_position_embeddings,
quant_config=quant_config,
layer_id=layer_id,
reduce_results=False,
@@ -1162,7 +1285,10 @@ class BailingMoELinearModel(nn.Module):
f"num_layers={self.num_layers} must be divided by layer_group_size={self.layer_group_size}"
)
- if self.pp_group.is_first_rank:
+ keep_word_embeddings = self.pp_group.is_first_rank or (
+ config.tie_word_embeddings and self.pp_group.is_last_rank
+ )
+ if keep_word_embeddings:
self.word_embeddings = VocabParallelEmbedding(
self.vocab_size,
self.embed_dim,
@@ -1226,14 +1352,14 @@ class BailingMoELinearModel(nn.Module):
input_ids: Optional[torch.Tensor],
positions: torch.Tensor,
forward_batch: Optional[ForwardBatch] = None,
- inputs_embeds: Optional[torch.Tensor] = None,
+ input_embeds: Optional[torch.Tensor] = None,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
) -> Union[torch.Tensor, PPProxyTensors]:
if self.pp_group.is_first_rank:
- if inputs_embeds is None:
+ if input_embeds is None:
hidden_states = self.word_embeddings(input_ids)
else:
- hidden_states = inputs_embeds
+ hidden_states = input_embeds
residual = None
else:
assert pp_proxy_tensors is not None
@@ -1241,7 +1367,7 @@ class BailingMoELinearModel(nn.Module):
residual = pp_proxy_tensors["residual"]
total_num_layers = self.end_layer - self.start_layer
- device = inputs_embeds.device if inputs_embeds is not None else input_ids.device
+ device = hidden_states.device
zero_allocator = BumpAllocator(
buffer_size=total_num_layers * 2 * (2 if forward_batch.can_run_tbo else 1),
dtype=torch.float32,
@@ -1361,7 +1487,7 @@ class BailingMoeV3ForCausalLM(nn.Module):
quant_config,
expected_architecture="BailingMoeV3ForCausalLM",
):
- num_shared_experts = getattr(hf_config, "num_shared_experts", 0)
+ num_shared_experts = _get_bailing_num_shared_experts(hf_config)
if num_shared_experts == 0:
return None
if not get_moe_a2a_backend().is_none():
@@ -1447,7 +1573,7 @@ class BailingMoeV3ForCausalLM(nn.Module):
self.num_fused_shared_experts = (
0
if is_shared_experts_fusion_disabled()
- else getattr(self.config, "num_shared_experts", 0)
+ else _get_bailing_num_shared_experts(self.config)
)
if self.num_fused_shared_experts == 0:
return
@@ -1501,6 +1627,11 @@ class BailingMoeV3ForCausalLM(nn.Module):
if not hasattr(self_attn, "kv_b_proj"):
continue
if hasattr(self_attn.kv_b_proj, "qweight"):
+ awq_dequantize = _get_awq_dequantize()
+ if awq_dequantize is None:
+ raise RuntimeError(
+ "AWQ dequantization is unavailable on this CPU platform"
+ )
if _is_cuda or _is_hip:
w = awq_dequantize(
self_attn.kv_b_proj.qweight,
@@ -1678,13 +1809,13 @@ class BailingMoeV3ForCausalLM(nn.Module):
input_ids: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
- inputs_embeds: Optional[torch.Tensor] = None,
+ input_embeds: Optional[torch.Tensor] = None,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
) -> Union[torch.Tensor, PPProxyTensors]:
hidden_states = self.model(
input_ids=input_ids,
positions=positions,
- inputs_embeds=inputs_embeds,
+ input_embeds=input_embeds,
forward_batch=forward_batch,
pp_proxy_tensors=pp_proxy_tensors,
)
diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py
index db6ba8427..53310d53b 100644
--- a/python/sglang/srt/models/qwen3_vl.py
+++ b/python/sglang/srt/models/qwen3_vl.py
@@ -295,43 +295,46 @@ class Qwen3VLMoeVisionPatchMerger(nn.Module):
use_data_parallel: bool = False,
tp_size: Optional[int] = None,
tp_rank: Optional[int] = None,
+ disable_merger_proj: bool = False,
) -> None:
super().__init__()
self.hidden_size = context_dim * (spatial_merge_size**2)
self.padded_context_dim = padded_context_dim * (spatial_merge_size**2)
self.use_postshuffle_norm = use_postshuffle_norm
+ self.disable_merger_proj = disable_merger_proj
if norm_layer is None:
norm_layer = partial(nn.LayerNorm, eps=1e-6)
self.norm = norm_layer(
self.hidden_size if use_postshuffle_norm else context_dim
)
- self.tp_size, self.tp_rank = _resolve_vision_tp(
- use_data_parallel=use_data_parallel,
- tp_size=tp_size,
- tp_rank=tp_rank,
- )
- self.linear_fc1 = ColumnParallelLinear(
- self.hidden_size,
- self.padded_context_dim,
- bias=True,
- quant_config=quant_config,
- prefix=add_prefix("linear_fc1", prefix),
- tp_size=self.tp_size,
- tp_rank=self.tp_rank,
- )
- self.act_fn = nn.GELU()
- self.linear_fc2 = RowParallelLinear(
- self.padded_context_dim,
- dim,
- bias=True,
- quant_config=quant_config,
- prefix=add_prefix("linear_fc2", prefix),
- tp_size=self.tp_size,
- tp_rank=self.tp_rank,
- use_dp_attention_reduce=is_dp_attention_enabled(),
- )
+ if not disable_merger_proj:
+ self.tp_size, self.tp_rank = _resolve_vision_tp(
+ use_data_parallel=use_data_parallel,
+ tp_size=tp_size,
+ tp_rank=tp_rank,
+ )
+ self.linear_fc1 = ColumnParallelLinear(
+ self.hidden_size,
+ self.padded_context_dim,
+ bias=True,
+ quant_config=quant_config,
+ prefix=add_prefix("linear_fc1", prefix),
+ tp_size=self.tp_size,
+ tp_rank=self.tp_rank,
+ )
+ self.act_fn = nn.GELU()
+ self.linear_fc2 = RowParallelLinear(
+ self.padded_context_dim,
+ dim,
+ bias=True,
+ quant_config=quant_config,
+ prefix=add_prefix("linear_fc2", prefix),
+ tp_size=self.tp_size,
+ tp_rank=self.tp_rank,
+ use_dp_attention_reduce=is_dp_attention_enabled(),
+ )
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.use_postshuffle_norm:
@@ -339,10 +342,11 @@ class Qwen3VLMoeVisionPatchMerger(nn.Module):
else:
x = self.norm(x).view(-1, self.hidden_size)
- x_parallel, _ = self.linear_fc1(x)
- x_parallel = self.act_fn(x_parallel)
- out, _ = self.linear_fc2(x_parallel)
- return out
+ if not self.disable_merger_proj:
+ x_parallel, _ = self.linear_fc1(x)
+ x_parallel = self.act_fn(x_parallel)
+ x, _ = self.linear_fc2(x_parallel)
+ return x
class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
@@ -369,9 +373,13 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
self.use_data_parallel = use_data_parallel
# layer indexes of which layer's output should be deep-stacked
self.deepstack_visual_indexes = vision_config.deepstack_visual_indexes
- self.out_hidden_size = vision_config.out_hidden_size * (
- 1 + len(self.deepstack_visual_indexes)
+ self.disable_merger_proj = getattr(vision_config, "disable_merger_proj", False)
+ merger_out_dim = (
+ self.hidden_size * self.spatial_merge_unit
+ if self.disable_merger_proj
+ else vision_config.out_hidden_size
)
+ self.out_hidden_size = merger_out_dim * (1 + len(self.deepstack_visual_indexes))
self.patch_embed = Qwen3VLVisionPatchEmbed(config=vision_config)
if self.pp_group.is_first_rank:
self.pos_embed = VocabParallelEmbedding(
@@ -441,6 +449,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
quant_config=quant_config,
prefix=add_prefix("merger", prefix),
use_data_parallel=use_data_parallel,
+ disable_merger_proj=self.disable_merger_proj,
)
self.deepstack_merger_list = nn.ModuleList(
@@ -455,6 +464,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
quant_config=quant_config,
prefix=add_prefix(f"deepstack_merger_list.{layer_idx}", prefix),
use_data_parallel=use_data_parallel,
+ disable_merger_proj=self.disable_merger_proj,
)
for layer_idx in range(len(self.deepstack_visual_indexes))
]
diff --git a/python/sglang/srt/multimodal/processors/bailing_mm.py b/python/sglang/srt/multimodal/processors/bailing_mm.py
new file mode 100644
index 000000000..011f65b56
--- /dev/null
+++ b/python/sglang/srt/multimodal/processors/bailing_mm.py
@@ -0,0 +1,261 @@
+# Copyright 2023 Antgroup and The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""Image and video processor for Bailing multimodal checkpoints."""
+
+from typing import Optional
+
+import torch
+from transformers import BaseImageProcessor
+
+from sglang.srt.layers.rotary_embedding.bailing_mrope import BailingMRotaryEmbedding
+from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput
+from sglang.srt.models.bailing_mm import (
+ BailingMM2NativeForConditionalGeneration,
+ BailingMMNativeForConditionalGeneration,
+)
+from sglang.srt.models.bailing_mm_v3 import BailingMoeV3VLForConditionalGeneration
+from sglang.srt.multimodal.processors.base_processor import (
+ BaseMultimodalProcessor,
+ BaseMultiModalProcessorOutput,
+ MultimodalSpecialTokens,
+)
+from sglang.srt.utils import ImageData, VideoData
+
+DEFAULT_IMAGE_PATCH_TOKEN = "<|image_pad|>"
+DEFAULT_FRAME_PATCH_TOKEN = "<|video_pad|>"
+DEFAULT_VISION_START_TOKEN = "<|vision_start|>"
+DEFAULT_VISION_END_TOKEN = "<|vision_end|>"
+DEFAULT_VIDEO_START_TOKEN = "<|video_start|>"
+DEFAULT_VIDEO_END_TOKEN = "<|video_end|>"
+
+
+class BailingMMMultimodalProcessor(BaseMultimodalProcessor):
+ """Prepare image/video features and Bailing three-axis positions."""
+
+ models = [
+ BailingMMNativeForConditionalGeneration,
+ BailingMM2NativeForConditionalGeneration,
+ BailingMoeV3VLForConditionalGeneration,
+ ]
+
+ def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
+ super().__init__(hf_config, server_args, _processor, *args, **kwargs)
+ tokenizer = getattr(_processor, "tokenizer", _processor)
+
+ self.image_token_id = self._resolve_token_id(
+ hf_config,
+ tokenizer,
+ "image_token_id",
+ "image_patch_token",
+ DEFAULT_IMAGE_PATCH_TOKEN,
+ )
+ self.video_token_id = self._resolve_token_id(
+ hf_config,
+ tokenizer,
+ "video_token_id",
+ "video_patch_token",
+ DEFAULT_FRAME_PATCH_TOKEN,
+ )
+ image_token = self._wrapped_token(
+ _processor,
+ tokenizer,
+ ("vision_start_token", "vision_bos_token"),
+ DEFAULT_VISION_START_TOKEN,
+ ("image_token",),
+ DEFAULT_IMAGE_PATCH_TOKEN,
+ ("vision_end_token", "vision_eos_token"),
+ DEFAULT_VISION_END_TOKEN,
+ )
+ video_token = self._wrapped_token(
+ _processor,
+ tokenizer,
+ ("video_start_token", "video_bos_token"),
+ DEFAULT_VIDEO_START_TOKEN,
+ ("video_token",),
+ DEFAULT_FRAME_PATCH_TOKEN,
+ ("video_end_token", "video_eos_token"),
+ DEFAULT_VIDEO_END_TOKEN,
+ )
+ self.mm_tokens = MultimodalSpecialTokens(
+ image_token=image_token,
+ video_token=video_token,
+ image_token_id=self.image_token_id,
+ video_token_id=self.video_token_id,
+ ).build(_processor)
+
+ @staticmethod
+ def _resolve_token_id(
+ hf_config,
+ tokenizer,
+ config_attr: str,
+ text_config_attr: str,
+ fallback_token: str,
+ ) -> int:
+ token_id = getattr(hf_config, config_attr, None)
+ if token_id is not None:
+ return token_id
+ for attr in ("text_config", "llm_config"):
+ text_config = getattr(hf_config, attr, None)
+ token_id = getattr(text_config, text_config_attr, None)
+ if token_id is not None:
+ return token_id
+ return tokenizer.convert_tokens_to_ids(fallback_token)
+
+ @staticmethod
+ def _token_string(processor, tokenizer, attrs, fallback: str) -> str:
+ for source in (processor, tokenizer):
+ for attr in attrs:
+ token = getattr(source, attr, None)
+ if token is not None:
+ return token
+ return fallback
+
+ @classmethod
+ def _wrapped_token(
+ cls,
+ processor,
+ tokenizer,
+ start_attrs,
+ start_fallback,
+ token_attrs,
+ token_fallback,
+ end_attrs,
+ end_fallback,
+ ) -> str:
+ return (
+ cls._token_string(processor, tokenizer, start_attrs, start_fallback)
+ + cls._token_string(processor, tokenizer, token_attrs, token_fallback)
+ + cls._token_string(processor, tokenizer, end_attrs, end_fallback)
+ )
+
+ def process_mm_data(
+ self,
+ input_text,
+ images=None,
+ videos=None,
+ audios=None,
+ processor=None,
+ **kwargs,
+ ) -> dict:
+ if audios:
+ raise ValueError("Audio inputs are not supported by Ling-3.0-flash-VL")
+ processor, _ = self._resolve_processor(processor)
+ processor_kwargs = {
+ "text": [input_text],
+ "return_tensors": "pt",
+ }
+ if images:
+ processor_kwargs["images"] = images
+ if videos:
+ processor_kwargs["videos"] = videos
+ image_processor = getattr(processor, "image_processor", None)
+ device: Optional[str] = None
+ if isinstance(image_processor, BaseImageProcessor):
+ device = self._fast_image_processor_device(processor)
+ if device is not None:
+ processor_kwargs["device"] = device
+
+ result = processor(**processor_kwargs)
+ for feature_name in self.FEATURE_NAMES:
+ feature = result.get(feature_name)
+ if not isinstance(feature, torch.Tensor):
+ continue
+ feature = feature.to(dtype=torch.bfloat16)
+ if not self.keep_mm_features_on_device:
+ feature = feature.cpu()
+ result[feature_name] = feature
+ return result
+
+ @staticmethod
+ def _request_url(item):
+ if isinstance(item, (ImageData, VideoData)):
+ return item.url
+ if isinstance(item, dict):
+ if "url" not in item:
+ raise ValueError("Bailing media dictionaries must contain a url")
+ return item["url"]
+ return item
+
+ def _processor_fetch_mm_input(self, prompt, image_data, video_data):
+ if isinstance(prompt, list):
+ if not prompt or not isinstance(prompt[0], int):
+ raise ValueError("Tokenized Bailing prompts must be a non-empty list")
+ prompt = self._tokenizer.decode(prompt)
+ if not isinstance(prompt, str):
+ raise TypeError(
+ f"Bailing prompt must be str or list[int], got {type(prompt)}"
+ )
+
+ contents = []
+ for item in image_data or []:
+ contents.append(
+ {
+ "type": "image_url",
+ "image_url": {"url": self._request_url(item)},
+ }
+ )
+ for item in video_data or []:
+ contents.append(
+ {
+ "type": "video_url",
+ "video_url": {"url": self._request_url(item)},
+ }
+ )
+ images, videos, audios = self._processor.process_vision_info(
+ conversations=[{"content": contents}]
+ )
+ if audios:
+ raise ValueError("Audio inputs are not supported by Ling-3.0-flash-VL")
+ return BaseMultiModalProcessorOutput(
+ images=images or [],
+ videos=videos or [],
+ audios=[],
+ input_text=prompt,
+ )
+
+ async def process_mm_data_async(
+ self,
+ image_data,
+ audio_data,
+ input_text,
+ request_obj,
+ **kwargs,
+ ):
+ if audio_data or getattr(request_obj, "audio_data", None):
+ raise ValueError("Audio inputs are not supported by Ling-3.0-flash-VL")
+ base_output = self._processor_fetch_mm_input(
+ input_text,
+ image_data,
+ getattr(request_obj, "video_data", None),
+ )
+ mm_items, input_ids, ret = await self.process_and_combine_mm_data_async(
+ base_output, self.mm_tokens
+ )
+ input_ids = input_ids.flatten()
+ mrope_positions, mrope_position_delta = (
+ BailingMRotaryEmbedding.bailing_3drope_get_input_positions_tensor(
+ input_ids,
+ self.hf_config,
+ image_grid_thw=ret.get("image_grid_thw"),
+ video_grid_thw=ret.get("video_grid_thw"),
+ )
+ )
+ return MultimodalProcessorOutput(
+ mm_items=mm_items,
+ input_ids=input_ids.tolist(),
+ im_token_id=self.image_token_id,
+ video_token_id=self.video_token_id,
+ mrope_positions=mrope_positions.squeeze(1),
+ mrope_position_delta=mrope_position_delta,
+ )
diff --git a/python/sglang/srt/parser/template_detection.py b/python/sglang/srt/parser/template_detection.py
index a662492f0..d2a6bc5b1 100644
--- a/python/sglang/srt/parser/template_detection.py
+++ b/python/sglang/srt/parser/template_detection.py
@@ -494,6 +494,16 @@ def _is_qwen3(ctx):
)
+def _is_ling3(ctx):
+ return (
+ ctx.has_text("SYSTEM")
+ and ctx.has_text("ASSISTANT")
+ and ctx.has_text("<|role_end|>")
+ and ctx.has_text("")
+ and ctx.has_text("")
+ )
+
+
def _is_deepseek_v3(ctx):
return ctx.reasoning_config == ReasoningToggleConfig(
toggle_param="thinking", default_enabled=False
@@ -535,6 +545,7 @@ REASONING_PARSER_RULES = (
DetectionRule(name="minimax", value="minimax", predicate=_is_minimax),
DetectionRule(name="step3p5", value="step3p5", predicate=_is_step3p5),
DetectionRule(name="step3", value="step3", predicate=_is_step3),
+ DetectionRule(name="ling3", value="ling3", predicate=_is_ling3),
DetectionRule(name="qwen3", value="qwen3", predicate=_is_qwen3),
DetectionRule(name="deepseek_v4", value="deepseek-v4", predicate=_is_deepseek_v4),
DetectionRule(name="deepseek_v3", value="deepseek-v3", predicate=_is_deepseek_v3),
@@ -573,6 +584,7 @@ TOOL_CALL_PARSER_RULES = (
DetectionRule(name="poolside_v1", value="poolside_v1", predicate=_is_poolside_v1),
DetectionRule(name="step3p5", value="step3p5", predicate=_is_step3p5),
DetectionRule(name="step3", value="step3", predicate=_is_step3),
+ DetectionRule(name="ling3", value="ling3", predicate=_is_ling3),
DetectionRule(
name="xml_kv_tool_call", value="glm45", predicate=_is_xml_kv_tool_call
),
@@ -768,6 +780,11 @@ def _architecture_auto_parsers(server_args, needs: Tuple[str, ...]) -> Dict[str,
if "KimiK3" in arch or model_type == "kimi_k3":
reasoning_parser, tool_call_parser = "kimi_k3", "kimi_k3"
+ elif arch in (
+ "BailingMoeV3ForCausalLM",
+ "BailingMoeV3VLForConditionalGeneration",
+ ) or model_type in ("bailing_hybrid", "bailing_moe_v3_vl"):
+ reasoning_parser, tool_call_parser = "ling3", "ling3"
elif "DeepseekV4" in arch:
reasoning_parser, tool_call_parser = "deepseek-v4", "deepseekv4"
elif "DeepseekV3" in arch:
diff --git a/python/sglang/srt/utils/hf_transformers/common.py b/python/sglang/srt/utils/hf_transformers/common.py
index 1d91b3a4a..c0ab4d66c 100644
--- a/python/sglang/srt/utils/hf_transformers/common.py
+++ b/python/sglang/srt/utils/hf_transformers/common.py
@@ -24,6 +24,8 @@ from huggingface_hub import snapshot_download
from sglang.srt.configs import (
AfmoeConfig,
BailingHybridConfig,
+ BailingMM2Config,
+ BailingMoeV3VLConfig,
ChatGLMConfig,
Cosmos3Config,
Cosmos3EdgeConfig,
@@ -111,6 +113,8 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
for cls in [
AfmoeConfig,
BailingHybridConfig,
+ BailingMM2Config,
+ BailingMoeV3VLConfig,
ChatGLMConfig,
DbrxConfig,
ExaoneConfig,
diff --git a/test/registered/unit/eplb/test_waterfill_eplb.py b/test/registered/unit/eplb/test_waterfill_eplb.py
index b060bdb20..99503d43e 100644
--- a/test/registered/unit/eplb/test_waterfill_eplb.py
+++ b/test/registered/unit/eplb/test_waterfill_eplb.py
@@ -98,8 +98,8 @@ class TestWaterfillEPLB(CustomTestCase):
self.assertTrue(torch.equal(processed_ids, torch.tensor([[0, 34, 270, 271]])))
self.assertTrue(torch.equal(recorder_ids, torch.tensor([[0, 33, 263]])))
- def test_topk_recorder_ids_match_dispatch_ids_without_per_rank_shared_slots(self):
- topk_ids = torch.tensor([[0, 33, 263, 256]], dtype=torch.int32)
+ def test_topk_recorder_ids_exclude_global_fused_shared_slot(self):
+ topk_ids = torch.tensor([[0, 33, 200, 256]], dtype=torch.int32)
topk_weights = torch.ones_like(topk_ids, dtype=torch.float32)
topk_config = TopKConfig(
top_k=4,
@@ -107,7 +107,7 @@ class TestWaterfillEPLB(CustomTestCase):
routed_scaling_factor=1.0,
)
dispatch_info = SimpleNamespace(
- num_physical_experts=264, ep_dispatch_algorithm="static"
+ num_physical_experts=256, ep_dispatch_algorithm="static"
)
def fake_eplb_postprocess(
@@ -136,8 +136,8 @@ class TestWaterfillEPLB(CustomTestCase):
expert_location_dispatch_info=dispatch_info,
)
- self.assertTrue(torch.equal(processed_ids, torch.tensor([[1, 34, 264, 257]])))
- self.assertTrue(torch.equal(recorder_ids, processed_ids))
+ self.assertTrue(torch.equal(processed_ids, torch.tensor([[1, 34, 201, 256]])))
+ self.assertTrue(torch.equal(recorder_ids, torch.tensor([[1, 34, 201]])))
if __name__ == "__main__":
diff --git a/test/registered/unit/layers/test_bailing_mrope_shape.py b/test/registered/unit/layers/test_bailing_mrope_shape.py
new file mode 100644
index 000000000..862dd5a7a
--- /dev/null
+++ b/test/registered/unit/layers/test_bailing_mrope_shape.py
@@ -0,0 +1,212 @@
+"""Regression tests for Bailing multimodal rotary positions and config bounds."""
+
+import unittest
+from types import SimpleNamespace
+from unittest import mock
+
+import torch
+
+from sglang.srt.configs.bailing_hybrid import BailingMoeV3VLConfig
+from sglang.srt.layers.rotary_embedding.bailing_mrope import BailingMRotaryEmbedding
+from sglang.test.ci.ci_register import register_cpu_ci
+from sglang.test.test_utils import CustomTestCase
+
+register_cpu_ci(est_time=5, suite="base-a-test-cpu")
+
+
+def _position_config(max_position_embeddings=131072):
+ return SimpleNamespace(
+ vision_config=SimpleNamespace(spatial_merge_size=2),
+ text_config=SimpleNamespace(
+ image_patch_token=11,
+ video_patch_token=12,
+ image_start_token=10,
+ video_start_token=13,
+ use_interleaved_frame_timestamp=False,
+ max_position_embeddings=max_position_embeddings,
+ ),
+ )
+
+
+class TestBailingMRotaryEmbedding(CustomTestCase):
+ def test_text_and_single_multimodal_position_shapes(self):
+ """A singleton sequence must retain the [3, batch, seq] contract."""
+ config = _position_config()
+ text_positions, text_delta = (
+ BailingMRotaryEmbedding.bailing_3drope_get_input_positions_tensor(
+ torch.tensor([7]), config, None, None
+ )
+ )
+ image_positions, image_delta = (
+ BailingMRotaryEmbedding.bailing_3drope_get_input_positions_tensor(
+ torch.tensor([10, 11]),
+ config,
+ image_grid_thw=torch.tensor([[1, 2, 2]]),
+ video_grid_thw=None,
+ )
+ )
+
+ self.assertEqual(text_positions.shape, (3, 1, 1))
+ self.assertEqual(text_delta.shape, (1, 1))
+ self.assertEqual(image_positions.shape, (3, 1, 2))
+ self.assertEqual(image_delta.shape, (1, 1))
+
+ def test_centered_height_positions_can_be_negative(self):
+ """Tall images require negative H coordinates instead of clamping to zero."""
+ config = _position_config()
+ input_ids = torch.tensor([10] + [11] * 7 + [99])
+
+ positions, _ = (
+ BailingMRotaryEmbedding.bailing_3drope_get_input_positions_tensor(
+ input_ids,
+ config,
+ image_grid_thw=torch.tensor([[1, 14, 2]]),
+ video_grid_thw=None,
+ )
+ )
+
+ self.assertEqual(positions.shape, (3, 1, 9))
+ self.assertLess(int(positions[1].min()), 0)
+
+ def test_checkpoint_position_bound_is_enforced(self):
+ """Media positions at or beyond the checkpoint context must fail clearly."""
+ config = _position_config(max_position_embeddings=4)
+ input_ids = torch.tensor([10] + [11] * 7)
+
+ with self.assertRaisesRegex(ValueError, "checkpoint bounds"):
+ BailingMRotaryEmbedding.bailing_3drope_get_input_positions_tensor(
+ input_ids,
+ config,
+ image_grid_thw=torch.tensor([[1, 14, 2]]),
+ video_grid_thw=None,
+ )
+
+ def test_negative_start_cache_growth_preserves_phase(self):
+ """Growing a negative-origin cache must append the next logical phase."""
+ with mock.patch("sglang.srt.layers.rotary_embedding.base._is_cpu", True):
+ rotary = BailingMRotaryEmbedding(
+ head_size=8,
+ rotary_dim=8,
+ max_position_embeddings=16,
+ base=10000,
+ is_neox_style=True,
+ dtype=torch.float32,
+ mrope_section=[2, 1, 1],
+ video_rope=True,
+ )
+ self.assertEqual(rotary.position_start, -16)
+ self.assertEqual(rotary.cos_sin_cache.shape[0], 32)
+
+ rotary._ensure_cos_sin_cache_length(32)
+ inv_freq = rotary._compute_inv_freq(rotary.base)
+ expected = torch.cat(((16 * inv_freq).cos(), (16 * inv_freq).sin()))
+ torch.testing.assert_close(rotary.cos_sin_cache[32], expected)
+
+ def test_yarn_scaling_extends_positive_cache_only(self):
+ """YaRN factor stretches the positive side; the negative side is fixed."""
+ with mock.patch("sglang.srt.layers.rotary_embedding.base._is_cpu", True):
+ rotary = BailingMRotaryEmbedding(
+ head_size=8,
+ rotary_dim=8,
+ max_position_embeddings=16,
+ base=10000,
+ is_neox_style=True,
+ dtype=torch.float32,
+ mrope_section=[2, 1, 1],
+ video_rope=True,
+ scaling_factor=2.0,
+ original_max_position_embeddings=16,
+ )
+ self.assertEqual(rotary.position_start, -16)
+ self.assertEqual(rotary.cos_sin_cache.shape[0], 16 + 32)
+ self.assertGreater(rotary.mscale, 1.0)
+
+ inv_freq = rotary._compute_inv_freq(rotary.base)
+ position = 20 # beyond the native bound of 16, row = 20 - (-16)
+ expected = torch.cat(
+ (
+ (position * inv_freq).cos() * rotary.mscale,
+ (position * inv_freq).sin() * rotary.mscale,
+ )
+ )
+ torch.testing.assert_close(rotary.cos_sin_cache[36], expected)
+
+ # Growth past the initial cache keeps the same scaled phase.
+ rotary._ensure_cos_sin_cache_length(48)
+ torch.testing.assert_close(rotary.cos_sin_cache[36], expected)
+
+ def test_yarn_scaling_relaxes_positive_position_bound(self):
+ """A yarn rope_parameters entry must raise the allowed positive bound."""
+ config = _position_config(max_position_embeddings=4)
+ config.text_config.rope_parameters = {
+ "rope_type": "yarn",
+ "factor": 2.0,
+ "original_max_position_embeddings": 4,
+ }
+ input_ids = torch.tensor([10] + [11] * 7)
+
+ # max position is 7: fails against the native bound of 4, passes once
+ # the yarn factor doubles the positive bound to 8.
+ positions, _ = (
+ BailingMRotaryEmbedding.bailing_3drope_get_input_positions_tensor(
+ input_ids,
+ config,
+ image_grid_thw=torch.tensor([[1, 14, 2]]),
+ video_grid_thw=None,
+ )
+ )
+ self.assertEqual(positions.shape, (3, 1, 8))
+
+ def test_rope_scaling_override_merges_into_text_config(self):
+ """A top-level rope_scaling override must not drop the mrope markers."""
+ config = BailingMoeV3VLConfig(
+ mrope_section=[8, 12, 12],
+ text_config={"max_position_embeddings": 131072},
+ vision_config={},
+ )
+ config.rope_scaling = {
+ "rope_type": "yarn",
+ "factor": 2.0,
+ "original_max_position_embeddings": 131072,
+ }
+
+ rope_parameters = config.text_config.rope_parameters
+ self.assertEqual(rope_parameters["rope_type"], "yarn")
+ self.assertEqual(rope_parameters["factor"], 2.0)
+ self.assertEqual(rope_parameters["mrope_section"], [8, 12, 12])
+ self.assertTrue(rope_parameters["video_rope"])
+
+ def test_public_checkpoint_config_contract(self):
+ """External Ling-3.0-flash-VL config literals must survive local parsing."""
+ config = BailingMoeV3VLConfig(
+ image_token_id=157157,
+ video_token_id=156909,
+ mrope_section=[8, 12, 12],
+ text_config={
+ "num_hidden_layers": 42,
+ "vocab_size": 157184,
+ "max_position_embeddings": 131072,
+ "moe_router_enable_expert_bias": True,
+ "num_experts": 512,
+ "num_experts_per_tok": 8,
+ "n_group": 8,
+ "topk_group": 4,
+ "score_function": "sigmoid",
+ "routed_scaling_factor": 2.5,
+ "short_conv_kernel_size": 4,
+ },
+ vision_config={"disable_merger_proj": True},
+ )
+
+ self.assertEqual(config.text_config.num_hidden_layers, 42)
+ self.assertEqual(config.text_config.max_position_embeddings, 131072)
+ self.assertTrue(config.text_config.moe_router_enable_expert_bias)
+ self.assertEqual(
+ config.text_config.rope_parameters["mrope_section"], [8, 12, 12]
+ )
+ self.assertTrue(config.text_config.rope_parameters["video_rope"])
+ self.assertTrue(config.vision_config.disable_merger_proj)
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/test/registered/unit/managers/test_bailing_modality_metadata.py b/test/registered/unit/managers/test_bailing_modality_metadata.py
new file mode 100644
index 000000000..2a75cad74
--- /dev/null
+++ b/test/registered/unit/managers/test_bailing_modality_metadata.py
@@ -0,0 +1,197 @@
+"""Regression tests for Bailing modality metadata and per-token routing bias."""
+
+import unittest
+from types import SimpleNamespace
+
+import torch
+
+from sglang.srt.configs.model_config import requires_mm_token_modalities
+from sglang.srt.layers.moe.topk import biased_grouped_topk_impl
+from sglang.srt.layers.multi_gate import create_multi_gate_mm_indices
+from sglang.srt.managers.schedule_batch import (
+ Modality,
+ MultimodalDataItem,
+ MultimodalInputs,
+ MultimodalProcessorOutput,
+)
+from sglang.srt.model_executor.forward_batch_info import (
+ _build_forward_token_modalities,
+ _maybe_build_forward_token_modalities,
+)
+from sglang.test.ci.ci_register import register_cpu_ci
+from sglang.test.test_utils import CustomTestCase
+
+register_cpu_ci(est_time=15, suite="base-a-test-cpu")
+
+
+class TestBailingModalityMetadata(CustomTestCase):
+ def test_offsets_survive_hash_padding(self):
+ """Hash-derived token replacement must not erase image/audio identity."""
+ items = [
+ MultimodalDataItem(
+ modality=Modality.IMAGE,
+ offsets=[(1, 2)],
+ feature=torch.ones(1),
+ ),
+ MultimodalDataItem(
+ modality=Modality.AUDIO,
+ offsets=[(4, 5)],
+ feature=torch.ones(1),
+ ),
+ ]
+ output = MultimodalProcessorOutput(
+ mm_items=items,
+ input_ids=[100, 11, 11, 101, 12, 12, 102],
+ )
+
+ inputs = MultimodalInputs.from_processor_output(
+ output, requires_mm_token_modalities=True
+ )
+
+ self.assertEqual(
+ inputs.token_modalities,
+ [
+ 0,
+ Modality.IMAGE.value,
+ Modality.IMAGE.value,
+ 0,
+ Modality.AUDIO.value,
+ Modality.AUDIO.value,
+ 0,
+ ],
+ )
+ self.assertNotEqual(items[0].pad_value, 11)
+ self.assertNotEqual(items[1].pad_value, 12)
+
+ def test_offset_validation_only_runs_for_multirouter(self):
+ item = MultimodalDataItem(
+ modality=Modality.IMAGE,
+ offsets=[(1, 3)],
+ feature=torch.ones(1),
+ )
+ output = MultimodalProcessorOutput(mm_items=[item], input_ids=[100, 101])
+
+ inputs = MultimodalInputs.from_processor_output(output)
+ self.assertIsNone(inputs.token_modalities)
+ with self.assertRaisesRegex(ValueError, "Invalid multimodal token offsets"):
+ MultimodalInputs.from_processor_output(
+ output, requires_mm_token_modalities=True
+ )
+
+ def test_chunked_metadata_is_identical_on_every_pp_stage(self):
+ """Each PP stage must independently receive the same active token map."""
+ mm_inputs = [
+ MultimodalInputs(
+ mm_items=[],
+ token_modalities=[0, Modality.IMAGE.value, Modality.IMAGE.value, 0],
+ ),
+ MultimodalInputs(
+ mm_items=[],
+ token_modalities=[Modality.AUDIO.value, Modality.AUDIO.value, 0],
+ ),
+ ]
+ expected = torch.tensor(
+ [
+ Modality.IMAGE.value,
+ Modality.IMAGE.value,
+ 0,
+ Modality.AUDIO.value,
+ Modality.AUDIO.value,
+ ],
+ dtype=torch.int8,
+ )
+
+ stage_maps = [
+ _build_forward_token_modalities(
+ mm_inputs,
+ extend_prefix_lens=[1, 0],
+ extend_seq_lens=[3, 2],
+ num_tokens=5,
+ device=torch.device("cpu"),
+ )
+ for _ in range(2)
+ ]
+
+ for stage_map in stage_maps:
+ torch.testing.assert_close(stage_map, expected)
+
+ def test_only_bailing_multirouter_requires_token_modalities(self):
+ bailing_arch = ["BailingMoeV3VLForConditionalGeneration"]
+ self.assertFalse(
+ requires_mm_token_modalities(
+ bailing_arch, SimpleNamespace(multi_gate=False, router_type="topN")
+ )
+ )
+ self.assertTrue(
+ requires_mm_token_modalities(
+ bailing_arch, SimpleNamespace(multi_gate=True, router_type="topN")
+ )
+ )
+ self.assertFalse(
+ requires_mm_token_modalities(
+ ["DeepseekV4ForCausalLM"],
+ SimpleNamespace(multi_gate=True, router_type="MultiRouter"),
+ )
+ )
+
+ def test_unrelated_model_skips_mismatched_metadata(self):
+ mm_inputs = [
+ MultimodalInputs(mm_items=[], token_modalities=[Modality.IMAGE.value])
+ ]
+ result = _maybe_build_forward_token_modalities(
+ SimpleNamespace(requires_mm_token_modalities=False),
+ mm_inputs,
+ extend_prefix_lens=[0],
+ extend_seq_lens=[1],
+ num_tokens=6,
+ device=torch.device("cpu"),
+ )
+
+ self.assertIsNone(result)
+ with self.assertRaisesRegex(ValueError, "does not match the forward batch"):
+ _maybe_build_forward_token_modalities(
+ SimpleNamespace(requires_mm_token_modalities=True),
+ mm_inputs,
+ extend_prefix_lens=[0],
+ extend_seq_lens=[1],
+ num_tokens=6,
+ device=torch.device("cpu"),
+ )
+
+ def test_mixed_modalities_select_reference_experts(self):
+ """Per-token bias must select image/audio experts after modality grouping."""
+ modalities = torch.tensor(
+ [
+ Modality.IMAGE.value,
+ Modality.IMAGE.value,
+ 0,
+ Modality.AUDIO.value,
+ Modality.AUDIO.value,
+ ],
+ dtype=torch.int8,
+ )
+ token_indices, modality_ids = create_multi_gate_mm_indices(modalities)
+ self.assertEqual(modality_ids.tolist(), [0, 1, 2])
+ self.assertEqual(token_indices[:1].tolist(), [2])
+ self.assertEqual(token_indices[64:66].tolist(), [0, 1])
+ self.assertEqual(token_indices[128:130].tolist(), [3, 4])
+
+ router_logits = torch.zeros(5, 8)
+ dynamic_bias = torch.zeros_like(router_logits)
+ expected_experts = torch.tensor([1, 1, 0, 6, 6], dtype=torch.int32)
+ dynamic_bias.scatter_(1, expected_experts.long().unsqueeze(1), 10.0)
+ _, expert_ids = biased_grouped_topk_impl(
+ hidden_states=torch.zeros(5, 4),
+ gating_output=router_logits,
+ correction_bias=dynamic_bias,
+ topk=1,
+ renormalize=True,
+ num_expert_group=2,
+ topk_group=1,
+ )
+
+ torch.testing.assert_close(expert_ids.squeeze(1), expected_experts)
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/test/registered/unit/models/test_bailing_vl_loader.py b/test/registered/unit/models/test_bailing_vl_loader.py
new file mode 100644
index 000000000..8403c4273
--- /dev/null
+++ b/test/registered/unit/models/test_bailing_vl_loader.py
@@ -0,0 +1,277 @@
+"""Regression tests for streaming Bailing multimodal weight dispatch."""
+
+import unittest
+
+import torch
+import torch.nn as nn
+
+from sglang.srt.configs.bailing_hybrid import (
+ BailingHybridConfig,
+ BailingMoeV3VLConfig,
+ is_bailing_multi_gate_enabled,
+)
+from sglang.srt.models.bailing_mm_v3 import (
+ BailingMoeV3VLForConditionalGeneration,
+)
+from sglang.test.ci.ci_register import register_cpu_ci
+from sglang.test.test_utils import CustomTestCase
+
+register_cpu_ci(est_time=5, suite="base-a-test-cpu")
+
+
+class _OneShotWeights:
+ def __init__(self, values):
+ self.values = values
+ self.iterations = 0
+
+ def __iter__(self):
+ self.iterations += 1
+ if self.iterations > 1:
+ raise AssertionError("checkpoint iterator was consumed more than once")
+ return iter(self.values)
+
+
+class _PublicRouter(nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.weight = nn.Parameter(torch.zeros(3, 2))
+ self.expert_bias = nn.Parameter(torch.zeros(3))
+
+
+class _PublicTextLayer(nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.mlp = nn.Module()
+ self.mlp.gate = _PublicRouter()
+
+
+class _TextModel(nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.model = nn.Module()
+ self.model.word_embeddings = nn.Embedding(3, 2)
+ self.model.layers = nn.ModuleList([_PublicTextLayer()])
+ self.model.norm = nn.LayerNorm(2, bias=False)
+ self.lm_head = nn.Linear(2, 3, bias=False)
+
+ def load_weights(self, weights):
+ params = dict(self.named_parameters())
+ loaded = set()
+ for name, value in weights:
+ params[name].data.copy_(value)
+ loaded.add(name)
+ return loaded
+
+
+class _PublicVisionBlock(nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.attn = nn.Module()
+ self.attn.qkv_proj = nn.Linear(2, 6)
+ self.attn.proj = nn.Linear(2, 2)
+ self.mlp = nn.Module()
+ self.mlp.linear_fc1 = nn.Linear(2, 4)
+ self.mlp.linear_fc2 = nn.Linear(4, 2)
+
+
+class _PublicVision(nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.patch_embed = nn.Module()
+ self.patch_embed.proj = nn.Linear(2, 2)
+ self.pos_embed = nn.Embedding(3, 2)
+ self.blocks = nn.ModuleList([_PublicVisionBlock()])
+ self.merger = nn.Module()
+ self.merger.norm = nn.LayerNorm(2)
+ # Simulate modules created by the inherited Qwen default. Coverage
+ # must ignore them when Bailing deepstack is disabled.
+ self.deepstack_merger_list = nn.ModuleList([nn.Linear(2, 2)])
+
+
+class TestBailingVLWeightLoading(CustomTestCase):
+ @staticmethod
+ def _wrapper():
+ wrapper = BailingMoeV3VLForConditionalGeneration.__new__(
+ BailingMoeV3VLForConditionalGeneration
+ )
+ nn.Module.__init__(wrapper)
+ wrapper.model = _TextModel()
+ wrapper._build_mm_encoders = True
+ wrapper.visual = _PublicVision()
+ wrapper.linear_proj = nn.Sequential(nn.Linear(2, 2), nn.GELU(), nn.Linear(2, 2))
+ wrapper.deepstack_visual_indexes = ()
+ wrapper.multi_gate_enabled = False
+ return wrapper
+
+ @staticmethod
+ def _filled_weights(wrapper, checkpoint_to_parameter):
+ params = dict(wrapper.named_parameters())
+ return [
+ (checkpoint_name, torch.full_like(params[parameter_name], value))
+ for checkpoint_name, parameter_name, value in checkpoint_to_parameter
+ ]
+
+ @classmethod
+ def _public_weights(cls, wrapper):
+ return cls._filled_weights(
+ wrapper,
+ [
+ (
+ "model.word_embeddings.weight",
+ "model.model.word_embeddings.weight",
+ 1,
+ ),
+ (
+ "model.layers.0.mlp.gate.weight",
+ "model.model.layers.0.mlp.gate.weight",
+ 2,
+ ),
+ (
+ "model.layers.0.mlp.gate.expert_bias",
+ "model.model.layers.0.mlp.gate.expert_bias",
+ 3,
+ ),
+ ("model.norm.weight", "model.model.norm.weight", 4),
+ ("lm_head.weight", "model.lm_head.weight", 5),
+ (
+ "model.visual.blocks.0.attn.qkv.weight",
+ "visual.blocks.0.attn.qkv_proj.weight",
+ 6,
+ ),
+ (
+ "model.visual.blocks.0.attn.qkv.bias",
+ "visual.blocks.0.attn.qkv_proj.bias",
+ 7,
+ ),
+ (
+ "model.visual.blocks.0.attn.proj.weight",
+ "visual.blocks.0.attn.proj.weight",
+ 8,
+ ),
+ (
+ "model.visual.blocks.0.attn.proj.bias",
+ "visual.blocks.0.attn.proj.bias",
+ 9,
+ ),
+ (
+ "model.visual.blocks.0.mlp.linear_fc1.weight",
+ "visual.blocks.0.mlp.linear_fc1.weight",
+ 10,
+ ),
+ (
+ "model.visual.blocks.0.mlp.linear_fc1.bias",
+ "visual.blocks.0.mlp.linear_fc1.bias",
+ 11,
+ ),
+ (
+ "model.visual.blocks.0.mlp.linear_fc2.weight",
+ "visual.blocks.0.mlp.linear_fc2.weight",
+ 12,
+ ),
+ (
+ "model.visual.blocks.0.mlp.linear_fc2.bias",
+ "visual.blocks.0.mlp.linear_fc2.bias",
+ 13,
+ ),
+ (
+ "model.visual.patch_embed.proj.weight",
+ "visual.patch_embed.proj.weight",
+ 14,
+ ),
+ (
+ "model.visual.patch_embed.proj.bias",
+ "visual.patch_embed.proj.bias",
+ 15,
+ ),
+ ("model.visual.pos_embed.weight", "visual.pos_embed.weight", 16),
+ ("model.visual.merger.norm.weight", "visual.merger.norm.weight", 17),
+ ("model.visual.merger.norm.bias", "visual.merger.norm.bias", 18),
+ ("linear_proj.0.weight", "linear_proj.0.weight", 19),
+ ("linear_proj.0.bias", "linear_proj.0.bias", 20),
+ ("linear_proj.2.weight", "linear_proj.2.weight", 21),
+ ("linear_proj.2.bias", "linear_proj.2.bias", 22),
+ ],
+ )
+
+ def test_v3_loader_accepts_public_checkpoint_names_once(self):
+ """The public checkpoint layout must load without a second iterator pass."""
+ wrapper = self._wrapper()
+ weights = _OneShotWeights(self._public_weights(wrapper))
+
+ wrapper.load_weights(weights)
+
+ self.assertEqual(weights.iterations, 1)
+ torch.testing.assert_close(
+ wrapper.model.model.layers[0].mlp.gate.expert_bias,
+ torch.full((3,), 3.0),
+ )
+ torch.testing.assert_close(
+ wrapper.visual.blocks[0].attn.qkv_proj.weight,
+ torch.full((6, 2), 6.0),
+ )
+ torch.testing.assert_close(wrapper.linear_proj[2].bias, torch.full((2,), 22.0))
+
+ def test_public_config_does_not_enable_qwen_deepstack_defaults(self):
+ """An omitted public deepstack field must not create random modules."""
+ config = BailingMoeV3VLConfig(vision_config={"disable_merger_proj": True})
+
+ self.assertEqual(config.vision_config.deepstack_visual_indexes, [])
+
+ def test_public_config_selects_standard_single_router(self):
+ """Absent MultiRouter evidence must retain the public single gate and bias."""
+ config = BailingMoeV3VLConfig(
+ text_config={
+ "score_function": "sigmoid",
+ "moe_router_enable_expert_bias": True,
+ "routed_scaling_factor": 2.5,
+ "n_group": 8,
+ "topk_group": 4,
+ "num_experts": 512,
+ "num_experts_per_tok": 8,
+ }
+ )
+
+ self.assertFalse(is_bailing_multi_gate_enabled(config.text_config))
+ self.assertTrue(config.text_config.moe_router_enable_expert_bias)
+ self.assertEqual(config.text_config.score_function, "sigmoid")
+
+ def test_multi_gate_requires_explicit_config_evidence(self):
+ """Internal MultiRouter checkpoints remain reachable only by declaration."""
+ for config in (
+ BailingHybridConfig(multi_gate=True),
+ BailingHybridConfig(router_type="MultiRouter"),
+ ):
+ with self.subTest(config=config):
+ self.assertTrue(is_bailing_multi_gate_enabled(config))
+
+ def test_required_multimodal_weight_coverage_is_enforced(self):
+ """A truncated public checkpoint must not leave random projection bias."""
+ wrapper = self._wrapper()
+ weights = _OneShotWeights(self._public_weights(wrapper)[:-1])
+
+ with self.assertRaisesRegex(
+ RuntimeError, "Missing required Bailing VL weights"
+ ):
+ wrapper.load_weights(weights)
+
+ def test_required_single_router_weight_coverage_is_enforced(self):
+ """Every public MoE layer must load its sole gate weight and expert bias."""
+ wrapper = self._wrapper()
+ public_weights = self._public_weights(wrapper)
+
+ for missing_name in (
+ "model.layers.0.mlp.gate.weight",
+ "model.layers.0.mlp.gate.expert_bias",
+ ):
+ with self.subTest(missing_name=missing_name):
+ weights = _OneShotWeights(
+ [item for item in public_weights if item[0] != missing_name]
+ )
+ with self.assertRaisesRegex(
+ RuntimeError, "Missing required Bailing VL router weights"
+ ):
+ wrapper.load_weights(weights)
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/test/registered/unit/models/test_shared_experts_fusion_gates.py b/test/registered/unit/models/test_shared_experts_fusion_gates.py
index 0e99a1817..7afae42ca 100644
--- a/test/registered/unit/models/test_shared_experts_fusion_gates.py
+++ b/test/registered/unit/models/test_shared_experts_fusion_gates.py
@@ -315,7 +315,14 @@ class TestBailingMoeV3Gate(_FusionGateCase):
packed_modules_mapping={},
)
- def _reason_on_cuda(self, quant_config):
+ def _width_only_config(self):
+ return SimpleNamespace(
+ architectures=["BailingMoeV3ForCausalLM"],
+ moe_intermediate_size=1024,
+ moe_shared_expert_intermediate_size=1024,
+ )
+
+ def _reason_on_cuda(self, quant_config, config=None, model_class=None):
bailing_moe_v3, _ = _import_bailing_modules()
self._seed()
@@ -328,11 +335,61 @@ class TestBailingMoeV3Gate(_FusionGateCase):
),
):
return self._reason(
- bailing_moe_v3.BailingMoeV3ForCausalLM,
- self._config(),
+ model_class or bailing_moe_v3.BailingMoeV3ForCausalLM,
+ config if config is not None else self._config(),
quant_config,
)
+ def test_width_only_fp4_mixed_experts_cannot_fuse(self):
+ quant_config = SimpleNamespace(get_name=lambda: "fp8", is_fp4_experts=True)
+ reason = self._reason_on_cuda(quant_config, self._width_only_config())
+ self.assertIn("different quant methods", reason)
+
+ def test_vl_wrapper_checks_the_width_on_its_text_config(self):
+ from sglang.srt.models.bailing_mm_v3 import (
+ BailingMoeV3VLForConditionalGeneration,
+ )
+
+ quant_config = SimpleNamespace(get_name=lambda: "fp8", is_fp4_experts=True)
+ config = SimpleNamespace(text_config=self._width_only_config())
+ reason = self._reason_on_cuda(
+ quant_config, config, BailingMoeV3VLForConditionalGeneration
+ )
+ self.assertIn("different quant methods", reason)
+
+ def test_width_only_bf16_experts_can_fuse(self):
+ self.assertIsNone(self._reason_on_cuda(None, self._width_only_config()))
+
+ def test_num_shared_experts_only_config_still_fuses(self):
+ self.assertIsNone(self._reason_on_cuda(None, self._config()))
+
+ def test_width_only_int4_mixed_experts_cannot_fuse(self):
+ reason = self._reason_on_cuda(
+ self._compressed_tensors(
+ [r"re:.*mlp\.shared_experts\.(gate|up|down)_proj.*"]
+ ),
+ self._width_only_config(),
+ )
+ self.assertIn("different quant methods", reason)
+
+ def test_width_controls_construction_count(self):
+ bailing_moe_v3, _ = _import_bailing_modules()
+ self.assertEqual(
+ bailing_moe_v3._get_bailing_num_shared_experts(self._width_only_config()),
+ 1,
+ )
+ self.assertEqual(
+ bailing_moe_v3._get_bailing_num_shared_experts(self._config()), 1
+ )
+ legacy_multi_shared = self._config()
+ legacy_multi_shared.num_shared_experts = 2
+ self.assertEqual(
+ bailing_moe_v3._get_bailing_num_shared_experts(legacy_multi_shared), 2
+ )
+ no_shared = self._width_only_config()
+ no_shared.moe_shared_expert_intermediate_size = 0
+ self.assertEqual(bailing_moe_v3._get_bailing_num_shared_experts(no_shared), 0)
+
def test_compressed_tensors_mixed_expert_layout_cannot_fuse(self):
reason = self._reason_on_cuda(
self._compressed_tensors(
diff --git a/test/registered/unit/multimodal/test_bailing_mm_processor.py b/test/registered/unit/multimodal/test_bailing_mm_processor.py
new file mode 100644
index 000000000..803d9dda6
--- /dev/null
+++ b/test/registered/unit/multimodal/test_bailing_mm_processor.py
@@ -0,0 +1,34 @@
+"""Regression tests for the Ling image/video-only processor contract."""
+
+import asyncio
+import unittest
+from types import SimpleNamespace
+
+from sglang.srt.multimodal.processors.bailing_mm import (
+ BailingMMMultimodalProcessor,
+)
+from sglang.test.ci.ci_register import register_cpu_ci
+from sglang.test.test_utils import CustomTestCase
+
+register_cpu_ci(est_time=2, suite="base-a-test-cpu")
+
+
+class TestBailingMMProcessor(CustomTestCase):
+ def test_audio_request_fails_before_preprocessing(self):
+ """The public image/video checkpoint must reject audio at the API boundary."""
+ processor = BailingMMMultimodalProcessor.__new__(BailingMMMultimodalProcessor)
+ request = SimpleNamespace(audio_data=["audio.wav"])
+
+ with self.assertRaisesRegex(ValueError, "Audio inputs are not supported"):
+ asyncio.run(
+ processor.process_mm_data_async(
+ image_data=[],
+ audio_data=request.audio_data,
+ input_text="test",
+ request_obj=request,
+ )
+ )
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/test/registered/unit/multimodal/test_gpu_feature_transport.py b/test/registered/unit/multimodal/test_gpu_feature_transport.py
index 151e0ba9d..f3413bd69 100644
--- a/test/registered/unit/multimodal/test_gpu_feature_transport.py
+++ b/test/registered/unit/multimodal/test_gpu_feature_transport.py
@@ -770,6 +770,7 @@ class TestSchedulerMmTransportBoundary(unittest.TestCase):
@staticmethod
def _prepare_scheduler(scheduler):
+ scheduler.model_config = SimpleNamespace(requires_mm_token_modalities=False)
scheduler.scheduler_stage_metrics = None
scheduler.session_controller = SimpleNamespace(maybe_reap=MagicMock())
scheduler._request_dispatcher = MagicMock(return_value=None)
@@ -785,6 +786,7 @@ class TestSchedulerMmTransportBoundary(unittest.TestCase):
self.mm_inputs = object()
scheduler = object.__new__(scheduler_module.Scheduler)
+ scheduler.model_config = SimpleNamespace(requires_mm_token_modalities=False)
scheduler.dp_tp_cpu_group = object()
request = TokenizedRequest()
@@ -861,7 +863,9 @@ class TestSchedulerMmTransportBoundary(unittest.TestCase):
):
scheduler.process_input_requests([request])
- build_inputs.assert_called_once_with(raw_inputs)
+ build_inputs.assert_called_once_with(
+ raw_inputs, requires_mm_token_modalities=False
+ )
self.assertIs(request.mm_inputs, materialized)
scheduler._request_dispatcher.assert_called_once_with(request)
cpu_broadcast.assert_not_called()
@@ -914,7 +918,7 @@ class TestSchedulerMmTransportBoundary(unittest.TestCase):
self.assertEqual(
build_inputs.call_args_list,
- [call(value) for value in raw_inputs],
+ [call(value, requires_mm_token_modalities=False) for value in raw_inputs],
)
self.assertEqual(
[inner.mm_inputs for inner in inner_requests],
@@ -940,6 +944,7 @@ class TestSchedulerMmTransportBoundary(unittest.TestCase):
from sglang.srt.managers import scheduler as scheduler_module
scheduler = object.__new__(scheduler_module.Scheduler)
+ scheduler.model_config = SimpleNamespace(requires_mm_token_modalities=False)
scheduler.dp_tp_group = SimpleNamespace(rank_in_group=0, first_rank=0)
scheduler.dp_tp_cpu_group = object()
diff --git a/test/registered/unit/parser/test_template_manager.py b/test/registered/unit/parser/test_template_manager.py
index 1d1da64fa..ba5282103 100644
--- a/test/registered/unit/parser/test_template_manager.py
+++ b/test/registered/unit/parser/test_template_manager.py
@@ -17,6 +17,7 @@ from sglang.srt.parser.template_detection import (
)
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
+from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2.0, suite="base-a-test-cpu")
@@ -58,7 +59,7 @@ def _glm53_template(concat):
)
-class TestTemplateManagerReasoningDetection(unittest.TestCase):
+class TestTemplateManagerReasoningDetection(CustomTestCase):
def _detect(self, template, vocab):
force, config = detect_reasoning_pattern(template)
parser = detect_reasoning_parser(
@@ -99,6 +100,25 @@ class TestTemplateManagerReasoningDetection(unittest.TestCase):
)
self.assertEqual(parser, "glm45")
+ def test_ling3_template_uses_ling3_parsers(self):
+ template = """
+ {% set enable_thinking = enable_thinking if enable_thinking is defined else true %}
+ {{ 'SYSTEM' }}
+ {{ 'ASSISTANT' }}
+ {{ '<|role_end|>' }}
+ {function-name}
+ {arg-key}
+ {arg-value}
+
+ """
+ force, config, reasoning_parser = self._detect(template, [])
+ tool_call_parser = detect_tool_call_parser(
+ template, _DummyTokenizer([]), config, force
+ )
+
+ self.assertEqual(reasoning_parser, "ling3")
+ self.assertEqual(tool_call_parser, "ling3")
+
def test_glm53_effort_template_resolves_glm_parsers(self):
# Without an enable_thinking toggle the GLM-4.5 rule misses, and the
# template used to fall through to deepseek-r1 + the xml_kv fallback
@@ -929,7 +949,7 @@ def _declared(server_args, field):
return resolution_result(server_args, field)
-class TestResolveAutoParsers(unittest.TestCase):
+class TestResolveAutoParsers(CustomTestCase):
"""Tests for resolve_auto_parsers()."""
qwen3_template = "{% set enable_thinking = enable_thinking if enable_thinking is defined else true %}"
@@ -1065,6 +1085,31 @@ class TestResolveAutoParsers(unittest.TestCase):
self.assertEqual(_declared(args, "reasoning_parser"), "kimi_k3")
self.assertEqual(_declared(args, "tool_call_parser"), "kimi_k3")
+ def test_bailing_architectures_and_model_types_use_ling3_parsers(self):
+ cases = (
+ (["BailingMoeV3VLForConditionalGeneration"], ""),
+ (None, "bailing_moe_v3_vl"),
+ (["BailingMoeV3ForCausalLM"], ""),
+ (None, "bailing_hybrid"),
+ )
+ for architectures, model_type in cases:
+ with self.subTest(architectures=architectures, model_type=model_type):
+ args = self._make_server_args(
+ reasoning_parser="auto", tool_call_parser="auto"
+ )
+ tokenizer = _DummyTokenizer([])
+ config = SimpleNamespace(
+ architectures=architectures, model_type=model_type
+ )
+
+ with _patch_hf_transformers_utils(
+ Mock(return_value=tokenizer), Mock(return_value=config)
+ ):
+ resolve_auto_parsers(args)
+
+ self.assertEqual(_declared(args, "reasoning_parser"), "ling3")
+ self.assertEqual(_declared(args, "tool_call_parser"), "ling3")
+
def test_deepseek_arch_fallback_runs_when_tokenizer_load_fails(self):
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
config = SimpleNamespace(architectures=["DeepseekV32ForCausalLM"])
diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py
index ab2c54d45..0dfeb162c 100644
--- a/test/registered/unit/test_model_overrides.py
+++ b/test/registered/unit/test_model_overrides.py
@@ -1777,6 +1777,74 @@ class TestGoldenModelOverrides(_IsolatedPublish):
"flashinfer_trtllm_routed",
)
+ def test_bailing_v3_mixed_mxfp4_selects_native_runner(self):
+ """Packed MXFP4 experts must not reach the FP8 Triton runner."""
+
+ def _args(**kw):
+ defaults = dict(
+ device="cuda",
+ moe_a2a_backend="none",
+ moe_runner_backend="auto",
+ _model_config=SimpleNamespace(quantization="fp8", is_fp4_experts=True),
+ )
+ defaults.update(kw)
+ return SimpleNamespace(**defaults)
+
+ with override_platform(
+ is_sm90=False, is_sm100=True, is_sm120=False, is_hip=False
+ ):
+ for architecture in (
+ "BailingMoeV3ForCausalLM",
+ "BailingMoeV3VLForConditionalGeneration",
+ ):
+ with self.subTest(architecture=architecture):
+ declarations = collect_model_override_declarations(
+ architecture,
+ _args(),
+ SimpleNamespace(architectures=[architecture]),
+ )
+ self.assertEqual(
+ declarations,
+ [
+ (
+ "_bailing_moe_v3_overrides",
+ {"moe_runner_backend": "flashinfer_mxfp4"},
+ )
+ ],
+ )
+
+ from sglang.srt.arg_groups.model_overrides.bailing_moe_v3 import (
+ _bailing_moe_v3_overrides,
+ )
+
+ hf = SimpleNamespace(
+ architectures=["BailingMoeV3VLForConditionalGeneration"]
+ )
+ self.assertEqual(
+ _bailing_moe_v3_overrides(_args(moe_runner_backend="triton"), hf),
+ {},
+ )
+ self.assertEqual(
+ _bailing_moe_v3_overrides(_args(moe_a2a_backend="deepep"), hf),
+ {},
+ )
+ self.assertEqual(
+ _bailing_moe_v3_overrides(
+ _args(
+ _model_config=SimpleNamespace(
+ quantization="fp8", is_fp4_experts=False
+ )
+ ),
+ hf,
+ ),
+ {},
+ )
+
+ with override_platform(
+ is_sm90=False, is_sm100=False, is_sm120=False, is_hip=False
+ ):
+ self.assertEqual(_bailing_moe_v3_overrides(_args(), hf), {})
+
def test_nemotron_h_overrides_at_callable_level(self):
from sglang.srt.arg_groups.model_overrides.nemotron_h import (
_nemotron_h_overrides,