diff --git a/docker/Dockerfile b/docker/Dockerfile
index 78b8054bc..49d56652f 100644
--- a/docker/Dockerfile
+++ b/docker/Dockerfile
@@ -744,6 +744,7 @@ RUN --mount=type=cache,target=/var/cache/apt,id=runtime-apt \
ca-certificates \
netcat-openbsd \
curl \
+ ffmpeg \
git \
# Runtime libraries
libopenmpi3 \
diff --git a/docs/docs/supported-models/multimodal_language_models.mdx b/docs/docs/supported-models/multimodal_language_models.mdx
index ac706a6c1..c9d2b396e 100644
--- a/docs/docs/supported-models/multimodal_language_models.mdx
+++ b/docs/docs/supported-models/multimodal_language_models.mdx
@@ -171,6 +171,18 @@ in the GitHub search bar.
NVIDIA Nemotron Nano v2 VL enables multi-image reasoning and video understanding, along with strong document intelligence, visual Q&A and summarization capabilities. It builds on Nemotron Nano V2, a hybrid Mamba-Transformer LLM, in order to achieve higher inference throughput in long document and video scenarios. |
Use --trust-remote-code. You may need to adjust --max-mamba-cache-size [default is 512] to fit memory constraints. |
+
+ | NVIDIA Cosmos3 Reasoner (Nano) |
+ nvidia/Cosmos3-Nano |
+ The understanding tower (Qwen3-VL-based reasoner) of NVIDIA's Cosmos3 world model, served as a standalone VLM for image and video understanding from the unified diffusers-layout checkpoint. The generation (diffusion) tower is dropped at load time. |
+ Launch with --model-type llm; see Cosmos3 LLM serving. Pass --language-model-only for text-only serving. |
+
+
+ | NVIDIA Cosmos3-Edge (4B) |
+ nvidia/Cosmos3-Edge |
+ The 4B dense understanding tower of Cosmos3-Edge: an Arcee-structure text model with a SigLIP2 vision tower and an Edge-specific spatial-merge projector, supporting image and video understanding. |
+ Launch with --model-type llm; see Cosmos3 LLM serving. Pass --language-model-only for text-only serving. |
+
| Ernie4.5-VL |
baidu/ERNIE-4.5-VL-28B-A3B-PT |
@@ -349,6 +361,45 @@ response = requests.post(url, json=data)
print(response.text)
```
+## Serving the LLM tower of Cosmos3 checkpoints
+
+`nvidia/Cosmos3-Nano` and `nvidia/Cosmos3-Edge` are unified world-model checkpoints in a diffusers layout, so `sglang serve` routes them to the diffusion runtime by default (see the [Cosmos3 cookbook page](/cookbook/diffusion/Cosmos/Cosmos3) for media generation). To serve the understanding tower as a vision-language model instead, force the LLM backend with `--model-type llm`:
+
+```bash Launch Server
+sglang serve \
+ --model-path nvidia/Cosmos3-Nano \
+ --model-type llm \
+ --host 0.0.0.0 --port 30000
+```
+
+The same command serves `nvidia/Cosmos3-Edge` by swapping the model path. Both accept image and video inputs through the OpenAI-compatible API:
+
+```bash Example Request
+curl http://localhost:30000/v1/chat/completions \
+ -H "Content-Type: application/json" \
+ -d '{
+ "model": "nvidia/Cosmos3-Nano",
+ "messages": [{
+ "role": "user",
+ "content": [
+ {"type": "text", "text": "Describe this image."},
+ {"type": "image_url", "image_url": {"url": "https://raw.githubusercontent.com/sgl-project/sglang/main/examples/assets/example_image.png"}}
+ ]
+ }],
+ "max_tokens": 128
+ }'
+```
+
+To serve only the text model (skipping the vision tower entirely), add `--language-model-only`. In this mode `/model_info` reports `has_image_understanding: false`, and requests containing image or video inputs are rejected:
+
+```bash Text-Only Launch
+sglang serve \
+ --model-path nvidia/Cosmos3-Nano \
+ --model-type llm \
+ --language-model-only \
+ --host 0.0.0.0 --port 30000
+```
+
## Usage Notes
### Performance Optimization
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/cosmos3.py b/python/sglang/multimodal_gen/configs/pipeline_configs/cosmos3.py
index 81e97360e..38c17b62c 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/cosmos3.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/cosmos3.py
@@ -90,11 +90,6 @@ def _distilled_sampler_config(model_path: str) -> dict | None:
return sampler
-def is_distilled_checkpoint(model_path: str) -> bool:
- """Whether the checkpoint is a few-step distilled variant."""
- return _distilled_sampler_config(model_path) is not None
-
-
def get_distilled_sigmas(model_path: str) -> list[float] | None:
"""The explicit fixed-step sigma schedule for a distilled checkpoint."""
sampler = _distilled_sampler_config(model_path)
diff --git a/python/sglang/srt/configs/__init__.py b/python/sglang/srt/configs/__init__.py
index bc6211b5e..fedbf7708 100644
--- a/python/sglang/srt/configs/__init__.py
+++ b/python/sglang/srt/configs/__init__.py
@@ -2,6 +2,13 @@ from sglang.srt.configs.afmoe import AfmoeConfig
from sglang.srt.configs.bailing_hybrid import BailingHybridConfig
from sglang.srt.configs.chatglm import ChatGLMConfig
from sglang.srt.configs.cohere2_moe import Cohere2MoeConfig
+from sglang.srt.configs.cosmos3 import (
+ Cosmos3Config,
+ Cosmos3EdgeConfig,
+ Cosmos3EdgeProjectorConfig,
+ Cosmos3EdgeTextConfig,
+ Cosmos3EdgeVisionConfig,
+)
from sglang.srt.configs.dbrx import DbrxConfig
from sglang.srt.configs.deepseekvl2 import DeepseekVL2Config
from sglang.srt.configs.dots3 import Dots3Config
@@ -74,6 +81,11 @@ __all__ = [
"BailingHybridConfig",
"ExaoneConfig",
"ChatGLMConfig",
+ "Cosmos3Config",
+ "Cosmos3EdgeConfig",
+ "Cosmos3EdgeTextConfig",
+ "Cosmos3EdgeVisionConfig",
+ "Cosmos3EdgeProjectorConfig",
"DbrxConfig",
"DeepseekVL2Config",
"LongcatFlashConfig",
diff --git a/python/sglang/srt/configs/cosmos3.py b/python/sglang/srt/configs/cosmos3.py
new file mode 100644
index 000000000..7069be3b8
--- /dev/null
+++ b/python/sglang/srt/configs/cosmos3.py
@@ -0,0 +1,240 @@
+# SPDX-License-Identifier: Apache-2.0
+"""Configuration for the Cosmos3 Reasoner (understanding tower).
+
+The Cosmos3 unified checkpoint stores a Qwen3-VL understanding tower alongside
+a generation (diffusion) tower. The Reasoner only serves the understanding
+tower, so it reuses the Qwen3-VL config schema and just declares its own
+``model_type`` so ``AutoConfig`` can resolve the checkpoint.
+"""
+
+from typing import Optional, Type
+
+from transformers import PretrainedConfig
+
+from sglang.srt.configs.qwen3_vl import Qwen3VLConfig
+
+
+class Cosmos3Config(Qwen3VLConfig):
+ model_type = "cosmos3_omni"
+
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ # The Qwen3-VL inference stack accesses ``config.vision_config`` and
+ # ``config.text_config`` as objects (e.g. ``config.vision_config.hidden_size``,
+ # ``config.vision_config.deepstack_visual_indexes``). Some transformers
+ # versions leave these sub-configs
+ # as raw dicts after construction, which would raise
+ # ``'dict' object has no attribute 'hidden_size'`` at model init. Coerce
+ # any dict-valued sub-config into its proper config object so the model
+ # loads regardless of the installed transformers version.
+ for attr, sub_cls in self.sub_configs.items():
+ sub = getattr(self, attr, None)
+ if isinstance(sub, dict):
+ setattr(self, attr, sub_cls(**sub))
+
+
+def _coerce_sub_config(
+ value: Optional[object],
+ config_cls: Type[PretrainedConfig],
+) -> PretrainedConfig:
+ if value is None:
+ return config_cls()
+ if isinstance(value, config_cls):
+ return value
+ if isinstance(value, dict):
+ return config_cls(**value)
+ if isinstance(value, PretrainedConfig):
+ return value
+ raise TypeError(f"Unsupported sub-config type: {type(value)!r}")
+
+
+def _normalize_edge_rope_parameters(value: Optional[dict]) -> Optional[dict]:
+ if value is None:
+ return None
+ rope_parameters = dict(value)
+ mrope_section = rope_parameters.get("mrope_section")
+ if mrope_section is not None:
+ rope_parameters["mrope_section"] = list(mrope_section)
+ rope_parameters.setdefault("mrope_interleaved", True)
+ return rope_parameters
+
+
+class Cosmos3EdgeTextConfig(PretrainedConfig):
+ model_type = "cosmos3_edge_text"
+ ignore_keys_at_rope_validation = {"mrope_section", "mrope_interleaved"}
+
+ def __init__(
+ self,
+ vocab_size: int = 131072,
+ hidden_size: int = 2048,
+ intermediate_size: int = 9216,
+ num_hidden_layers: int = 28,
+ num_attention_heads: int = 16,
+ num_key_value_heads: int = 8,
+ head_dim: int = 128,
+ hidden_act: str = "relu2",
+ rms_norm_eps: float = 1e-5,
+ max_position_embeddings: int = 131072,
+ initializer_range: float = 0.02,
+ attention_bias: bool = False,
+ attention_dropout: float = 0.0,
+ mlp_bias: bool = False,
+ rope_parameters: Optional[dict] = None,
+ rope_scaling: Optional[dict] = None,
+ use_cache: bool = True,
+ bos_token_id: Optional[int] = None,
+ eos_token_id: Optional[int] = 11,
+ pad_token_id: Optional[int] = None,
+ tie_word_embeddings: bool = False,
+ **kwargs,
+ ) -> None:
+ super().__init__(
+ bos_token_id=bos_token_id,
+ eos_token_id=eos_token_id,
+ pad_token_id=pad_token_id,
+ tie_word_embeddings=tie_word_embeddings,
+ **kwargs,
+ )
+ self.vocab_size = vocab_size
+ self.hidden_size = hidden_size
+ self.intermediate_size = intermediate_size
+ self.num_hidden_layers = num_hidden_layers
+ self.num_attention_heads = num_attention_heads
+ self.num_key_value_heads = num_key_value_heads
+ self.head_dim = head_dim
+ self.hidden_act = hidden_act
+ self.rms_norm_eps = rms_norm_eps
+ self.max_position_embeddings = max_position_embeddings
+ self.initializer_range = initializer_range
+ self.attention_bias = attention_bias
+ self.attention_dropout = attention_dropout
+ self.mlp_bias = mlp_bias
+ self.use_cache = use_cache
+
+ if rope_parameters is None:
+ rope_parameters = rope_scaling
+ if rope_parameters is None:
+ rope_parameters = {
+ "mrope_section": [24, 20, 20],
+ "rope_theta": 100000000,
+ "rope_type": "default",
+ }
+ rope_parameters = _normalize_edge_rope_parameters(rope_parameters)
+ self.rope_parameters = rope_parameters
+ # SGLang's RoPE factory accepts the v5-style rope_parameters schema, but
+ # several generic paths still probe rope_scaling.
+ self.rope_scaling = _normalize_edge_rope_parameters(
+ rope_scaling if rope_scaling is not None else rope_parameters
+ )
+
+
+class Cosmos3EdgeVisionConfig(PretrainedConfig):
+ model_type = "cosmos3_edge_vision"
+
+ def __init__(
+ self,
+ hidden_size: int = 1152,
+ intermediate_size: int = 4304,
+ num_hidden_layers: int = 27,
+ num_attention_heads: int = 16,
+ num_channels: int = 3,
+ num_patches: int = 256,
+ patch_size: int = 16,
+ hidden_act: str = "gelu_pytorch_tanh",
+ layer_norm_eps: float = 1e-6,
+ attention_dropout: float = 0.0,
+ **kwargs,
+ ) -> None:
+ super().__init__(**kwargs)
+ self.hidden_size = hidden_size
+ self.intermediate_size = intermediate_size
+ self.num_hidden_layers = num_hidden_layers
+ self.num_attention_heads = num_attention_heads
+ self.num_channels = num_channels
+ self.num_patches = num_patches
+ self.patch_size = patch_size
+ self.hidden_act = hidden_act
+ self.layer_norm_eps = layer_norm_eps
+ self.attention_dropout = attention_dropout
+
+
+class Cosmos3EdgeProjectorConfig(PretrainedConfig):
+ model_type = "cosmos3_edge_projector"
+
+ def __init__(
+ self,
+ input_hidden_size: int = 1152,
+ merger_intermediate_size: int = 11520,
+ out_hidden_size: int = 2048,
+ spatial_merge_size: int = 2,
+ use_postshuffle_norm: bool = False,
+ **kwargs,
+ ) -> None:
+ super().__init__(**kwargs)
+ self.input_hidden_size = input_hidden_size
+ self.merger_intermediate_size = merger_intermediate_size
+ self.out_hidden_size = out_hidden_size
+ self.spatial_merge_size = spatial_merge_size
+ self.use_postshuffle_norm = use_postshuffle_norm
+
+
+class Cosmos3EdgeConfig(PretrainedConfig):
+ model_type = "cosmos3_edge"
+ sub_configs = {
+ "text_config": Cosmos3EdgeTextConfig,
+ "vision_config": Cosmos3EdgeVisionConfig,
+ "projector_config": Cosmos3EdgeProjectorConfig,
+ }
+
+ def __init__(
+ self,
+ text_config: Optional[object] = None,
+ vision_config: Optional[object] = None,
+ projector_config: Optional[object] = None,
+ image_token_id: int = 19,
+ video_token_id: int = 18,
+ vision_start_token_id: int = 20,
+ vision_end_token_id: int = 21,
+ tie_word_embeddings: bool = False,
+ bos_token_id: Optional[int] = None,
+ eos_token_id: Optional[int] = None,
+ pad_token_id: Optional[int] = None,
+ **kwargs,
+ ) -> None:
+ super().__init__(
+ bos_token_id=bos_token_id,
+ eos_token_id=eos_token_id,
+ pad_token_id=pad_token_id,
+ tie_word_embeddings=tie_word_embeddings,
+ **kwargs,
+ )
+ self.text_config = _coerce_sub_config(text_config, Cosmos3EdgeTextConfig)
+ self.vision_config = _coerce_sub_config(vision_config, Cosmos3EdgeVisionConfig)
+ self.projector_config = _coerce_sub_config(
+ projector_config, Cosmos3EdgeProjectorConfig
+ )
+
+ # Qwen-style multimodal processing reads these from vision_config, while
+ # the Cosmos3-Edge checkpoint stores them under projector_config.
+ self.vision_config.spatial_merge_size = self.projector_config.spatial_merge_size
+ self.vision_config.temporal_patch_size = 1
+ self.vision_config.out_hidden_size = self.projector_config.out_hidden_size
+
+ 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
+ self.tie_word_embeddings = tie_word_embeddings
+
+ if getattr(self, "architectures", None) is None:
+ self.architectures = ["Cosmos3EdgeForConditionalGeneration"]
+
+ for attr in ("bos_token_id", "eos_token_id", "pad_token_id"):
+ parent_value = getattr(self, attr, None)
+ text_value = getattr(self.text_config, attr, None)
+ if parent_value is None and text_value is not None:
+ setattr(self, attr, text_value)
+ elif parent_value is not None and text_value is None:
+ setattr(self.text_config, attr, parent_value)
+ if not hasattr(self.text_config, "tie_word_embeddings"):
+ self.text_config.tie_word_embeddings = tie_word_embeddings
diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py
index 2a87a7e0b..590d0dd34 100644
--- a/python/sglang/srt/configs/model_config.py
+++ b/python/sglang/srt/configs/model_config.py
@@ -353,7 +353,13 @@ class ModelConfig:
rope_scaling = getattr(self.hf_text_config, "rope_parameters", None) or getattr(
self.hf_text_config, "rope_scaling", {}
)
- self.is_lm_only = getattr(self.hf_config, "language_model_only", False)
+ # Text-only serving comes from either the checkpoint's own declaration
+ # or the --language-model-only flag; every capability flag below
+ # (is_multimodal, is_*_understandable_model) must see both sources,
+ # since /model_info advertises them and drives media warmup requests.
+ self.is_lm_only = language_model_only or getattr(
+ self.hf_config, "language_model_only", False
+ )
self.model_is_mrope = (
not self.is_lm_only
and rope_scaling is not None
@@ -588,7 +594,7 @@ class ModelConfig:
# Checkpoints declare this one themselves (hf_transformers/processor.py),
# so the flag may only turn it on: writing the default back would build a
# vision tower with no weights to fill.
- self.hf_config.language_model_only = language_model_only or self.is_lm_only
+ self.hf_config.language_model_only = self.is_lm_only
# matryoshka embeddings
self.matryoshka_dimensions = getattr(
@@ -1353,8 +1359,23 @@ class ModelConfig:
return quant_cfg
def _parse_modelopt_quant_config(self, quant_config_dict: dict) -> Optional[dict]:
- """Parse ModelOpt quantization config and return the appropriate quant_method."""
- json_quant_configs = quant_config_dict["quantization"]
+ """Parse ModelOpt quantization config and return the appropriate quant_method.
+
+ Supports both nested LLM ``hf_quant_config.json``::
+
+ {"quantization": {"quant_algo": "FP8", ...}}
+
+ and flat formats (``config.json`` ``quantization_config``, or diffusion /
+ unified ModelOpt exports such as Cosmos3)::
+
+ {"quant_algo": "FP8", "quant_method": "modelopt", ...}
+ """
+ if "quantization" in quant_config_dict:
+ json_quant_configs = quant_config_dict["quantization"]
+ elif "quant_algo" in quant_config_dict:
+ json_quant_configs = quant_config_dict
+ else:
+ return None
quant_algo = json_quant_configs.get("quant_algo", None)
if quant_algo == "MIXED_PRECISION":
@@ -1886,6 +1907,7 @@ def is_generation_model(model_architectures: List[str], is_embedding: bool = Fal
multimodal_model_archs = [
"CLIPModel",
"Cohere2VisionForConditionalGeneration",
+ "Cosmos3EdgeForConditionalGeneration",
"DeepseekVL2ForCausalLM",
"Ernie4_5_VLMoeForConditionalGeneration",
"MiniMaxM3SparseForConditionalGeneration",
diff --git a/python/sglang/srt/disaggregation/encoder/preprocessor.py b/python/sglang/srt/disaggregation/encoder/preprocessor.py
index f05229d8b..364747746 100644
--- a/python/sglang/srt/disaggregation/encoder/preprocessor.py
+++ b/python/sglang/srt/disaggregation/encoder/preprocessor.py
@@ -550,6 +550,7 @@ class EncoderPreprocessor:
"qwen3_5",
"qwen3_5_moe",
"intern_s2_preview",
+ "cosmos3_omni",
]
and video_processor_kwargs.get("video_metadata", None) is not None
):
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 c3ed2b4cd..e074bd84a 100644
--- a/python/sglang/srt/layers/rotary_embedding/mrope_rope_index.py
+++ b/python/sglang/srt/layers/rotary_embedding/mrope_rope_index.py
@@ -75,6 +75,8 @@ def get_rope_index(
or model_type.startswith("qwen3_vl_moe")
or model_type.startswith("qwen3_5")
or model_type.startswith("interns2_mobius")
+ or model_type.startswith("cosmos3_omni")
+ or model_type.startswith("cosmos3_edge")
) and video_grid_thw is not None:
video_grid_thw = torch.repeat_interleave(
video_grid_thw, video_grid_thw[:, 0], dim=0
@@ -162,6 +164,8 @@ def get_rope_index(
"qwen3_5_moe",
"intern_s2_preview",
"interns2_mobius",
+ "cosmos3_omni",
+ "cosmos3_edge",
):
t_index = (
torch.arange(llm_grid_t, device=position_ids.device)
diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py
index 59f5ea01a..011c2cc35 100644
--- a/python/sglang/srt/model_loader/loader.py
+++ b/python/sglang/srt/model_loader/loader.py
@@ -378,6 +378,12 @@ class DefaultModelLoader(BaseModelLoader):
fall_back_to_pt: bool = True
"""Whether .pt weights can be used."""
+ allow_patterns_overrides: Optional[list[str]] = None
+ """If defined, weights will load exclusively using these patterns.
+
+ Used by checkpoints whose weights live in subfolders (e.g. the Cosmos3
+ diffusers-style layout with ``transformer/`` and ``vision_encoder/``)."""
+
model_config: Optional[ModelConfig] = None
"""The model configuration (for checking architecture, etc)."""
@@ -388,6 +394,9 @@ class DefaultModelLoader(BaseModelLoader):
model_config.revision,
prefix="",
fall_back_to_pt=getattr(model, "fall_back_to_pt_during_load", True),
+ allow_patterns_overrides=getattr(
+ model, "allow_patterns_overrides", None
+ ),
model_config=model_config,
)
@@ -437,7 +446,11 @@ class DefaultModelLoader(BaseModelLoader):
return model
def _prepare_weights(
- self, model_name_or_path: str, revision: Optional[str], fall_back_to_pt: bool
+ self,
+ model_name_or_path: str,
+ revision: Optional[str],
+ fall_back_to_pt: bool,
+ allow_patterns_overrides: Optional[list[str]] = None,
) -> Tuple[str, List[str], bool]:
"""Prepare weights for the model.
@@ -477,6 +490,9 @@ class DefaultModelLoader(BaseModelLoader):
if fall_back_to_pt:
allow_patterns += ["*.pt"]
+ if allow_patterns_overrides is not None:
+ allow_patterns = allow_patterns_overrides
+
if not is_local:
hf_folder = download_weights_from_hf(
model_name_or_path,
@@ -499,7 +515,7 @@ class DefaultModelLoader(BaseModelLoader):
for pattern in allow_patterns:
hf_weights_files += glob.glob(os.path.join(hf_folder, pattern))
if len(hf_weights_files) > 0:
- if pattern == "*.safetensors":
+ if pattern.endswith(".safetensors"):
use_safetensors = True
break
@@ -517,7 +533,12 @@ class DefaultModelLoader(BaseModelLoader):
revision,
)
hf_weights_files = filter_duplicate_safetensors_files(
- hf_weights_files, hf_folder, index_file
+ hf_weights_files,
+ hf_folder,
+ index_file,
+ allow_patterns=(
+ allow_patterns if allow_patterns_overrides is not None else None
+ ),
)
else:
hf_weights_files = filter_files_not_needed_for_inference(hf_weights_files)
@@ -557,9 +578,13 @@ class DefaultModelLoader(BaseModelLoader):
"""Get an iterator for the model weights based on the load format."""
extra_config = self.load_config.model_loader_extra_config
use_multithread = extra_config.get("enable_multithread_load", True)
+
if resolved_source is None:
hf_folder, hf_weights_files, use_safetensors = self._prepare_weights(
- source.model_or_path, source.revision, source.fall_back_to_pt
+ source.model_or_path,
+ source.revision,
+ source.fall_back_to_pt,
+ source.allow_patterns_overrides,
)
if use_safetensors and source.model_config is not None:
hf_weights_files = maybe_add_mtp_safetensors(
@@ -727,6 +752,7 @@ class DefaultModelLoader(BaseModelLoader):
source.model_or_path,
source.revision,
source.fall_back_to_pt,
+ source.allow_patterns_overrides,
)
if use_safetensors and source.model_config is not None:
weight_files = maybe_add_mtp_safetensors(
diff --git a/python/sglang/srt/model_loader/weight_utils.py b/python/sglang/srt/model_loader/weight_utils.py
index fd659e7c9..a784d02da 100644
--- a/python/sglang/srt/model_loader/weight_utils.py
+++ b/python/sglang/srt/model_loader/weight_utils.py
@@ -259,6 +259,23 @@ def _resolve_explicit_draft_quant_config(
return quant_config
+def _modelopt_quant_section(config: dict) -> dict:
+ """Return ModelOpt quant settings from nested or flat ``hf_quant_config.json``.
+
+ Nested LLM format::
+
+ {"quantization": {"quant_algo": "FP8", "exclude_modules": [...]}}
+
+ Flat format (``config.json`` ``quantization_config`` / Cosmos3-style exports)::
+
+ {"quant_algo": "FP8", "ignore": [...], "quant_method": "modelopt", ...}
+ """
+ quantization = config.get("quantization")
+ if isinstance(quantization, dict):
+ return quantization
+ return config
+
+
# TODO(woosuk): Move this to other place.
def get_quant_config(
model_config: ModelConfig,
@@ -386,12 +403,17 @@ def get_quant_config(
quant_config_file = quant_config_files[0]
with open(quant_config_file) as f:
config = json.load(f)
+ quant_section = _modelopt_quant_section(config)
if remap_prefix is not None:
- exclude_modules = [
- replace_prefix(key, remap_prefix)
- for key in config["quantization"]["exclude_modules"]
- ]
- config["quantization"]["exclude_modules"] = exclude_modules
+ # Nested configs use ``exclude_modules``; flat ModelOpt exports use ``ignore``.
+ exclude_key = (
+ "exclude_modules" if "exclude_modules" in quant_section else "ignore"
+ )
+ if exclude_key in quant_section:
+ quant_section[exclude_key] = [
+ replace_prefix(key, remap_prefix)
+ for key in quant_section[exclude_key]
+ ]
config["packed_modules_mapping"] = packed_modules_mapping
if model_config.quantization == "bitsandbytes":
@@ -399,7 +421,7 @@ def get_quant_config(
elif model_config.quantization.startswith("modelopt") and (
config.get("producer", {}).get("name", "").startswith("modelopt")
):
- quant_algo = config["quantization"]["quant_algo"]
+ quant_algo = quant_section.get("quant_algo")
if quant_algo is None:
# (yizhang2077) workaround for nvidia/Llama-4-Maverick-17B-128E-Eagle3
if model_config.hf_config.architectures[0] != "LlamaForCausalLMEagle3":
@@ -423,15 +445,19 @@ def get_quant_config(
)
-def _check_index_files_exist(snapshot_dir: str) -> Tuple[bool, Optional[str]]:
+def _check_index_files_exist(
+ snapshot_dir: str, allow_patterns: Optional[List[str]] = None
+) -> Tuple[bool, Optional[str]]:
"""
- Check if all files listed in safetensors index files actually exist on disk.
+ Check if files listed in safetensors index files actually exist on disk.
This catches cases where the snapshot directory exists but files are missing
- (e.g., due to incomplete downloads or corrupted cache).
+ (e.g., due to incomplete downloads or corrupted cache). If allow_patterns is
+ provided, only indexed files matching those patterns are validated.
Args:
snapshot_dir: Path to the model snapshot directory
+ allow_patterns: Optional source patterns to scope validation.
Returns:
Tuple of (all_exist, error_message)
@@ -453,6 +479,15 @@ def _check_index_files_exist(snapshot_dir: str) -> Tuple[bool, Optional[str]]:
if not weight_map:
continue
required_files = set(weight_map.values())
+ if allow_patterns is not None:
+ required_files = {
+ fn
+ for fn in required_files
+ if any(
+ fnmatch.fnmatch(fn.replace(os.sep, "/"), pattern)
+ for pattern in allow_patterns
+ )
+ }
missing_files = [
fn
for fn in required_files
@@ -555,7 +590,9 @@ def _find_local_hf_snapshot_dir_unlocked(
# Check for missing files from index (lightweight, for all users)
# This catches incomplete downloads before they cause cryptic load errors
if local_weight_files:
- is_complete, error_msg = _check_index_files_exist(found_local_snapshot_dir)
+ is_complete, error_msg = _check_index_files_exist(
+ found_local_snapshot_dir, allow_patterns
+ )
if not is_complete:
log_info_on_rank0(
logger,
@@ -715,7 +752,10 @@ def download_safetensors_index_file_from_hf(
# So, we use the index_file to
# look up which safetensors files should be used.
def filter_duplicate_safetensors_files(
- hf_weights_files: List[str], hf_folder: str, index_file: str
+ hf_weights_files: List[str],
+ hf_folder: str,
+ index_file: str,
+ allow_patterns: Optional[List[str]] = None,
) -> List[str]:
# model.safetensors.index.json is a mapping from keys in the
# torch state_dict to safetensors file holding that weight.
@@ -739,9 +779,18 @@ def filter_duplicate_safetensors_files(
for weight_name in weight_map:
weight_files_in_index.add(os.path.join(hf_folder, weight_map[weight_name]))
# Fail fast if the index references shard files that are not on disk (e.g. an
- # incomplete or interrupted download). Otherwise those shards are silently
- # dropped and the model loads with uninitialized weights.
- missing_files = sorted(f for f in weight_files_in_index if not os.path.isfile(f))
+ # incomplete or interrupted download). For subfolder-scoped loads, only
+ # validate the indexed shards that match the requested source patterns.
+ if allow_patterns is None:
+ files_to_validate = weight_files_in_index
+ else:
+ files_to_validate = set()
+ for f in weight_files_in_index:
+ rel_path = os.path.relpath(f, hf_folder).replace(os.sep, "/")
+ if any(fnmatch.fnmatch(rel_path, pattern) for pattern in allow_patterns):
+ files_to_validate.add(f)
+
+ missing_files = sorted(f for f in files_to_validate if not os.path.isfile(f))
if missing_files:
raise RuntimeError(
f"{index_file} references {len(missing_files)} shard file(s) missing "
diff --git a/python/sglang/srt/models/cosmos3.py b/python/sglang/srt/models/cosmos3.py
new file mode 100644
index 000000000..3ac172d8d
--- /dev/null
+++ b/python/sglang/srt/models/cosmos3.py
@@ -0,0 +1,116 @@
+# Copyright 2023-2025 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.
+# ==============================================================================
+"""Inference-only Cosmos3 Reasoner (understanding tower) model.
+
+Cosmos3 ships a unified diffusers-layout checkpoint that stores a Qwen3-VL
+understanding tower alongside a generation (diffusion) tower. The Reasoner
+serves only the understanding tower, so it reuses the Qwen3-VL inference stack
+and drops the generation-tower weights at load time.
+
+The checkpoint keeps the LLM weights under ``transformer/`` and the vision
+encoder weights under ``vision_encoder/``, so the two are loaded from separate
+subfolders via ``allow_patterns_overrides`` / ``secondary_weights``.
+"""
+
+from typing import Iterable, Optional, Tuple
+
+import torch
+
+from sglang.srt.configs.cosmos3 import Cosmos3Config
+from sglang.srt.layers.quantization.base_config import QuantizationConfig
+from sglang.srt.model_loader.loader import DefaultModelLoader
+from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration
+from sglang.srt.models.utils import WeightsMapper
+from sglang.srt.runtime_context import get_model
+
+
+class Cosmos3ForConditionalGeneration(Qwen3VLForConditionalGeneration):
+ # Cosmos3 unified checkpoints store a Qwen3-VL understanding tower alongside
+ # a generation tower in a flat key layout. This mapper drops the generation
+ # tower weights and rewrites the understanding tower keys into the nested
+ # Qwen3-VL checkpoint form consumed by the parent ``load_weights``.
+ hf_to_sglang_mapper = WeightsMapper(
+ orig_to_new_substr={
+ # Drop ModelOpt calibration buffers. The FP8 export already ships
+ # inference ``weight_scale`` / ``input_scale``; transformers restores
+ # ``*_quantizer._amax`` via ModelOpt HF checkpointing, but SGLang's
+ # ModelOptFp8 path does not register those modules. Same drop as the
+ # diffusion Cosmos3 loader.
+ "_quantizer.": None,
+ # Drop the generation (diffusion) tower.
+ "_moe_gen": None,
+ ".add_q_proj.": None,
+ ".add_k_proj.": None,
+ ".add_v_proj.": None,
+ ".to_add_out.": None,
+ ".norm_added_q.": None,
+ ".norm_added_k.": None,
+ # Understanding-tower attention projections -> Qwen3 names.
+ ".to_q.": ".q_proj.",
+ ".to_k.": ".k_proj.",
+ ".to_v.": ".v_proj.",
+ ".to_out.": ".o_proj.",
+ ".norm_q.": ".q_norm.",
+ ".norm_k.": ".k_norm.",
+ },
+ orig_to_new_prefix={
+ # Understanding-tower (LLM) keys -> nested language-model namespace.
+ "layers.": "model.language_model.layers.",
+ "embed_tokens.": "model.language_model.embed_tokens.",
+ "norm.": "model.language_model.norm.",
+ # Vision-encoder keys -> visual namespace.
+ "blocks.": "model.visual.blocks.",
+ "merger.": "model.visual.merger.",
+ "patch_embed.": "model.visual.patch_embed.",
+ "pos_embed.": "model.visual.pos_embed.",
+ "deepstack_merger_list.": "model.visual.deepstack_merger_list.",
+ # Diffusion-only latent/timestep/modality heads -> dropped.
+ "proj_in.": None,
+ "proj_out.": None,
+ "time_embedder.": None,
+ "audio_": None,
+ "action_": None,
+ },
+ )
+
+ # The understanding-tower LLM weights live in the ``transformer/`` subfolder
+ # of the diffusers-layout checkpoint.
+ allow_patterns_overrides = ["transformer/*.safetensors"]
+
+ def __init__(
+ self,
+ config: Cosmos3Config,
+ quant_config: Optional[QuantizationConfig] = None,
+ prefix: str = "",
+ ) -> None:
+ super().__init__(config=config, quant_config=quant_config, prefix=prefix)
+
+ # The vision encoder weights live in a separate ``vision_encoder/``
+ # subfolder, so load them as a secondary weight source.
+ self.secondary_weights = []
+ if not self.language_model_only:
+ self.secondary_weights.append(
+ DefaultModelLoader.Source(
+ model_or_path=get_model().model_path,
+ revision=get_model().revision,
+ prefix="",
+ allow_patterns_overrides=["vision_encoder/*.safetensors"],
+ )
+ )
+
+ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
+ return super().load_weights(self.hf_to_sglang_mapper.apply(weights))
+
+
+EntryClass = Cosmos3ForConditionalGeneration
diff --git a/python/sglang/srt/models/cosmos3_edge.py b/python/sglang/srt/models/cosmos3_edge.py
new file mode 100644
index 000000000..67bd7db42
--- /dev/null
+++ b/python/sglang/srt/models/cosmos3_edge.py
@@ -0,0 +1,527 @@
+# Copyright 2023-2025 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.
+# ==============================================================================
+"""Inference-only Cosmos3-Edge VLM.
+
+Cosmos3-Edge stores a dense UND text tower in ``transformer/`` and a SigLIP2
+vision tower plus projector in ``vision_encoder/``. The text tower matches the
+Arcee causal-LM structure used by SGLang, while the vision path uses the native
+SigLIP2 implementation and an Edge-specific spatial-merge projector.
+"""
+
+from typing import Iterable, List, Optional, Tuple
+
+import numpy as np
+import torch
+from torch import nn
+
+from sglang.srt.configs.cosmos3 import (
+ Cosmos3EdgeConfig,
+ Cosmos3EdgeProjectorConfig,
+ Cosmos3EdgeTextConfig,
+)
+from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
+from sglang.srt.layers.logits_processor import LogitsProcessorOutput
+from sglang.srt.layers.quantization.base_config import QuantizationConfig
+from sglang.srt.managers.mm_utils import (
+ MultiModalityDataPaddingPatternMultimodalTokens,
+ embed_mm_inputs,
+)
+from sglang.srt.managers.schedule_batch import (
+ MultimodalDataItem,
+ MultimodalInputFormat,
+ 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.arcee import ArceeForCausalLM
+from sglang.srt.models.siglip2 import Siglip2Model
+from sglang.srt.models.utils import WeightsMapper
+from sglang.srt.utils import add_prefix
+
+
+class Cosmos3EdgeVisionProjector(nn.Module):
+ def __init__(
+ self,
+ config: Cosmos3EdgeProjectorConfig,
+ quant_config: Optional[QuantizationConfig] = None,
+ prefix: str = "",
+ ) -> None:
+ super().__init__()
+ self.spatial_merge_size = config.spatial_merge_size
+ self.use_postshuffle_norm = config.use_postshuffle_norm
+ self.context_dim = config.input_hidden_size
+ self.hidden_size = self.context_dim * (self.spatial_merge_size**2)
+
+ norm_dim = self.hidden_size if self.use_postshuffle_norm else self.context_dim
+ self.norm = nn.LayerNorm(norm_dim, eps=1e-6)
+ self.linear_fc1 = ColumnParallelLinear(
+ self.hidden_size,
+ config.merger_intermediate_size,
+ bias=True,
+ quant_config=quant_config,
+ prefix=add_prefix("linear_fc1", prefix),
+ )
+ self.act_fn = nn.GELU()
+ self.linear_fc2 = RowParallelLinear(
+ config.merger_intermediate_size,
+ config.out_hidden_size,
+ bias=True,
+ quant_config=quant_config,
+ prefix=add_prefix("linear_fc2", prefix),
+ )
+
+ def _spatial_merge(
+ self, vision_features: torch.Tensor, spatial_shapes: torch.Tensor
+ ) -> torch.Tensor:
+ merge_size = self.spatial_merge_size
+ hidden_size = vision_features.shape[-1]
+ lengths = (spatial_shapes[:, 0] * spatial_shapes[:, 1]).tolist()
+ tiles = torch.split(vision_features, lengths, dim=0)
+
+ merged_parts = []
+ for tile, (height, width) in zip(tiles, spatial_shapes.tolist()):
+ height = int(height)
+ width = int(width)
+ if height == 0 or width == 0:
+ continue
+ if height % merge_size != 0 or width % merge_size != 0:
+ raise ValueError(
+ "Cosmos3-Edge vision grid must be divisible by "
+ f"spatial_merge_size={merge_size}, got {(height, width)}."
+ )
+ tile = tile.view(height, width, hidden_size)
+ tile = tile.view(
+ height // merge_size,
+ merge_size,
+ width // merge_size,
+ merge_size,
+ hidden_size,
+ )
+ tile = tile.permute(0, 2, 1, 3, 4).reshape(
+ (height // merge_size) * (width // merge_size),
+ merge_size * merge_size * hidden_size,
+ )
+ merged_parts.append(tile)
+
+ if not merged_parts:
+ return vision_features.new_empty((0, merge_size * merge_size * hidden_size))
+ return torch.cat(merged_parts, dim=0)
+
+ def forward(
+ self, vision_features: torch.Tensor, spatial_shapes: torch.Tensor
+ ) -> torch.Tensor:
+ if self.use_postshuffle_norm:
+ vision_features = self._spatial_merge(vision_features, spatial_shapes)
+ vision_features = self.norm(vision_features)
+ else:
+ vision_features = self.norm(vision_features)
+ vision_features = self._spatial_merge(vision_features, spatial_shapes)
+
+ hidden_states, _ = self.linear_fc1(vision_features)
+ hidden_states = self.act_fn(hidden_states)
+ hidden_states, _ = self.linear_fc2(hidden_states)
+ return hidden_states
+
+
+class Cosmos3EdgeForConditionalGeneration(ArceeForCausalLM):
+ # Multimodal serving needs both text and vision subfolders. In
+ # --language-model-only mode __init__ narrows this instance attribute to the
+ # transformer shards so the vision files are not downloaded or loaded.
+ allow_patterns_overrides = ["[tv]*er/*.safetensors"]
+
+ hf_to_sglang_mapper = WeightsMapper(
+ orig_to_new_substr={
+ # Drop ModelOpt calibration buffers and generation-side tensors.
+ "_quantizer.": None,
+ "_moe_gen": None,
+ "k_norm_und_for_gen": None,
+ ".add_q_proj.": None,
+ ".add_k_proj.": None,
+ ".add_v_proj.": None,
+ ".to_add_out.": None,
+ ".norm_added_q.": None,
+ ".norm_added_k.": None,
+ # Text attention projection names -> SGLang/Arcee names.
+ ".to_q.": ".q_proj.",
+ ".to_k.": ".k_proj.",
+ ".to_v.": ".v_proj.",
+ ".to_out.": ".o_proj.",
+ },
+ orig_to_new_prefix={
+ "embed_tokens.": "model.embed_tokens.",
+ "layers.": "model.layers.",
+ "norm.": "model.norm.",
+ # Diffusion-only top-level modules.
+ "proj_in.": None,
+ "proj_out.": None,
+ "time_embedder.": None,
+ "action_": None,
+ "audio_": None,
+ # Vision/projector weights are routed before applying this mapper.
+ "model.visual.": None,
+ "model.projector.": None,
+ },
+ )
+
+ def __init__(
+ self,
+ config: Cosmos3EdgeConfig,
+ quant_config: Optional[QuantizationConfig] = None,
+ prefix: str = "",
+ ) -> None:
+ self.root_config = config
+ self.language_model_only = bool(getattr(config, "language_model_only", False))
+ self.allow_patterns_overrides = (
+ ["transformer/*.safetensors"]
+ if self.language_model_only
+ else ["[tv]*er/*.safetensors"]
+ )
+
+ text_config = getattr(config, "text_config", config)
+ if isinstance(text_config, dict):
+ text_config = Cosmos3EdgeTextConfig(**text_config)
+ super().__init__(
+ config=text_config,
+ quant_config=quant_config,
+ prefix=prefix,
+ )
+
+ self.image_token_id = getattr(config, "image_token_id", None)
+ self.video_token_id = getattr(config, "video_token_id", None)
+
+ if not self.language_model_only:
+ self.visual = Siglip2Model(
+ config=config.vision_config,
+ quant_config=quant_config,
+ prefix=add_prefix("visual", prefix),
+ )
+ self.projector = Cosmos3EdgeVisionProjector(
+ config=config.projector_config,
+ quant_config=quant_config,
+ prefix=add_prefix("projector", prefix),
+ )
+
+ def pad_input_ids(
+ self, input_ids: List[int], mm_inputs: MultimodalInputs
+ ) -> List[int]:
+ pattern = MultiModalityDataPaddingPatternMultimodalTokens()
+ return pattern.pad_input_tokens(input_ids, mm_inputs)
+
+ @staticmethod
+ def _as_tensor(value) -> Optional[torch.Tensor]:
+ if value is None:
+ return None
+ if isinstance(value, torch.Tensor):
+ return value
+ if isinstance(value, np.ndarray):
+ return torch.from_numpy(value)
+ return torch.as_tensor(value)
+
+ @classmethod
+ def _get_item_value(cls, item: MultimodalDataItem, *names: str):
+ for name in names:
+ try:
+ value = getattr(item, name)
+ except AttributeError:
+ value = None
+ if value is not None:
+ return value
+ return None
+
+ @classmethod
+ def _get_spatial_shapes(cls, item: MultimodalDataItem) -> torch.Tensor:
+ spatial_shapes = cls._get_item_value(
+ item, "spatial_shapes", "image_grid_hws", "grid_hws"
+ )
+ if spatial_shapes is not None:
+ spatial_shapes = cls._as_tensor(spatial_shapes).to(dtype=torch.long)
+ if spatial_shapes.ndim == 1:
+ spatial_shapes = spatial_shapes.view(1, -1)
+ if spatial_shapes.shape[-1] == 3:
+ rows = []
+ for t, height, width in spatial_shapes.view(-1, 3).tolist():
+ rows.extend([[height, width]] * int(t))
+ return torch.tensor(rows, dtype=torch.long)
+ if spatial_shapes.shape[-1] != 2:
+ raise ValueError(
+ "Cosmos3-Edge spatial_shapes must have shape (..., 2) or (..., 3), "
+ f"got {tuple(spatial_shapes.shape)}."
+ )
+ return spatial_shapes.view(-1, 2)
+
+ grid = cls._get_item_value(item, "image_grid_thw", "video_grid_thw")
+ grid = cls._as_tensor(grid).to(dtype=torch.long) if grid is not None else None
+ if grid is None:
+ raise ValueError(
+ "Cosmos3-Edge vision item is missing spatial_shapes or *_grid_thw."
+ )
+ if grid.ndim == 1:
+ grid = grid.view(1, -1)
+ if grid.shape[-1] != 3:
+ raise ValueError(
+ "Cosmos3-Edge grid metadata must have shape (..., 3), "
+ f"got {tuple(grid.shape)}."
+ )
+ rows = []
+ for t, height, width in grid.view(-1, 3).tolist():
+ rows.extend([[height, width]] * int(t))
+ return torch.tensor(rows, dtype=torch.long)
+
+ def _pack_visual_items(
+ self, items: List[MultimodalDataItem]
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
+ packed_features = []
+ all_spatial_shapes = []
+
+ for item in items:
+ pixel_values = self._as_tensor(item.feature)
+ if pixel_values is None:
+ raise ValueError("Cosmos3-Edge vision item is missing pixel values.")
+
+ spatial_shapes = self._get_spatial_shapes(item).cpu()
+ lengths = (spatial_shapes[:, 0] * spatial_shapes[:, 1]).tolist()
+ expected_tokens = int(sum(lengths))
+
+ if pixel_values.ndim == 2:
+ if pixel_values.shape[0] != expected_tokens:
+ raise ValueError(
+ "Cosmos3-Edge packed pixel count does not match "
+ f"spatial_shapes: {pixel_values.shape[0]} vs "
+ f"{expected_tokens}."
+ )
+ packed_features.append(pixel_values)
+ elif pixel_values.ndim == 3:
+ if pixel_values.shape[0] != len(lengths):
+ raise ValueError(
+ "Cosmos3-Edge padded pixel batch does not match "
+ f"spatial_shapes: {pixel_values.shape[0]} vs "
+ f"{len(lengths)}."
+ )
+ attention_mask = self._as_tensor(
+ self._get_item_value(item, "pixel_attention_mask", "attention_mask")
+ )
+ for idx, length in enumerate(lengths):
+ if attention_mask is None:
+ packed_features.append(pixel_values[idx, : int(length)])
+ else:
+ mask = attention_mask[idx].reshape(-1).bool()
+ packed_features.append(pixel_values[idx][mask])
+ else:
+ raise ValueError(
+ "Cosmos3-Edge pixel_values must be packed 2D or padded 3D, "
+ f"got {tuple(pixel_values.shape)}."
+ )
+
+ all_spatial_shapes.append(spatial_shapes)
+
+ spatial_shapes_cpu = torch.cat(all_spatial_shapes, dim=0)
+ pixel_values_packed = torch.cat(packed_features, dim=0).to(
+ device=self.visual.device,
+ dtype=self.visual.dtype,
+ )
+ lengths = (spatial_shapes_cpu[:, 0] * spatial_shapes_cpu[:, 1]).to(
+ dtype=torch.int32, device=pixel_values_packed.device
+ )
+ cu_seqlens = torch.zeros(
+ lengths.numel() + 1, dtype=torch.int32, device=pixel_values_packed.device
+ )
+ cu_seqlens[1:] = torch.cumsum(lengths, dim=0)
+ max_seqlen = lengths.max()
+ return pixel_values_packed, spatial_shapes_cpu, cu_seqlens, max_seqlen
+
+ def _get_visual_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
+ if self.language_model_only:
+ raise RuntimeError("Cosmos3-Edge was loaded with --language-model-only.")
+ if not items:
+ return torch.empty(0, device=self.visual.device, dtype=self.visual.dtype)
+
+ if any(
+ item.format == MultimodalInputFormat.PRECOMPUTED_EMBEDDING for item in items
+ ):
+ if not all(
+ item.format == MultimodalInputFormat.PRECOMPUTED_EMBEDDING
+ for item in items
+ ):
+ raise ValueError(
+ "Cosmos3-Edge cannot mix raw features and precomputed "
+ "embeddings within the same modality."
+ )
+ embeddings = [self._as_tensor(item.feature) for item in items]
+ if any(embedding is None for embedding in embeddings):
+ raise ValueError(
+ "Cosmos3-Edge precomputed embedding items must contain feature."
+ )
+ result = torch.cat(embeddings, dim=0)
+ return result.reshape(-1, result.shape[-1])
+
+ pixel_values_packed, spatial_shapes, cu_seqlens, max_seqlen = (
+ self._pack_visual_items(items)
+ )
+ vision_outputs = self.visual(
+ pixel_values_packed=pixel_values_packed,
+ spatial_shapes=spatial_shapes,
+ cu_seqlens=cu_seqlens,
+ max_seqlen=max_seqlen,
+ )
+ if vision_outputs.dim() == 3:
+ vision_outputs = vision_outputs[0]
+ return self.projector(vision_outputs, spatial_shapes)
+
+ def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
+ return self._get_visual_feature(items)
+
+ def get_video_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
+ return self._get_visual_feature(items)
+
+ def _embed_multimodal_inputs(
+ self, input_ids: torch.Tensor, forward_batch: ForwardBatch
+ ) -> Optional[torch.Tensor]:
+ if self.pp_group.is_first_rank:
+ if (
+ not forward_batch.forward_mode.is_decode()
+ and not forward_batch.forward_mode.is_target_verify()
+ and forward_batch.contains_mm_inputs()
+ ):
+ mm_inputs_list = [
+ mm_input
+ for mm_input in forward_batch.mm_inputs
+ if mm_input is not None
+ ]
+ extend_prefix_lens = [
+ prefix_len
+ for i, prefix_len in enumerate(forward_batch.extend_prefix_lens_cpu)
+ if forward_batch.mm_inputs[i] is not None
+ ]
+ extend_seq_lens = [
+ seq_len
+ for i, seq_len in enumerate(forward_batch.extend_seq_lens_cpu)
+ if forward_batch.mm_inputs[i] is not None
+ ]
+ input_embeds, _ = embed_mm_inputs(
+ mm_inputs_list=mm_inputs_list,
+ extend_prefix_lens=extend_prefix_lens,
+ extend_seq_lens=extend_seq_lens,
+ input_ids=input_ids,
+ input_embedding=self.get_input_embeddings(),
+ multimodal_model=self,
+ )
+
+ for mm_input in mm_inputs_list:
+ if mm_input and hasattr(mm_input, "mm_items"):
+ for item in mm_input.mm_items:
+ feature = getattr(item, "feature", None)
+ if isinstance(feature, torch.Tensor) and feature.is_cuda:
+ item.feature = feature.to("cpu", non_blocking=True)
+ forward_batch.mm_inputs = None
+ forward_batch.mm_input_embeds = input_embeds
+ else:
+ input_embeds = self.get_input_embeddings()(input_ids)
+
+ if forward_batch.input_embeds is not None:
+ forward_batch.input_embeds.copy_(input_embeds)
+ input_embeds = forward_batch.input_embeds
+ return input_embeds
+ return None
+
+ @torch.no_grad()
+ def forward(
+ self,
+ input_ids: torch.Tensor,
+ positions: torch.Tensor,
+ forward_batch: ForwardBatch,
+ input_embeds: torch.Tensor = None,
+ get_embedding: bool = False,
+ pp_proxy_tensors: Optional[PPProxyTensors] = None,
+ ) -> LogitsProcessorOutput:
+ if forward_batch.mrope_positions is not None:
+ positions = forward_batch.mrope_positions
+
+ needs_mm_embedding = (
+ not forward_batch.forward_mode.is_decode()
+ and not forward_batch.forward_mode.is_target_verify()
+ and forward_batch.contains_mm_inputs()
+ )
+ if (
+ input_embeds is not None
+ or self.language_model_only
+ or not needs_mm_embedding
+ ):
+ return super().forward(
+ input_ids=input_ids,
+ positions=positions,
+ forward_batch=forward_batch,
+ input_embeds=input_embeds,
+ get_embedding=get_embedding,
+ pp_proxy_tensors=pp_proxy_tensors,
+ )
+
+ input_embeds = self._embed_multimodal_inputs(input_ids, forward_batch)
+ hidden_states = self.model(
+ input_ids=input_ids,
+ positions=positions,
+ forward_batch=forward_batch,
+ input_embeds=input_embeds,
+ pp_proxy_tensors=pp_proxy_tensors,
+ )
+
+ aux_hidden_states = None
+ if self.capture_aux_hidden_states:
+ hidden_states, aux_hidden_states = hidden_states
+
+ if self.pp_group.is_last_rank:
+ if not get_embedding:
+ return self.logits_processor(
+ input_ids,
+ hidden_states,
+ self.lm_head,
+ forward_batch,
+ aux_hidden_states,
+ )
+ return self.pooler(hidden_states, forward_batch)
+ return hidden_states
+
+ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
+ text_weights = []
+ visual_weights = []
+ projector_weights = []
+
+ for name, loaded_weight in weights:
+ if name.startswith("model.visual."):
+ if not self.language_model_only:
+ new_name = name.replace("model.visual.", "vision_model.", 1)
+ visual_weights.append((new_name, loaded_weight))
+ elif name.startswith("model.projector."):
+ if not self.language_model_only:
+ new_name = name.replace("model.projector.", "projector.", 1)
+ projector_weights.append((new_name, loaded_weight))
+ else:
+ text_weights.append((name, loaded_weight))
+
+ super().load_weights(self.hf_to_sglang_mapper.apply(text_weights))
+
+ if self.language_model_only:
+ return
+
+ self.visual.load_weights(visual_weights)
+
+ params_dict = dict(self.named_parameters())
+ for name, loaded_weight in projector_weights:
+ if name not in params_dict:
+ continue
+ param = params_dict[name]
+ weight_loader = getattr(param, "weight_loader", default_weight_loader)
+ weight_loader(param, loaded_weight)
+
+
+EntryClass = Cosmos3EdgeForConditionalGeneration
diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py
index ad7b5ce01..5d2faac6f 100644
--- a/python/sglang/srt/models/qwen3_vl.py
+++ b/python/sglang/srt/models/qwen3_vl.py
@@ -1595,10 +1595,9 @@ class Qwen3VLForConditionalGeneration(nn.Module):
# Skip loading extra bias for GPTQ models.
if name.endswith(".bias") and name not in params_dict:
continue
- # Skip loading visual/language model weights
- if (
- self.config.encoder_only or self.config.language_only
- ) and name not in params_dict:
+ # Skip unexpected stacked names (e.g. ModelOpt quantizer buffers
+ # that were remapped gate_proj -> gate_up_proj but are not params).
+ if name not in params_dict:
continue
param = params_dict[name]
weight_loader = param.weight_loader
diff --git a/python/sglang/srt/multimodal/processors/cosmos3_edge.py b/python/sglang/srt/multimodal/processors/cosmos3_edge.py
new file mode 100644
index 000000000..a99245885
--- /dev/null
+++ b/python/sglang/srt/multimodal/processors/cosmos3_edge.py
@@ -0,0 +1,723 @@
+# Copyright 2023-2025 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.
+# ==============================================================================
+"""Cosmos3-Edge multimodal processor.
+
+The current supported Transformers release may not ship the Cosmos3-Edge
+processor classes yet. This processor keeps serving unblocked by applying the
+checkpoint's SigLIP2-style image/video patchification directly in SGLang.
+"""
+
+import math
+from typing import Any, List, Optional, Union
+
+import numpy as np
+import torch
+from PIL import Image
+
+from sglang.srt.layers.rotary_embedding import MRotaryEmbedding
+from sglang.srt.managers.schedule_batch import (
+ Modality,
+ MultimodalDataItem,
+ MultimodalProcessorOutput,
+)
+from sglang.srt.models.cosmos3_edge import Cosmos3EdgeForConditionalGeneration
+from sglang.srt.multimodal.processors.base_processor import (
+ BaseMultimodalProcessor as SGLangBaseProcessor,
+)
+from sglang.srt.multimodal.processors.base_processor import (
+ MultimodalSpecialTokens,
+)
+from sglang.srt.utils.video_decoder import VideoDecoderWrapper
+
+IMAGE_MIN_PIXELS = 256 * 256
+IMAGE_MAX_PIXELS = 4096 * 4096
+VIDEO_MIN_PIXELS = 64 * 64
+VIDEO_TOTAL_PIXELS = 6144 * 4096
+MAX_RATIO = 200
+DEFAULT_SOURCE_VIDEO_FPS = 24.0
+DEFAULT_TARGET_VIDEO_FPS = 2.0
+DEFAULT_MIN_FRAMES = 4
+DEFAULT_MAX_FRAMES = 768
+
+
+def _round_by_factor(number: float, factor: int) -> int:
+ return round(number / factor) * factor
+
+
+def _ceil_by_factor(number: float, factor: int) -> int:
+ return math.ceil(number / factor) * factor
+
+
+def _floor_by_factor(number: float, factor: int) -> int:
+ return math.floor(number / factor) * factor
+
+
+def _smart_resize(
+ height: int,
+ width: int,
+ *,
+ factor: int,
+ min_pixels: int,
+ max_pixels: int,
+ num_frames: int = 1,
+) -> tuple[int, int]:
+ if num_frames <= 0:
+ raise ValueError(f"num_frames must be positive, got {num_frames}")
+ if max(height, width) / min(height, width) > MAX_RATIO:
+ raise ValueError(
+ "absolute aspect ratio must be smaller than "
+ f"{MAX_RATIO}, got {max(height, width) / min(height, width)}"
+ )
+
+ h_bar = max(factor, _round_by_factor(height, factor))
+ w_bar = max(factor, _round_by_factor(width, factor))
+ if num_frames * h_bar * w_bar > max_pixels:
+ beta = math.sqrt((num_frames * height * width) / max_pixels)
+ h_bar = max(factor, _floor_by_factor(height / beta, factor))
+ w_bar = max(factor, _floor_by_factor(width / beta, factor))
+ elif num_frames * h_bar * w_bar < min_pixels:
+ beta = math.sqrt(min_pixels / (num_frames * height * width))
+ h_bar = _ceil_by_factor(height * beta, factor)
+ w_bar = _ceil_by_factor(width * beta, factor)
+ return h_bar, w_bar
+
+
+def _as_pil_image(image: Any) -> Image.Image:
+ if isinstance(image, Image.Image):
+ return image.convert("RGB")
+
+ if isinstance(image, torch.Tensor):
+ image = image.detach().cpu()
+ if image.ndim == 3 and image.shape[0] in (1, 3, 4):
+ image = image.permute(1, 2, 0)
+ image = image.numpy()
+
+ image = np.asarray(image)
+ if image.ndim == 2:
+ image = np.stack([image] * 3, axis=-1)
+ if image.ndim != 3:
+ raise ValueError(
+ f"Expected an image with 2 or 3 dimensions, got {image.shape}."
+ )
+ if image.shape[0] in (1, 3, 4) and image.shape[-1] not in (1, 3, 4):
+ image = np.moveaxis(image, 0, -1)
+ if image.shape[-1] == 1:
+ image = np.repeat(image, 3, axis=-1)
+ if image.shape[-1] == 4:
+ image = image[..., :3]
+
+ if image.dtype != np.uint8:
+ image = image.astype(np.float32)
+ if image.size and image.max() <= 1.0:
+ image = image * 255.0
+ image = np.clip(image, 0, 255).astype(np.uint8)
+ return Image.fromarray(image).convert("RGB")
+
+
+class Cosmos3EdgeProcessor(SGLangBaseProcessor):
+ models = [Cosmos3EdgeForConditionalGeneration]
+ gpu_image_decode = False
+
+ @staticmethod
+ def _get_processor_output_value(output, key: str):
+ if output is None:
+ return None
+ return output.get(key) if hasattr(output, "get") else getattr(output, key, None)
+
+ @staticmethod
+ def _as_grid_batch(value) -> Optional[torch.Tensor]:
+ if value is None:
+ return None
+ grid = torch.as_tensor(value, dtype=torch.long)
+ return grid.unsqueeze(0) if grid.ndim == 1 else grid
+
+ @classmethod
+ def _get_grid_from_output_or_items(
+ cls,
+ output,
+ mm_items: list[MultimodalDataItem],
+ key: str,
+ modality: Modality,
+ ) -> Optional[torch.Tensor]:
+ grid = cls._as_grid_batch(cls._get_processor_output_value(output, key))
+ if grid is not None:
+ return grid
+
+ grids = []
+ for item in mm_items:
+ if not item.is_modality(modality):
+ continue
+ item_grid = cls._as_grid_batch(item.model_specific_data.get(key))
+ if item_grid is not None:
+ grids.append(item_grid)
+ return torch.cat(grids, dim=0) if grids else None
+
+ def _get_precomputed_mrope(self, output):
+ positions = self._get_processor_output_value(output, "mrope_positions")
+ delta = self._get_processor_output_value(output, "mrope_position_delta")
+ if positions is None or delta is None:
+ return None
+
+ positions = torch.as_tensor(positions)
+ if positions.ndim == 3:
+ if positions.shape[1] != 1:
+ return None
+ positions = positions.squeeze(1)
+ if positions.ndim != 2 or positions.shape[0] != 3:
+ return None
+
+ delta = torch.as_tensor(delta)
+ if delta.ndim <= 1:
+ delta = delta.reshape(-1, 1)
+ return positions, delta
+
+ def _make_processor_output(
+ self,
+ input_ids: Union[list[int], torch.Tensor],
+ mm_items: list[MultimodalDataItem],
+ image_grid_thw: Optional[torch.Tensor],
+ video_grid_thw: Optional[torch.Tensor],
+ processor_output=None,
+ ) -> MultimodalProcessorOutput:
+ input_ids = torch.as_tensor(input_ids, dtype=torch.long).flatten()
+ mrope_result = self._get_precomputed_mrope(processor_output)
+ if mrope_result is None:
+ has_images = any(item.is_image() for item in mm_items)
+ has_videos = any(item.is_video() for item in mm_items)
+ if has_images and image_grid_thw is None:
+ raise ValueError(
+ "Cosmos3-Edge processed image input requires image_grid_thw "
+ "or precomputed MRoPE positions."
+ )
+ if has_videos and video_grid_thw is None:
+ raise ValueError(
+ "Cosmos3-Edge processed video input requires video_grid_thw "
+ "or precomputed MRoPE positions."
+ )
+ mrope_result = MRotaryEmbedding.get_rope_index(
+ spatial_merge_size=self.spatial_merge_size,
+ image_token_id=self.mm_tokens.image_token_id,
+ video_token_id=self.mm_tokens.video_token_id,
+ vision_start_token_id=self.vision_start_token_id,
+ model_type=self.model_type,
+ input_ids=input_ids.unsqueeze(0),
+ image_grid_thw=image_grid_thw,
+ video_grid_thw=video_grid_thw,
+ )
+
+ mrope_positions, mrope_position_delta = mrope_result
+ if mrope_positions.ndim == 3:
+ mrope_positions = mrope_positions.squeeze(1)
+
+ return MultimodalProcessorOutput(
+ input_ids=input_ids.tolist(),
+ mm_items=mm_items,
+ im_start_id=self.IM_START_TOKEN_ID,
+ im_end_id=self.IM_END_TOKEN_ID,
+ im_token_id=self.IMAGE_TOKEN_ID,
+ video_token_id=self.VIDEO_TOKEN_ID,
+ mrope_positions=mrope_positions,
+ mrope_position_delta=mrope_position_delta,
+ )
+
+ async def _process_preprocessed_mm_data(self, base_output):
+ mm_items, input_ids, processor_output = (
+ await self.process_and_combine_mm_data_async(base_output, self.mm_tokens)
+ )
+ image_grid_thw = self._get_grid_from_output_or_items(
+ processor_output,
+ mm_items,
+ "image_grid_thw",
+ Modality.IMAGE,
+ )
+ video_grid_thw = self._get_grid_from_output_or_items(
+ processor_output,
+ mm_items,
+ "video_grid_thw",
+ Modality.VIDEO,
+ )
+ return self._make_processor_output(
+ input_ids=input_ids,
+ mm_items=mm_items,
+ image_grid_thw=image_grid_thw,
+ video_grid_thw=video_grid_thw,
+ processor_output=processor_output,
+ )
+
+ def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
+ super().__init__(hf_config, server_args, _processor, *args, **kwargs)
+
+ self.IM_TOKEN_ID = hf_config.image_token_id
+ self.IMAGE_TOKEN_ID = hf_config.image_token_id
+ self.VIDEO_TOKEN_ID = hf_config.video_token_id
+ self.IM_START_TOKEN_ID = hf_config.vision_start_token_id
+ self.IM_END_TOKEN_ID = hf_config.vision_end_token_id
+ self.vision_start_token_id = hf_config.vision_start_token_id
+ self.model_type = hf_config.model_type
+
+ self.patch_size = hf_config.vision_config.patch_size
+ self._spatial_merge_size = hf_config.projector_config.spatial_merge_size
+ self.temporal_patch_size = 1
+
+ image_token = self._tokenizer.convert_ids_to_tokens([self.IMAGE_TOKEN_ID])[0]
+ video_token = self._tokenizer.convert_ids_to_tokens([self.VIDEO_TOKEN_ID])[0]
+ 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(self._processor)
+
+ self.ATTR_NAME_TO_MODALITY["pixel_attention_mask"] = Modality.IMAGE
+ self.ATTR_NAME_TO_MODALITY["spatial_shapes"] = Modality.IMAGE
+ self.ATTR_NAME_TO_MODALITY["pixel_attention_mask_videos"] = Modality.VIDEO
+ self.ATTR_NAME_TO_MODALITY["spatial_shapes_videos"] = Modality.VIDEO
+
+ @property
+ def spatial_merge_size(self):
+ return self._spatial_merge_size
+
+ def _tokenize_prompt(self, prompt: Union[str, list[int]]) -> list[int]:
+ if isinstance(prompt, list):
+ return list(prompt)
+ add_special_tokens = True
+ bos = getattr(self._tokenizer, "bos_token", None)
+ if self._tokenizer_auto_adds_specials and bos and prompt.startswith(bos):
+ add_special_tokens = False
+ return self._tokenizer.encode(prompt, add_special_tokens=add_special_tokens)
+
+ def _size_limits(
+ self, config: dict, default_min_pixels: int, default_max_pixels: int
+ ) -> tuple[int, int]:
+ size_value = config.get("size", {})
+ size = size_value if isinstance(size_value, dict) else {}
+ min_pixels = config.get(
+ "min_pixels",
+ config.get("shortest_edge", size.get("shortest_edge", default_min_pixels)),
+ )
+ max_pixels = config.get(
+ "max_pixels",
+ config.get("longest_edge", size.get("longest_edge", default_max_pixels)),
+ )
+ return int(min_pixels), int(max_pixels)
+
+ def _preprocess_pil_image(
+ self,
+ image: Image.Image,
+ *,
+ min_pixels: int,
+ max_pixels: int,
+ resized_size: Optional[tuple[int, int]] = None,
+ ) -> tuple[torch.Tensor, tuple[int, int]]:
+ factor = self.patch_size * self.spatial_merge_size
+ if resized_size is None:
+ resized_height, resized_width = _smart_resize(
+ image.height,
+ image.width,
+ factor=factor,
+ min_pixels=min_pixels,
+ max_pixels=max_pixels,
+ )
+ else:
+ resized_height, resized_width = resized_size
+
+ if image.size != (resized_width, resized_height):
+ image = image.resize(
+ (resized_width, resized_height), Image.Resampling.BICUBIC
+ )
+
+ array = np.asarray(image, dtype=np.float32) / 255.0
+ array = (array - 0.5) / 0.5
+
+ patch_size = self.patch_size
+ grid_h = resized_height // patch_size
+ grid_w = resized_width // patch_size
+ patches = array.reshape(grid_h, patch_size, grid_w, patch_size, 3)
+ patches = patches.transpose(0, 2, 1, 3, 4).reshape(grid_h * grid_w, -1)
+ return torch.from_numpy(patches), (grid_h, grid_w)
+
+ def _preprocess_image_item(self, image: Any) -> MultimodalDataItem:
+ min_pixels, max_pixels = self._size_limits(
+ self.image_config, IMAGE_MIN_PIXELS, IMAGE_MAX_PIXELS
+ )
+ patches, (grid_h, grid_w) = self._preprocess_pil_image(
+ _as_pil_image(image), min_pixels=min_pixels, max_pixels=max_pixels
+ )
+ spatial_shapes = torch.tensor([[grid_h, grid_w]], dtype=torch.long)
+ image_grid_thw = torch.tensor([[1, grid_h, grid_w]], dtype=torch.long)
+ return MultimodalDataItem(
+ modality=Modality.IMAGE,
+ feature=patches,
+ model_specific_data={
+ "spatial_shapes": spatial_shapes,
+ "image_grid_thw": image_grid_thw,
+ },
+ )
+
+ def _select_frame_indices(self, total_frames: int, video_fps: float) -> list[int]:
+ if total_frames <= 0:
+ raise ValueError("Video must contain at least one frame.")
+
+ has_num_frames = "num_frames" in self.video_config
+ has_nframes = "nframes" in self.video_config
+ if has_num_frames and has_nframes:
+ raise ValueError("Specify only one of num_frames and nframes")
+
+ explicit_num_frames = self.video_config.get(
+ "num_frames", self.video_config.get("nframes")
+ )
+ if explicit_num_frames is not None:
+ if "fps" in self.video_config:
+ raise ValueError("Specify only one of num_frames/nframes and fps")
+ nframes = int(explicit_num_frames)
+ else:
+ fps = float(self.video_config.get("fps", DEFAULT_TARGET_VIDEO_FPS))
+ source_fps = video_fps if video_fps > 0 else DEFAULT_SOURCE_VIDEO_FPS
+ nframes = int(total_frames / source_fps * fps)
+
+ min_frames = int(self.video_config.get("min_frames", DEFAULT_MIN_FRAMES))
+ max_frames = int(self.video_config.get("max_frames", DEFAULT_MAX_FRAMES))
+ if min_frames <= 0 or max_frames < min_frames:
+ raise ValueError(
+ "Video frame limits must satisfy 0 < min_frames <= max_frames"
+ )
+ nframes = max(min_frames, min(max_frames, nframes))
+
+ nframes = max(1, min(total_frames, nframes))
+ if nframes == total_frames:
+ return list(range(total_frames))
+ return (
+ np.linspace(0, total_frames - 1, num=nframes)
+ .round()
+ .astype(np.int64)
+ .tolist()
+ )
+
+ def _timestamps_from_indices(
+ self, frame_indices: list[int], video_fps: float
+ ) -> list[float]:
+ if video_fps > 0:
+ return [float(idx) / video_fps for idx in frame_indices]
+ return [float(idx) for idx in range(len(frame_indices))]
+
+ def _default_sampled_video_fps(self) -> float:
+ fps = float(self.video_config.get("fps", DEFAULT_TARGET_VIDEO_FPS))
+ return fps if fps > 0 else 0.0
+
+ def _coerce_video_frames(self, video: Any) -> list[Image.Image]:
+ if isinstance(video, torch.Tensor):
+ video = video.detach().cpu()
+ if video.ndim == 3:
+ video = video.unsqueeze(0)
+ if video.ndim != 4:
+ raise ValueError(
+ f"Expected video tensor with 4 dimensions, got {video.shape}."
+ )
+ if video.shape[-1] in (1, 3, 4):
+ frames = [video[i] for i in range(video.shape[0])]
+ elif video.shape[1] in (1, 3, 4):
+ frames = [video[i].permute(1, 2, 0) for i in range(video.shape[0])]
+ else:
+ raise ValueError(
+ f"Cannot infer video channel dimension from {video.shape}."
+ )
+ return [_as_pil_image(frame) for frame in frames]
+
+ if isinstance(video, np.ndarray):
+ if video.ndim == 3:
+ video = video[None, ...]
+ if video.ndim != 4:
+ raise ValueError(
+ f"Expected video array with 4 dimensions, got {video.shape}."
+ )
+ return [_as_pil_image(frame) for frame in video]
+
+ if isinstance(video, (list, tuple)):
+ return [_as_pil_image(frame) for frame in video]
+
+ raise ValueError(f"Unsupported video input type: {type(video)!r}.")
+
+ def _video_to_frames_and_timestamps(
+ self, video: Any
+ ) -> tuple[list[Image.Image], list[float]]:
+ metadata = None
+ if isinstance(video, tuple) and len(video) == 2 and isinstance(video[1], dict):
+ video, metadata = video
+
+ if isinstance(video, VideoDecoderWrapper):
+ fps = float(video.avg_fps or DEFAULT_SOURCE_VIDEO_FPS)
+ indices = self._select_frame_indices(len(video), fps)
+ frames = video.get_frames_as_tensor(indices)
+ frame_images = [_as_pil_image(frame) for frame in frames]
+ timestamps = self._timestamps_from_indices(indices, fps)
+ return frame_images, timestamps
+
+ frames = self._coerce_video_frames(video)
+ fps = self._default_sampled_video_fps()
+ frame_indices = list(range(len(frames)))
+ if metadata is not None:
+ fps = float(metadata.get("fps", fps) or 0.0)
+ metadata_indices = metadata.get("frames_indices")
+ if metadata_indices is not None:
+ metadata_indices = np.asarray(metadata_indices).reshape(-1).tolist()
+ if len(metadata_indices) == len(frames):
+ frame_indices = [int(idx) for idx in metadata_indices]
+ return frames, self._timestamps_from_indices(frame_indices, fps)
+
+ def _preprocess_video_item(self, video: Any) -> MultimodalDataItem:
+ frames, timestamps = self._video_to_frames_and_timestamps(video)
+ total_min_pixels, total_max_pixels = self._size_limits(
+ self.video_config, VIDEO_MIN_PIXELS, VIDEO_TOTAL_PIXELS
+ )
+ first_frame = frames[0]
+ factor = self.patch_size * self.spatial_merge_size
+ resized_size = _smart_resize(
+ first_frame.height,
+ first_frame.width,
+ factor=factor,
+ min_pixels=total_min_pixels,
+ max_pixels=total_max_pixels,
+ num_frames=len(frames),
+ )
+
+ frame_patches = []
+ spatial_shapes = []
+ for frame in frames:
+ patches, (grid_h, grid_w) = self._preprocess_pil_image(
+ frame,
+ min_pixels=total_min_pixels,
+ max_pixels=total_max_pixels,
+ resized_size=resized_size,
+ )
+ frame_patches.append(patches)
+ spatial_shapes.append([grid_h, grid_w])
+
+ grid_h, grid_w = spatial_shapes[0]
+ feature = torch.cat(frame_patches, dim=0)
+ spatial_shapes_tensor = torch.tensor(spatial_shapes, dtype=torch.long)
+ video_grid_thw = torch.tensor([[len(frames), grid_h, grid_w]], dtype=torch.long)
+ return MultimodalDataItem(
+ modality=Modality.VIDEO,
+ feature=feature,
+ model_specific_data={
+ "spatial_shapes": spatial_shapes_tensor,
+ "video_grid_thw": video_grid_thw,
+ "timestamps": timestamps,
+ },
+ )
+
+ def _timestamp_token_ids(self, timestamp: float) -> list[int]:
+ return self._tokenizer.encode(
+ f"<{timestamp:.1f} seconds>", add_special_tokens=False
+ )
+
+ def _build_input_ids(
+ self,
+ prompt: Union[str, list[int]],
+ img_grid_thw: Optional[torch.Tensor],
+ video_grid_thw: Optional[torch.Tensor],
+ video_timestamps: Optional[list[list[float]]],
+ ):
+ if not isinstance(prompt, list):
+ prompt = self._tokenize_prompt(prompt)
+
+ input_ids = []
+ offsets = []
+ modality_list = []
+ cur_idx = 0
+ spatial_merge_size = self.spatial_merge_size
+
+ vision_start_indices = []
+ for i in range(len(prompt) - 1):
+ if prompt[i + 1] == self.IMAGE_TOKEN_ID:
+ vision_start_indices.append((i, Modality.IMAGE))
+ elif prompt[i + 1] == self.VIDEO_TOKEN_ID:
+ vision_start_indices.append((i, Modality.VIDEO))
+
+ img_idx = 0
+ video_idx = 0
+ for mm_start_idx, modality in vision_start_indices:
+ if modality == Modality.IMAGE:
+ if img_grid_thw is None:
+ raise ValueError(
+ "Missing image grid metadata for image placeholder."
+ )
+ mm_token_num = int(
+ img_grid_thw[img_idx].prod().item() // (spatial_merge_size**2)
+ )
+ assert cur_idx <= mm_start_idx
+ input_ids.extend(prompt[cur_idx : mm_start_idx + 1])
+ mm_offset_start = len(input_ids)
+ input_ids.extend([self.IMAGE_TOKEN_ID] * mm_token_num)
+ offsets.append((mm_offset_start, len(input_ids) - 1))
+ modality_list.append(Modality.IMAGE)
+ cur_idx = mm_start_idx + 2
+ img_idx += 1
+ continue
+
+ if video_grid_thw is None:
+ raise ValueError("Missing video grid metadata for video placeholder.")
+ num_frames = int(video_grid_thw[video_idx][0].item())
+ tokens_per_frame = int(
+ video_grid_thw[video_idx][1:].prod().item() // (spatial_merge_size**2)
+ )
+ timestamps = (
+ video_timestamps[video_idx]
+ if video_timestamps is not None
+ else [float(i) for i in range(num_frames)]
+ )
+ if len(timestamps) != num_frames:
+ raise ValueError(
+ "Cosmos3-Edge video timestamps must match video frame count: "
+ f"got {len(timestamps)} vs {num_frames}."
+ )
+
+ has_start = prompt[mm_start_idx] == self.IM_START_TOKEN_ID
+ has_end = (
+ mm_start_idx + 2 < len(prompt)
+ and prompt[mm_start_idx + 2] == self.IM_END_TOKEN_ID
+ )
+ target_start = mm_start_idx if has_start else mm_start_idx + 1
+ target_end = mm_start_idx + 2 if has_start and has_end else mm_start_idx + 1
+ assert cur_idx <= target_start
+
+ input_ids.extend(prompt[cur_idx:target_start])
+ frame_offsets = []
+ for timestamp in timestamps:
+ input_ids.extend(self._timestamp_token_ids(float(timestamp)))
+ input_ids.append(self.IM_START_TOKEN_ID)
+ mm_offset_start = len(input_ids)
+ input_ids.extend([self.VIDEO_TOKEN_ID] * tokens_per_frame)
+ frame_offsets.append((mm_offset_start, len(input_ids) - 1))
+ input_ids.append(self.IM_END_TOKEN_ID)
+
+ offsets.append(frame_offsets)
+ modality_list.append(Modality.VIDEO)
+ cur_idx = target_end + 1
+ video_idx += 1
+ else:
+ input_ids.extend(prompt[cur_idx:])
+
+ return input_ids, offsets, modality_list
+
+ def _assign_offsets(
+ self,
+ modality_list: list[Modality],
+ offsets: list,
+ image_items: list[MultimodalDataItem],
+ video_items: list[MultimodalDataItem],
+ ) -> list[MultimodalDataItem]:
+ image_idx = 0
+ video_idx = 0
+ mm_items = []
+ for modality, offset in zip(modality_list, offsets):
+ if modality == Modality.IMAGE:
+ item = image_items[image_idx]
+ image_idx += 1
+ elif modality == Modality.VIDEO:
+ item = video_items[video_idx]
+ video_idx += 1
+ else:
+ continue
+ item.offsets = offset if isinstance(offset, list) else [offset]
+ mm_items.append(item)
+
+ if image_idx != len(image_items) or video_idx != len(video_items):
+ raise ValueError(
+ "Cosmos3-Edge prompt media placeholders do not match provided media."
+ )
+ return mm_items
+
+ async def process_mm_data_async(
+ self,
+ image_data: List[Union[str, bytes]],
+ audio_data,
+ input_text: str,
+ request_obj,
+ **kwargs,
+ ):
+ video_data = getattr(request_obj, "video_data", None)
+ if video_data is not None and not isinstance(video_data, list):
+ video_data = [video_data]
+
+ if not image_data and not video_data:
+ input_ids = self._tokenize_prompt(input_text)
+ return MultimodalProcessorOutput(
+ input_ids=input_ids,
+ mm_items=[],
+ im_start_id=self.IM_START_TOKEN_ID,
+ im_end_id=self.IM_END_TOKEN_ID,
+ im_token_id=self.IMAGE_TOKEN_ID,
+ video_token_id=self.VIDEO_TOKEN_ID,
+ )
+
+ base_output = await self.load_mm_data(
+ prompt=input_text,
+ image_data=image_data,
+ video_data=video_data,
+ multimodal_tokens=self.mm_tokens,
+ )
+
+ if self._all_mm_data_is_preprocessed(base_output.images, base_output.videos):
+ return await self._process_preprocessed_mm_data(base_output)
+
+ if any(
+ self._is_preprocessed_input(item)
+ for item in [*base_output.images, *base_output.videos]
+ ):
+ raise ValueError(
+ "Cosmos3-Edge does not support mixing raw and preprocessed media "
+ "in the same request."
+ )
+
+ image_items = [
+ self._preprocess_image_item(image) for image in base_output.images
+ ]
+ video_items = [
+ self._preprocess_video_item(video) for video in base_output.videos
+ ]
+
+ image_grid_thw = (
+ torch.cat([item.image_grid_thw for item in image_items], dim=0)
+ if image_items
+ else None
+ )
+ video_grid_thw = (
+ torch.cat([item.video_grid_thw for item in video_items], dim=0)
+ if video_items
+ else None
+ )
+
+ prompt_ids = base_output.input_ids
+ if prompt_ids is None:
+ prompt_ids = self._tokenize_prompt(base_output.input_text)
+ video_timestamps = (
+ [item.timestamps for item in video_items] if video_items else None
+ )
+ input_ids, offsets, modality_list = self._build_input_ids(
+ prompt_ids,
+ img_grid_thw=image_grid_thw,
+ video_grid_thw=video_grid_thw,
+ video_timestamps=video_timestamps,
+ )
+ mm_items = self._assign_offsets(
+ modality_list, offsets, image_items=image_items, video_items=video_items
+ )
+
+ return self._make_processor_output(
+ input_ids=input_ids,
+ mm_items=mm_items,
+ image_grid_thw=image_grid_thw,
+ video_grid_thw=video_grid_thw,
+ )
diff --git a/python/sglang/srt/multimodal/processors/qwen_vl.py b/python/sglang/srt/multimodal/processors/qwen_vl.py
index 875d42de6..ebe693b5f 100644
--- a/python/sglang/srt/multimodal/processors/qwen_vl.py
+++ b/python/sglang/srt/multimodal/processors/qwen_vl.py
@@ -17,6 +17,7 @@ from sglang.srt.managers.schedule_batch import (
MultimodalDataItem,
MultimodalProcessorOutput,
)
+from sglang.srt.models.cosmos3 import Cosmos3ForConditionalGeneration
from sglang.srt.models.interns2_mobius import (
InternS2MobiusForConditionalGeneration,
)
@@ -298,6 +299,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
InternS2PreviewForConditionalGeneration,
InternS2MobiusForConditionalGeneration,
Qwen3OmniMoeForConditionalGeneration,
+ Cosmos3ForConditionalGeneration,
]
def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
@@ -521,6 +523,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
"qwen3_5_moe",
"intern_s2_preview",
"interns2_mobius",
+ "cosmos3_omni",
):
return None
@@ -656,6 +659,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
"qwen3_5",
"qwen3_5_moe",
"intern_s2_preview",
+ "cosmos3_omni",
]
and video_timestamps is not None
):
@@ -765,6 +769,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
"qwen3_5_moe",
"intern_s2_preview",
"interns2_mobius",
+ "cosmos3_omni",
):
processor_kwargs.update(
video_metadata=video_metadata,
diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py
index 1b21c9cad..0e9597b6d 100644
--- a/python/sglang/srt/server_args.py
+++ b/python/sglang/srt/server_args.py
@@ -3791,7 +3791,11 @@ class ServerArgs:
# ===== END TO BE REFACTORED ====
- LANGUAGE_MODEL_ONLY_ARCHITECTURES = ("MuseGlimmerForConditionalGeneration",)
+ LANGUAGE_MODEL_ONLY_ARCHITECTURES = (
+ "MuseGlimmerForConditionalGeneration",
+ "Cosmos3ForConditionalGeneration",
+ "Cosmos3EdgeForConditionalGeneration",
+ )
# The attention-backend allow-list is enforced via
# --enable-page-major-kv-layout (implied by the unified pool in
diff --git a/python/sglang/srt/utils/hf_transformers/common.py b/python/sglang/srt/utils/hf_transformers/common.py
index fd101dabb..26d141ce2 100644
--- a/python/sglang/srt/utils/hf_transformers/common.py
+++ b/python/sglang/srt/utils/hf_transformers/common.py
@@ -25,6 +25,11 @@ from sglang.srt.configs import (
AfmoeConfig,
BailingHybridConfig,
ChatGLMConfig,
+ Cosmos3Config,
+ Cosmos3EdgeConfig,
+ Cosmos3EdgeProjectorConfig,
+ Cosmos3EdgeTextConfig,
+ Cosmos3EdgeVisionConfig,
DbrxConfig,
DeepseekVL2Config,
Dots3Config,
@@ -223,6 +228,42 @@ for name, cls in _CONFIG_REGISTRY.items():
if "already registered" not in err and "already used" not in err:
logger.warning("Failed to register config %s: %s", name, e)
+# Cosmos3 (understanding tower) reuses the Qwen3-VL config schema. Register it
+# with AutoConfig only (not `_CONFIG_REGISTRY`), so the nested `text_config` is
+# flattened onto the top-level config in `get_config` — the same path the base
+# Qwen3-VL config relies on. Adding it to `_CONFIG_REGISTRY` would trigger a
+# `from_pretrained` reload that drops that flattening.
+try:
+ AutoConfig.register(Cosmos3Config.model_type, Cosmos3Config)
+except ValueError as e:
+ err = str(e).lower()
+ if "already registered" not in err and "already used" not in err:
+ logger.warning("Failed to register config %s: %s", Cosmos3Config.model_type, e)
+
+# Cosmos3-Edge native text support starts from the checkpoint root config, then
+# consumes ``text_config`` in ``sglang.srt.models.cosmos3_edge``. Keep it out of
+# `_CONFIG_REGISTRY` so the generic parser can flatten text attributes onto the
+# root config after `AutoConfig.from_pretrained`, matching other multimodal
+# configs that use a text sub-config.
+for _cosmos3_edge_config_cls in (
+ Cosmos3EdgeTextConfig,
+ Cosmos3EdgeVisionConfig,
+ Cosmos3EdgeProjectorConfig,
+ Cosmos3EdgeConfig,
+):
+ try:
+ AutoConfig.register(
+ _cosmos3_edge_config_cls.model_type, _cosmos3_edge_config_cls
+ )
+ except ValueError as e:
+ err = str(e).lower()
+ if "already registered" not in err and "already used" not in err:
+ logger.warning(
+ "Failed to register config %s: %s",
+ _cosmos3_edge_config_cls.model_type,
+ e,
+ )
+
# ---------------------------------------------------------------------------
# Download / path helpers
diff --git a/test/registered/unit/model_loader/test_modelopt_loader.py b/test/registered/unit/model_loader/test_modelopt_loader.py
index 613960f5d..48a9eddbf 100644
--- a/test/registered/unit/model_loader/test_modelopt_loader.py
+++ b/test/registered/unit/model_loader/test_modelopt_loader.py
@@ -35,7 +35,10 @@ from sglang.srt.model_loader.loader import (
ModelOptModelLoader,
get_model_loader,
)
-from sglang.srt.model_loader.weight_utils import get_quant_config
+from sglang.srt.model_loader.weight_utils import (
+ _modelopt_quant_section,
+ get_quant_config,
+)
from sglang.srt.models.minimax_m3 import MiniMaxM3SparseForCausalLM
from sglang.srt.models.utils import WeightsMapper
from sglang.srt.utils import get_device
@@ -925,6 +928,50 @@ class TestModelOptMixedPrecisionConfig(CustomTestCase):
self.assertEqual(result["quant_method"], "modelopt_mixed")
+ def test_flat_hf_quant_config_without_quantization_key(self):
+ """Diffusion/unified ModelOpt exports use a flat hf_quant_config.json.
+
+ Regression for Cosmos3-style checkpoints that put quant_algo at the top
+ level (no nested ``quantization`` key).
+ """
+ model_config = ModelConfig.__new__(ModelConfig)
+
+ result = model_config._parse_modelopt_quant_config(
+ {
+ "quant_method": "modelopt",
+ "quant_algo": "FP8",
+ "quant_type": "FP8_FP8",
+ "ignore": ["lm_head", "visual*"],
+ }
+ )
+
+ self.assertEqual(result["quant_method"], "modelopt_fp8")
+ self.assertEqual(result["quant_algo"], "FP8")
+
+ def test_hf_quant_config_missing_quant_algo_returns_none(self):
+ model_config = ModelConfig.__new__(ModelConfig)
+ self.assertIsNone(
+ model_config._parse_modelopt_quant_config(
+ {"quant_method": "modelopt", "producer": {"name": "modelopt"}}
+ )
+ )
+
+ def test_modelopt_quant_section_supports_nested_and_flat(self):
+ nested = {"quantization": {"quant_algo": "FP8", "exclude_modules": ["lm_head"]}}
+ self.assertEqual(
+ _modelopt_quant_section(nested)["quant_algo"],
+ "FP8",
+ )
+
+ flat = {
+ "quant_method": "modelopt",
+ "quant_algo": "FP8",
+ "ignore": ["lm_head"],
+ "producer": {"name": "modelopt"},
+ }
+ self.assertIs(_modelopt_quant_section(flat), flat)
+ self.assertEqual(_modelopt_quant_section(flat)["quant_algo"], "FP8")
+
def test_mixed_precision_override_does_not_hijack_w4afp8(self):
self.assertIsNone(
ModelOptMixedPrecisionConfig.override_quantization_method(
diff --git a/test/registered/unit/model_loader/test_prefetch_checkpoints.py b/test/registered/unit/model_loader/test_prefetch_checkpoints.py
index dda326265..63bd66827 100644
--- a/test/registered/unit/model_loader/test_prefetch_checkpoints.py
+++ b/test/registered/unit/model_loader/test_prefetch_checkpoints.py
@@ -379,13 +379,12 @@ class TestPrefetchDispatch(CustomTestCase):
return DefaultModelLoader(load_config)
def _make_source(self):
- # model_config=None skips maybe_add_mtp_safetensors.
- return SimpleNamespace(
+ # model_config=None skips maybe_add_mtp_safetensors. A real Source
+ # (not a stand-in) so new fields with defaults are picked up.
+ return DefaultModelLoader.Source(
model_or_path="/dummy",
revision=None,
fall_back_to_pt=False,
- model_config=None,
- prefix="",
)
def _server_args(self, prefetch, disable_mmap=False, drop_cache=False):
diff --git a/test/registered/unit/model_loader/test_weight_utils.py b/test/registered/unit/model_loader/test_weight_utils.py
index 0e4ba5ed0..443f32da1 100644
--- a/test/registered/unit/model_loader/test_weight_utils.py
+++ b/test/registered/unit/model_loader/test_weight_utils.py
@@ -21,6 +21,7 @@ def _write_index(folder, weight_map):
def _touch(folder, name):
path = os.path.join(folder, name)
+ os.makedirs(os.path.dirname(path), exist_ok=True)
open(path, "w").close()
return path
@@ -72,6 +73,52 @@ class TestFilterDuplicateSafetensorsFiles(CustomTestCase):
)
self.assertEqual(sorted(result), sorted([shard1, shard2]))
+ def test_missing_shard_outside_allow_patterns_is_ignored(self):
+ # Cosmos3-style checkpoints use one root index for multiple subfolder
+ # weight sources. Loading the transformer source should not require the
+ # vision encoder shard to already be present; the secondary source
+ # downloads and loads it separately.
+ _write_index(
+ self.folder,
+ {
+ "llm": "transformer/diffusion_pytorch_model.safetensors",
+ "vit": "vision_encoder/model.safetensors",
+ },
+ )
+ transformer = _touch(
+ self.folder, "transformer/diffusion_pytorch_model.safetensors"
+ )
+
+ result = filter_duplicate_safetensors_files(
+ hf_weights_files=[transformer],
+ hf_folder=self.folder,
+ index_file=INDEX_NAME,
+ allow_patterns=["transformer/*.safetensors"],
+ )
+ self.assertEqual(result, [transformer])
+
+ def test_missing_shard_inside_allow_patterns_raises(self):
+ _write_index(
+ self.folder,
+ {
+ "llm1": "transformer/model-00001-of-00002.safetensors",
+ "llm2": "transformer/model-00002-of-00002.safetensors",
+ "vit": "vision_encoder/model.safetensors",
+ },
+ )
+ transformer = _touch(
+ self.folder, "transformer/model-00001-of-00002.safetensors"
+ )
+
+ with self.assertRaises(RuntimeError) as cm:
+ filter_duplicate_safetensors_files(
+ hf_weights_files=[transformer],
+ hf_folder=self.folder,
+ index_file=INDEX_NAME,
+ allow_patterns=["transformer/*.safetensors"],
+ )
+ self.assertIn("model-00002-of-00002.safetensors", str(cm.exception))
+
def test_single_file_model_no_index_returns_unchanged(self):
# No index on disk (single-file / dummy / object-storage): early return.
single = _touch(self.folder, "model.safetensors")
diff --git a/test/registered/unit/models/test_cosmos3.py b/test/registered/unit/models/test_cosmos3.py
new file mode 100644
index 000000000..4d91de871
--- /dev/null
+++ b/test/registered/unit/models/test_cosmos3.py
@@ -0,0 +1,310 @@
+"""Unit tests for the Cosmos3 reasoner (understanding tower).
+
+Covers the three pieces of load-time logic that make the diffusers-layout
+Cosmos3 checkpoint loadable through the Qwen3-VL inference stack:
+
+1. ``Cosmos3ForConditionalGeneration.hf_to_sglang_mapper`` - renames the
+ understanding-tower keys into the nested Qwen3-VL checkpoint form and drops
+ the generation-tower weights.
+2. ``Cosmos3Config`` - reuses the Qwen3-VL schema under ``model_type
+ "cosmos3_omni"`` and is registered with ``AutoConfig``.
+3. ``DefaultModelLoader`` ``allow_patterns_overrides`` - globs weights from the
+ ``transformer/`` and ``vision_encoder/`` subfolders rather than the repo root.
+
+All of this is pure CPU logic (no server / engine launch).
+"""
+
+from sglang.test.ci.ci_register import register_cpu_ci
+
+register_cpu_ci(est_time=15, suite="base-a-test-cpu")
+
+import os
+import tempfile
+import unittest
+
+import safetensors.torch
+import torch
+
+from sglang.srt.configs import Cosmos3Config
+from sglang.srt.configs.load_config import LoadConfig, LoadFormat
+from sglang.srt.model_loader.loader import DefaultModelLoader
+from sglang.srt.models.cosmos3 import Cosmos3ForConditionalGeneration
+from sglang.srt.runtime_context import get_context
+from sglang.test.test_utils import CustomTestCase
+
+
+class TestCosmos3WeightsMapper(CustomTestCase):
+ """Validate the understanding-tower rename + generation-tower drop rules."""
+
+ def setUp(self):
+ self.mapper = Cosmos3ForConditionalGeneration.hf_to_sglang_mapper
+
+ def test_understanding_tower_is_renamed(self):
+ # Flat Cosmos3 keys -> nested Qwen3-VL checkpoint keys. The parent
+ # Qwen3VLForConditionalGeneration.load_weights then strips the
+ # `model.language_model.`/`model.visual.` prefixes and fuses q/k/v.
+ inputs = [
+ "layers.0.self_attn.to_q.weight",
+ "layers.0.self_attn.to_k.weight",
+ "layers.0.self_attn.to_v.weight",
+ "layers.0.self_attn.to_out.weight",
+ "layers.0.self_attn.norm_q.weight",
+ "layers.0.self_attn.norm_k.weight",
+ "layers.0.mlp.gate_proj.weight",
+ "layers.0.mlp.up_proj.weight",
+ "layers.0.mlp.down_proj.weight",
+ "layers.0.input_layernorm.weight",
+ "embed_tokens.weight",
+ "norm.weight",
+ "lm_head.weight",
+ ]
+ expected = [
+ "model.language_model.layers.0.self_attn.q_proj.weight",
+ "model.language_model.layers.0.self_attn.k_proj.weight",
+ "model.language_model.layers.0.self_attn.v_proj.weight",
+ "model.language_model.layers.0.self_attn.o_proj.weight",
+ "model.language_model.layers.0.self_attn.q_norm.weight",
+ "model.language_model.layers.0.self_attn.k_norm.weight",
+ "model.language_model.layers.0.mlp.gate_proj.weight",
+ "model.language_model.layers.0.mlp.up_proj.weight",
+ "model.language_model.layers.0.mlp.down_proj.weight",
+ "model.language_model.layers.0.input_layernorm.weight",
+ "model.language_model.embed_tokens.weight",
+ "model.language_model.norm.weight",
+ "lm_head.weight",
+ ]
+ self.assertEqual(self.mapper.apply_list(inputs), expected)
+
+ def test_vision_encoder_is_prefixed(self):
+ inputs = [
+ "blocks.0.attn.qkv.weight",
+ "merger.norm.weight",
+ "patch_embed.proj.weight",
+ "pos_embed.weight",
+ "deepstack_merger_list.0.norm.weight",
+ ]
+ expected = [
+ "model.visual.blocks.0.attn.qkv.weight",
+ "model.visual.merger.norm.weight",
+ "model.visual.patch_embed.proj.weight",
+ "model.visual.pos_embed.weight",
+ "model.visual.deepstack_merger_list.0.norm.weight",
+ ]
+ self.assertEqual(self.mapper.apply_list(inputs), expected)
+
+ def test_generation_tower_is_dropped(self):
+ dropped = [
+ "layers.0.self_attn.add_q_proj.weight",
+ "layers.0.self_attn.add_k_proj.weight",
+ "layers.0.self_attn.add_v_proj.weight",
+ "layers.0.self_attn.to_add_out.weight",
+ "layers.0.self_attn.norm_added_q.weight",
+ "layers.0.self_attn.norm_added_k.weight",
+ "layers.0.self_attn.q_proj_moe_gen.weight",
+ "layers.0.mlp_moe_gen.gate_up_proj.weight",
+ "norm_moe_gen.weight",
+ "proj_in.weight",
+ "proj_out.weight",
+ "time_embedder.linear_1.weight",
+ "audio_proj_in.weight",
+ "audio_proj_out.weight",
+ "action_proj_in.weight",
+ "action_proj_out.weight",
+ "audio_modality_embed",
+ "action_modality_embed",
+ ]
+ self.assertEqual(self.mapper.apply_list(dropped), [])
+
+ def test_modelopt_quantizer_buffers_are_dropped(self):
+ # FP8 Cosmos3 exports keep ModelOpt TensorQuantizer state alongside the
+ # inference weight_scale/input_scale tensors. SGLang uses the latter;
+ # transformers restores the former via ModelOpt HF checkpointing.
+ dropped = [
+ "layers.0.mlp.gate_proj.input_quantizer._amax",
+ "layers.0.mlp.gate_proj.weight_quantizer._amax",
+ "layers.0.mlp.gate_proj.weight_quantizer._scale",
+ "layers.0.self_attn.to_q.input_quantizer._amax",
+ "layers.0.self_attn.to_q.weight_quantizer._amax",
+ "layers.0.self_attn.to_q.weight_quantizer._scale",
+ ]
+ self.assertEqual(self.mapper.apply_list(dropped), [])
+
+ def test_fp8_inference_scales_are_kept(self):
+ inputs = [
+ "layers.0.mlp.gate_proj.weight",
+ "layers.0.mlp.gate_proj.weight_scale",
+ "layers.0.mlp.gate_proj.input_scale",
+ "layers.0.self_attn.to_q.weight",
+ "layers.0.self_attn.to_q.weight_scale",
+ "layers.0.self_attn.to_q.input_scale",
+ ]
+ expected = [
+ "model.language_model.layers.0.mlp.gate_proj.weight",
+ "model.language_model.layers.0.mlp.gate_proj.weight_scale",
+ "model.language_model.layers.0.mlp.gate_proj.input_scale",
+ "model.language_model.layers.0.self_attn.q_proj.weight",
+ "model.language_model.layers.0.self_attn.q_proj.weight_scale",
+ "model.language_model.layers.0.self_attn.q_proj.input_scale",
+ ]
+ self.assertEqual(self.mapper.apply_list(inputs), expected)
+
+ def test_moe_gen_substring_wins_over_norm_prefix(self):
+ # `norm_moe_gen.weight` must be dropped (generation), not routed to the
+ # final `norm.` -> language-model norm.
+ self.assertEqual(self.mapper.apply_list(["norm_moe_gen.weight"]), [])
+ self.assertEqual(
+ self.mapper.apply_list(["norm.weight"]),
+ ["model.language_model.norm.weight"],
+ )
+
+
+class TestCosmos3Config(CustomTestCase):
+ def test_model_type(self):
+ self.assertEqual(Cosmos3Config.model_type, "cosmos3_omni")
+
+ def test_subconfigs_are_objects(self):
+ # The Qwen3-VL inference stack reads sub-configs as objects, e.g.
+ # `config.vision_config.hidden_size` and
+ # `config.vision_config.deepstack_visual_indexes`. Some transformers
+ # versions leave sub-configs as raw dicts after construction, which would
+ # raise `'dict' object has no attribute 'hidden_size'` at model init.
+ # Cosmos3Config coerces them into config objects, so assert that here.
+ cfg = Cosmos3Config(
+ text_config={
+ "hidden_size": 128,
+ "num_hidden_layers": 2,
+ "num_attention_heads": 4,
+ },
+ vision_config={"depth": 3, "hidden_size": 64},
+ )
+ self.assertNotIsInstance(cfg.text_config, dict)
+ self.assertNotIsInstance(cfg.vision_config, dict)
+ self.assertEqual(cfg.text_config.hidden_size, 128)
+ self.assertEqual(cfg.text_config.num_hidden_layers, 2)
+ self.assertEqual(cfg.text_config.num_attention_heads, 4)
+ self.assertEqual(cfg.vision_config.hidden_size, 64)
+ self.assertEqual(cfg.vision_config.depth, 3)
+
+ def test_registered_with_autoconfig(self):
+ # Importing common runs the AutoConfig registration side effects.
+ # Newer transformers ship a native Cosmos3OmniConfig owning this model
+ # type, and sglang's registration deliberately yields to it — so assert
+ # the contract the Qwen3-VL stack needs from whichever class resolves,
+ # not one implementation class.
+ from transformers import AutoConfig
+
+ import sglang.srt.utils.hf_transformers.common # noqa: F401
+
+ cfg = AutoConfig.for_model(
+ "cosmos3_omni",
+ text_config={"hidden_size": 128, "num_hidden_layers": 2},
+ vision_config={"hidden_size": 64},
+ )
+ self.assertEqual(cfg.model_type, "cosmos3_omni")
+ # Sub-configs must resolve to attribute-accessible objects, not dicts:
+ # the inference stack reads e.g. `config.vision_config.hidden_size`.
+ self.assertEqual(cfg.text_config.hidden_size, 128)
+ self.assertEqual(cfg.vision_config.hidden_size, 64)
+
+
+class TestCosmos3MropeIndex(CustomTestCase):
+ """Cosmos3 reuses the Qwen3-VL mrope path.
+
+ The multimodal processor calls ``MRotaryEmbedding.get_rope_index`` with the
+ config's ``model_type``. Because Cosmos3 declares its own ``cosmos3_omni``
+ type (for AutoConfig resolution), the mrope dispatch must recognize it as a
+ Qwen3-VL-family model or it raises ``RuntimeError: Unimplemented model type:
+ cosmos3_omni``. This guards that regression.
+ """
+
+ def _get_rope_index(self, model_type):
+ from sglang.srt.layers.rotary_embedding.mrope import MRotaryEmbedding
+
+ # A single image: grid [t=1, h=2, w=2] with spatial_merge_size 2 expands
+ # to exactly one image placeholder token after the vision-start token.
+ input_ids = torch.tensor([[1, 99, 100, 2]], dtype=torch.long)
+ image_grid_thw = torch.tensor([[1, 2, 2]], dtype=torch.long)
+ return MRotaryEmbedding.get_rope_index(
+ spatial_merge_size=2,
+ image_token_id=100,
+ video_token_id=101,
+ vision_start_token_id=99,
+ model_type=model_type,
+ input_ids=input_ids,
+ image_grid_thw=image_grid_thw,
+ )
+
+ def test_cosmos3_omni_is_supported(self):
+ # Must not raise "Unimplemented model type".
+ positions, delta = self._get_rope_index("cosmos3_omni")
+ self.assertEqual(positions.shape[0], 3)
+
+ def test_cosmos3_matches_qwen3_vl(self):
+ # Cosmos3 must produce identical rope indices to the qwen3_vl path it
+ # reuses, so behavior can't silently diverge.
+ pos_cosmos, delta_cosmos = self._get_rope_index("cosmos3_omni")
+ pos_qwen, delta_qwen = self._get_rope_index("qwen3_vl")
+ self.assertTrue(torch.equal(pos_cosmos, pos_qwen))
+ self.assertTrue(torch.equal(delta_cosmos, delta_qwen))
+
+
+class TestAllowPatternsOverrides(CustomTestCase):
+ """Validate DefaultModelLoader subfolder globbing for diffusers layouts."""
+
+ def setUp(self):
+ # Publish a default config so _prepare_weights reads real bags; the
+ # default model_checksum=None skips checksum verification.
+ override = get_context().override_server_args()
+ override.install()
+ self.addCleanup(override.restore)
+ self.loader = DefaultModelLoader(LoadConfig(load_format=LoadFormat.AUTO))
+
+ def _make_checkpoint(self, root):
+ os.makedirs(os.path.join(root, "transformer"))
+ os.makedirs(os.path.join(root, "vision_encoder"))
+ # A root-level file that must be ignored when an override is given.
+ safetensors.torch.save_file(
+ {"root": torch.zeros(2)}, os.path.join(root, "model.safetensors")
+ )
+ safetensors.torch.save_file(
+ {"llm": torch.zeros(2)},
+ os.path.join(root, "transformer", "diffusion_pytorch_model.safetensors"),
+ )
+ safetensors.torch.save_file(
+ {"vit": torch.zeros(2)},
+ os.path.join(root, "vision_encoder", "model.safetensors"),
+ )
+
+ def test_override_selects_transformer_subfolder(self):
+ with tempfile.TemporaryDirectory() as root:
+ self._make_checkpoint(root)
+ _, files, use_safetensors = self.loader._prepare_weights(
+ root, None, True, ["transformer/*.safetensors"]
+ )
+ self.assertTrue(use_safetensors)
+ self.assertEqual(len(files), 1)
+ self.assertTrue(
+ files[0].endswith("transformer/diffusion_pytorch_model.safetensors")
+ )
+
+ def test_override_selects_vision_subfolder(self):
+ with tempfile.TemporaryDirectory() as root:
+ self._make_checkpoint(root)
+ _, files, _ = self.loader._prepare_weights(
+ root, None, True, ["vision_encoder/*.safetensors"]
+ )
+ self.assertEqual(len(files), 1)
+ self.assertTrue(files[0].endswith("vision_encoder/model.safetensors"))
+
+ def test_no_override_globs_repo_root(self):
+ with tempfile.TemporaryDirectory() as root:
+ self._make_checkpoint(root)
+ _, files, _ = self.loader._prepare_weights(root, None, True)
+ # Without an override only the root-level file is discovered.
+ self.assertEqual(
+ [os.path.basename(f) for f in files], ["model.safetensors"]
+ )
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/test/registered/unit/models/test_cosmos3_edge.py b/test/registered/unit/models/test_cosmos3_edge.py
new file mode 100644
index 000000000..24cc291ac
--- /dev/null
+++ b/test/registered/unit/models/test_cosmos3_edge.py
@@ -0,0 +1,442 @@
+# Copyright 2023-2025 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.
+# ==============================================================================
+# CPU coverage for Cosmos3-Edge checkpoint mapping and video prompt preparation.
+
+from sglang.test.ci.ci_register import register_cpu_ci
+
+register_cpu_ci(est_time=10, suite="base-a-test-cpu")
+
+import asyncio
+import unittest
+
+import numpy as np
+import torch
+
+from sglang.srt.configs.cosmos3 import Cosmos3EdgeConfig
+from sglang.srt.layers.rotary_embedding import MRotaryEmbedding
+from sglang.srt.managers.schedule_batch import (
+ Modality,
+ MultimodalDataItem,
+ MultimodalInputFormat,
+)
+from sglang.srt.models.cosmos3_edge import Cosmos3EdgeForConditionalGeneration
+from sglang.srt.multimodal.processors.base_processor import (
+ BaseMultiModalProcessorOutput,
+ MultimodalSpecialTokens,
+)
+from sglang.srt.multimodal.processors.cosmos3_edge import (
+ Cosmos3EdgeProcessor,
+ _smart_resize,
+)
+from sglang.test.test_utils import CustomTestCase
+
+
+class TestCosmos3EdgeConfig(CustomTestCase):
+ def test_checkpoint_without_pad_token_has_explicit_none(self):
+ config = Cosmos3EdgeConfig(text_config={"eos_token_id": 11})
+
+ self.assertTrue(hasattr(config.text_config, "pad_token_id"))
+ self.assertIsNone(config.text_config.pad_token_id)
+ self.assertEqual(config.text_config.eos_token_id, 11)
+ self.assertFalse(config.text_config.tie_word_embeddings)
+ self.assertTrue(hasattr(config, "pad_token_id"))
+ self.assertIsNone(config.pad_token_id)
+ self.assertEqual(config.eos_token_id, 11)
+
+
+class TestCosmos3EdgeWeightsMapper(CustomTestCase):
+ def setUp(self):
+ self.mapper = Cosmos3EdgeForConditionalGeneration.hf_to_sglang_mapper
+
+ def test_text_tower_is_renamed_for_arcee(self):
+ inputs = [
+ "embed_tokens.weight",
+ "layers.0.input_layernorm.weight",
+ "layers.0.self_attn.to_q.weight",
+ "layers.0.self_attn.to_k.weight",
+ "layers.0.self_attn.to_v.weight",
+ "layers.0.self_attn.to_out.weight",
+ "layers.0.mlp.up_proj.weight",
+ "layers.0.mlp.down_proj.weight",
+ "norm.weight",
+ "lm_head.weight",
+ ]
+ expected = [
+ "model.embed_tokens.weight",
+ "model.layers.0.input_layernorm.weight",
+ "model.layers.0.self_attn.q_proj.weight",
+ "model.layers.0.self_attn.k_proj.weight",
+ "model.layers.0.self_attn.v_proj.weight",
+ "model.layers.0.self_attn.o_proj.weight",
+ "model.layers.0.mlp.up_proj.weight",
+ "model.layers.0.mlp.down_proj.weight",
+ "model.norm.weight",
+ "lm_head.weight",
+ ]
+ self.assertEqual(self.mapper.apply_list(inputs), expected)
+
+ def test_generation_and_routed_vision_weights_are_dropped(self):
+ dropped = [
+ "layers.0.self_attn.k_norm_und_for_gen.weight",
+ "layers.0.self_attn.add_q_proj.weight",
+ "layers.0.self_attn.to_add_out.weight",
+ "layers.0.mlp_moe_gen.up_proj.weight",
+ "proj_in.weight",
+ "time_embedder.linear_1.weight",
+ "model.visual.encoder.layers.0.self_attn.q_proj.weight",
+ "model.projector.linear_fc1.weight",
+ ]
+ self.assertEqual(self.mapper.apply_list(dropped), [])
+
+
+class TestCosmos3EdgeVideoSampling(CustomTestCase):
+ def _processor(self, video_config=None):
+ processor = object.__new__(Cosmos3EdgeProcessor)
+ processor.video_config = video_config or {}
+ return processor
+
+ def test_default_sampling_matches_qwen3_vl(self):
+ processor = self._processor()
+ indices = processor._select_frame_indices(total_frames=300, video_fps=30.0)
+ expected = np.linspace(0, 299, num=20).round().astype(np.int64).tolist()
+ self.assertEqual(indices, expected)
+
+ def test_default_sampling_clamps_to_frame_limits(self):
+ processor = self._processor()
+ self.assertEqual(
+ len(processor._select_frame_indices(total_frames=30, video_fps=30.0)),
+ 4,
+ )
+ self.assertEqual(
+ len(processor._select_frame_indices(total_frames=15_000, video_fps=10.0)),
+ 768,
+ )
+
+ def test_num_frames_and_legacy_nframes(self):
+ num_frames = self._processor({"num_frames": 5})
+ nframes = self._processor({"nframes": 5})
+ expected = [0, 25, 50, 74, 99]
+ self.assertEqual(num_frames._select_frame_indices(100, 30.0), expected)
+ self.assertEqual(nframes._select_frame_indices(100, 30.0), expected)
+
+ def test_explicit_frame_count_and_fps_are_mutually_exclusive(self):
+ processor = self._processor({"num_frames": 5, "fps": 2.0})
+ with self.assertRaisesRegex(ValueError, "Specify only one"):
+ processor._select_frame_indices(100, 30.0)
+
+ def test_missing_source_fps_uses_24_fps(self):
+ processor = self._processor()
+ self.assertEqual(
+ len(processor._select_frame_indices(total_frames=240, video_fps=0.0)),
+ 20,
+ )
+
+
+class TestCosmos3EdgeResize(CustomTestCase):
+ def test_video_max_pixels_is_a_total_temporal_budget(self):
+ height, width = _smart_resize(
+ 1024,
+ 1024,
+ factor=32,
+ min_pixels=4096,
+ max_pixels=4 * 1024 * 1024,
+ num_frames=16,
+ )
+ self.assertEqual((height, width), (512, 512))
+ self.assertLessEqual(16 * height * width, 4 * 1024 * 1024)
+
+ def test_image_resize_keeps_single_frame_semantics(self):
+ height, width = _smart_resize(
+ 1024,
+ 1024,
+ factor=32,
+ min_pixels=4096,
+ max_pixels=512 * 512,
+ )
+ self.assertEqual((height, width), (512, 512))
+
+
+class _TimestampTokenizer:
+ def encode(self, text, add_special_tokens=False):
+ del add_special_tokens
+ timestamp = float(text.removeprefix("<").split()[0])
+ return [1000 + int(timestamp * 10)]
+
+
+class TestCosmos3EdgePromptExpansion(CustomTestCase):
+ def _processor(self):
+ processor = object.__new__(Cosmos3EdgeProcessor)
+ processor.IMAGE_TOKEN_ID = 19
+ processor.VIDEO_TOKEN_ID = 18
+ processor.IM_START_TOKEN_ID = 20
+ processor.IM_END_TOKEN_ID = 21
+ processor._spatial_merge_size = 2
+ processor._tokenizer = _TimestampTokenizer()
+ return processor
+
+ def test_video_placeholder_expands_once_per_frame(self):
+ processor = self._processor()
+ prompt = [7, 20, 18, 21, 8]
+ video_grid_thw = torch.tensor([[3, 4, 4]], dtype=torch.long)
+ timestamps = [[0.0, 0.5, 1.0]]
+
+ input_ids, offsets, modalities = processor._build_input_ids(
+ prompt,
+ img_grid_thw=None,
+ video_grid_thw=video_grid_thw,
+ video_timestamps=timestamps,
+ )
+
+ expected = [7]
+ expected_offsets = []
+ for timestamp_id in (1000, 1005, 1010):
+ expected.extend([timestamp_id, 20])
+ offset_start = len(expected)
+ expected.extend([18] * 4)
+ expected_offsets.append((offset_start, len(expected) - 1))
+ expected.append(21)
+ expected.append(8)
+
+ self.assertEqual(input_ids, expected)
+ self.assertEqual(offsets, [expected_offsets])
+ self.assertEqual(modalities, [Modality.VIDEO])
+
+ def test_expanded_video_mrope_matches_qwen3_vl(self):
+ processor = self._processor()
+ video_grid_thw = torch.tensor([[3, 4, 4]], dtype=torch.long)
+ input_ids, _, _ = processor._build_input_ids(
+ [7, 20, 18, 21, 8],
+ img_grid_thw=None,
+ video_grid_thw=video_grid_thw,
+ video_timestamps=[[0.0, 0.5, 1.0]],
+ )
+ input_ids = torch.tensor([input_ids], dtype=torch.long)
+
+ kwargs = dict(
+ spatial_merge_size=2,
+ image_token_id=19,
+ video_token_id=18,
+ vision_start_token_id=20,
+ input_ids=input_ids,
+ video_grid_thw=video_grid_thw,
+ )
+ edge_positions, edge_delta = MRotaryEmbedding.get_rope_index(
+ model_type="cosmos3_edge", **kwargs
+ )
+ qwen_positions, qwen_delta = MRotaryEmbedding.get_rope_index(
+ model_type="qwen3_vl", **kwargs
+ )
+
+ self.assertEqual(edge_positions.shape, (3, 1, len(input_ids[0])))
+ self.assertTrue(torch.equal(edge_positions, qwen_positions))
+ self.assertTrue(torch.equal(edge_delta, qwen_delta))
+
+
+def _processed_input_processor():
+ processor = object.__new__(Cosmos3EdgeProcessor)
+ processor.IMAGE_TOKEN_ID = 19
+ processor.VIDEO_TOKEN_ID = 18
+ processor.IM_START_TOKEN_ID = 20
+ processor.IM_END_TOKEN_ID = 21
+ processor.vision_start_token_id = 20
+ processor.model_type = "cosmos3_edge"
+ processor._spatial_merge_size = 2
+ processor._tokenizer = _TimestampTokenizer()
+ processor._processor = processor._tokenizer
+ processor.mm_tokens = MultimodalSpecialTokens(
+ image_token="",
+ video_token="