[Feature] Add Muse Glimmer model support (#34262)
Co-authored-by: sglang-bot <232288953+sglang-bot@users.noreply.github.com> Co-authored-by: Brayden Zhong <brayden.zhong@radixark.ai> Co-authored-by: Jimmy Shong <69131491+Jiminator@users.noreply.github.com> Co-authored-by: hnyls2002 <lsyincs@gmail.com> Co-authored-by: Alex Nails <alex.nails@radixark.ai> Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
This commit is contained in:
co-authored by
sglang-bot
Brayden Zhong
Jimmy Shong
hnyls2002
Alex Nails
Liangsheng Yin
parent
9c1517df4a
commit
fde9ad2531
@@ -48,9 +48,7 @@ def get_tokenizer(
|
||||
pretrained_model_name_or_path is not None
|
||||
and pretrained_model_name_or_path != ""
|
||||
)
|
||||
if pretrained_model_name_or_path.endswith(
|
||||
".json"
|
||||
) or pretrained_model_name_or_path.endswith(".model"):
|
||||
if pretrained_model_name_or_path.endswith((".json", ".model", ".gguf")):
|
||||
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||
|
||||
return get_tokenizer(pretrained_model_name_or_path)
|
||||
|
||||
@@ -1930,6 +1930,14 @@ def _deepseek_v4_sm120_moe(view: Any) -> dict:
|
||||
return {}
|
||||
|
||||
|
||||
@_register_for("MuseGlimmerForConditionalGeneration", "MuseGlimmerForCausalLM")
|
||||
def _muse_glimmer_fp4_gemm_runner_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if is_sm120_supported() and server_args.fp4_gemm_runner_backend == "auto":
|
||||
logger.info("Use marlin as FP4 GEMM runner backend on SM120 for Muse Glimmer")
|
||||
return {"fp4_gemm_runner_backend": "marlin"}
|
||||
return {}
|
||||
|
||||
|
||||
@register_post_process
|
||||
def _sparse_head_overlap_disable(view: Any) -> dict:
|
||||
|
||||
|
||||
@@ -37,6 +37,10 @@ from sglang.srt.configs.locate_anything import LocateAnythingConfig
|
||||
from sglang.srt.configs.longcat_flash import LongcatFlashConfig
|
||||
from sglang.srt.configs.minicpmv4_6 import MiniCPMV4_6Config, MiniCPMV4_6VisionConfig
|
||||
from sglang.srt.configs.minimax_vl import MiniMaxM3VLConfig
|
||||
from sglang.srt.configs.muse_glimmer import (
|
||||
MuseGlimmerAssistantConfig,
|
||||
MuseGlimmerConfig,
|
||||
)
|
||||
from sglang.srt.configs.nano_nemotron_vl import (
|
||||
NemotronH_Nano_Omni_Reasoning_V3_Config,
|
||||
NemotronH_Nano_VL_V2_Config,
|
||||
@@ -76,6 +80,8 @@ __all__ = [
|
||||
"Step3TextConfig",
|
||||
"Step3VisionEncoderConfig",
|
||||
"Olmo3Config",
|
||||
"MuseGlimmerConfig",
|
||||
"MuseGlimmerAssistantConfig",
|
||||
"KimiLinearConfig",
|
||||
"KimiK3Config",
|
||||
"KimiK25Config",
|
||||
|
||||
@@ -264,6 +264,7 @@ class ModelConfig:
|
||||
is_multi_layer_eagle: bool = False,
|
||||
encoder_only: bool = False,
|
||||
language_only: bool = False,
|
||||
language_model_only: bool = False,
|
||||
disable_hybrid_swa_memory: bool = False,
|
||||
model_config_parser: str = "auto",
|
||||
speculative_algorithm: Optional[str] = None,
|
||||
@@ -451,7 +452,8 @@ class ModelConfig:
|
||||
)
|
||||
# TODO: requires further polishing
|
||||
# Key on the tower, not the attribute: several config classes default
|
||||
# vision_config to None, which presence alone would read as image-capable.
|
||||
# vision_config to None, which presence alone would read as image-capable
|
||||
# (MuseGlimmerConfig's text-only layouts are one such case).
|
||||
self.is_image_understandable_model = (
|
||||
enable_multimodal
|
||||
and not self.is_lm_only
|
||||
@@ -534,6 +536,10 @@ class ModelConfig:
|
||||
|
||||
self.hf_config.encoder_only = encoder_only
|
||||
self.hf_config.language_only = language_only
|
||||
# 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
|
||||
|
||||
# matryoshka embeddings
|
||||
self.matryoshka_dimensions = getattr(
|
||||
@@ -582,6 +588,7 @@ class ModelConfig:
|
||||
override_config_file=override_config_file,
|
||||
is_multi_layer_eagle=server_args.enable_multi_layer_eagle,
|
||||
language_only=server_args.language_only,
|
||||
language_model_only=server_args.language_model_only,
|
||||
encoder_only=server_args.encoder_only,
|
||||
is_draft_model=is_draft_model,
|
||||
is_draft_quantization_explicit=(
|
||||
@@ -1830,6 +1837,7 @@ multimodal_model_archs = [
|
||||
"MossVLForConditionalGeneration",
|
||||
"NemotronH_Nano_VL_V2",
|
||||
"NemotronH_Nano_Omni_Reasoning_V3",
|
||||
"MuseGlimmerForConditionalGeneration",
|
||||
"PixtralForConditionalGeneration",
|
||||
"Qwen2AudioForConditionalGeneration",
|
||||
"Qwen2VLForConditionalGeneration",
|
||||
@@ -1893,6 +1901,7 @@ multimodal_breakable_cuda_graph_supported_model_archs = [
|
||||
"InternS2MobiusForConditionalGeneration",
|
||||
"Qwen3_5ForConditionalGeneration",
|
||||
"Qwen3_5MoeForConditionalGeneration",
|
||||
"MuseGlimmerForConditionalGeneration",
|
||||
]
|
||||
|
||||
if external_mm_model_arch := envs.SGLANG_EXTERNAL_MM_MODEL_ARCH.get():
|
||||
@@ -2036,6 +2045,8 @@ def is_hybrid_swa_model(
|
||||
"Gemma4UnifiedForConditionalGeneration",
|
||||
"LagunaForCausalLM",
|
||||
"MellumForCausalLM",
|
||||
"MuseGlimmerForCausalLM",
|
||||
"MuseGlimmerForConditionalGeneration",
|
||||
"InklingForConditionalGeneration",
|
||||
"InklingForConditionalGenerationMTP",
|
||||
"UnlimitedOCRForCausalLM",
|
||||
@@ -2111,6 +2122,8 @@ def get_hybrid_layer_ids(
|
||||
or "Gemma4UnifiedForConditionalGeneration" in model_architectures
|
||||
or "LagunaForCausalLM" in model_architectures
|
||||
or "MellumForCausalLM" in model_architectures
|
||||
or "MuseGlimmerForCausalLM" in model_architectures
|
||||
or "MuseGlimmerForConditionalGeneration" in model_architectures
|
||||
):
|
||||
layer_types = getattr(hf_text_config, "layer_types", [])
|
||||
swa_attention_layer_ids = [
|
||||
|
||||
@@ -0,0 +1,282 @@
|
||||
# Copyright 2023-2026 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
|
||||
import math
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from sglang.srt.configs.muse_glimmer_processing import MuseGlimmerProcessor
|
||||
from sglang.srt.multimodal.customized_mm_processor_utils import (
|
||||
register_customized_processor,
|
||||
)
|
||||
|
||||
_ARCH = "muse-glimmer"
|
||||
|
||||
|
||||
class MuseGlimmerAssistantConfig(PretrainedConfig):
|
||||
|
||||
model_type = "muse_glimmer_assistant"
|
||||
# The DFlash draft has no head; draft_worker_common borrows the target's.
|
||||
vocab_size = None
|
||||
|
||||
|
||||
class MuseGlimmerVisionConfig(PretrainedConfig):
|
||||
|
||||
model_type = "muse_glimmer_vision"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int = 1536,
|
||||
intermediate_size: int = 8960,
|
||||
num_hidden_layers: int = 50,
|
||||
num_attention_heads: int = 16,
|
||||
hidden_act: str = "gelu",
|
||||
layer_norm_eps: float = 1e-5,
|
||||
attention_types: Optional[List[str]] = None,
|
||||
max_position_embeddings: int = 1024,
|
||||
merge_size: int = 2,
|
||||
patch_size: int = 14,
|
||||
patch_temporal: int = 2,
|
||||
pos_emb_height: int = 32,
|
||||
pos_emb_width: int = 32,
|
||||
rope_theta: float = 10000.0,
|
||||
**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.hidden_act = hidden_act
|
||||
self.layer_norm_eps = layer_norm_eps
|
||||
self.attention_types = attention_types
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.merge_size = merge_size
|
||||
self.patch_size = patch_size
|
||||
self.patch_temporal = patch_temporal
|
||||
self.pos_emb_height = pos_emb_height
|
||||
self.pos_emb_width = pos_emb_width
|
||||
self.rope_theta = rope_theta
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
||||
@register_customized_processor(MuseGlimmerProcessor)
|
||||
class MuseGlimmerConfig(PretrainedConfig):
|
||||
model_type = "muse_glimmer"
|
||||
sub_configs = {"vision_config": MuseGlimmerVisionConfig}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int = 202048,
|
||||
hidden_size: int = 6656,
|
||||
intermediate_size: int = 19968,
|
||||
num_hidden_layers: int = 52,
|
||||
num_attention_heads: int = 32,
|
||||
num_key_value_heads: int = 2,
|
||||
head_dim: int = 128,
|
||||
hidden_act: str = "silu",
|
||||
max_position_embeddings: int = 16384,
|
||||
rms_norm_eps: float = 1e-5,
|
||||
post_norm_eps: float = 1e-8,
|
||||
rope_theta: float = 500000.0,
|
||||
sliding_window: int = 2048,
|
||||
layer_types: Optional[List[str]] = None,
|
||||
no_rope_layers: Optional[List[int]] = None,
|
||||
use_qk_norm: bool = True,
|
||||
use_attn_output_gate: bool = True,
|
||||
qk_scale_factor: float = 43.7840518911,
|
||||
rope_is_neox_style: bool = False,
|
||||
normalize_tok_embeddings: bool = True,
|
||||
output_multiplier: float = 0.19611613513818404,
|
||||
output_soft_cap_temp: Optional[float] = 20.0,
|
||||
tie_word_embeddings: bool = False,
|
||||
bos_token_id: int = 200000,
|
||||
eos_token_id: int = 200001,
|
||||
vision_config: Optional[Dict[str, Any]] = None,
|
||||
image_token_id: Optional[int] = None,
|
||||
video_token_id: Optional[int] = None,
|
||||
out_hidden_size: int = 6144,
|
||||
projector_hidden_act: str = "gelu",
|
||||
projector_hidden_size: int = 4096,
|
||||
**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.max_position_embeddings = max_position_embeddings
|
||||
self.rms_norm_eps = rms_norm_eps
|
||||
self.post_norm_eps = post_norm_eps
|
||||
self.rope_theta = rope_theta
|
||||
self.sliding_window = sliding_window
|
||||
self.layer_types = layer_types
|
||||
self.no_rope_layers = no_rope_layers
|
||||
self.use_qk_norm = use_qk_norm
|
||||
self.use_attn_output_gate = use_attn_output_gate
|
||||
self.qk_scale_factor = qk_scale_factor
|
||||
self.rope_is_neox_style = rope_is_neox_style
|
||||
self.normalize_tok_embeddings = normalize_tok_embeddings
|
||||
self.output_multiplier = output_multiplier
|
||||
self.output_soft_cap_temp = output_soft_cap_temp
|
||||
if isinstance(vision_config, dict):
|
||||
vision_config = MuseGlimmerVisionConfig(**vision_config)
|
||||
self.vision_config = vision_config
|
||||
self.image_token_id = image_token_id
|
||||
self.video_token_id = video_token_id
|
||||
self.out_hidden_size = out_hidden_size
|
||||
self.projector_hidden_act = projector_hidden_act
|
||||
self.projector_hidden_size = projector_hidden_size
|
||||
super().__init__(
|
||||
tie_word_embeddings=tie_word_embeddings,
|
||||
bos_token_id=bos_token_id,
|
||||
eos_token_id=eos_token_id,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_gguf(cls, gguf_path: str) -> "MuseGlimmerConfig":
|
||||
return cls(**muse_glimmer_config_kwargs_from_gguf(gguf_path))
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config_dict: Dict[str, Any], **kwargs):
|
||||
return super().from_dict(
|
||||
muse_glimmer_config_kwargs_from_hf(config_dict), **kwargs
|
||||
)
|
||||
|
||||
|
||||
_HF_TEXT_KEYS_TRANSLATED = frozenset(
|
||||
{
|
||||
"final_logit_softcapping",
|
||||
"hidden_activation",
|
||||
"layer_rope_theta",
|
||||
"model_type",
|
||||
"qk_scale_factor",
|
||||
"rope_parameters",
|
||||
}
|
||||
)
|
||||
|
||||
_HF_NESTED_KEYS = ("text_config", "vision_config")
|
||||
|
||||
_HF_VISION_KEYS_TRANSLATED = frozenset({"layer_types", "model_type", "rope_parameters"})
|
||||
|
||||
|
||||
def muse_glimmer_config_kwargs_from_hf(config_dict: Dict[str, Any]) -> Dict[str, Any]:
|
||||
if "text_config" not in config_dict:
|
||||
return config_dict
|
||||
|
||||
text = config_dict["text_config"]
|
||||
kwargs = {k: v for k, v in config_dict.items() if k not in _HF_NESTED_KEYS}
|
||||
kwargs.update({k: v for k, v in text.items() if k not in _HF_TEXT_KEYS_TRANSLATED})
|
||||
kwargs.update(
|
||||
hidden_act=text["hidden_activation"],
|
||||
rope_theta=text["rope_parameters"]["rope_theta"],
|
||||
no_rope_layers=[1 if theta else 0 for theta in text["layer_rope_theta"]],
|
||||
output_soft_cap_temp=text["final_logit_softcapping"],
|
||||
qk_scale_factor=text["qk_scale_factor"] * math.sqrt(text["head_dim"]),
|
||||
rope_is_neox_style=True,
|
||||
)
|
||||
if "vision_config" in config_dict:
|
||||
kwargs["vision_config"] = muse_glimmer_vision_config_kwargs_from_hf(
|
||||
config_dict["vision_config"]
|
||||
)
|
||||
return kwargs
|
||||
|
||||
|
||||
def muse_glimmer_vision_config_kwargs_from_hf(
|
||||
vision_config_dict: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""Translate the vendor's ``vision_config`` into ``MuseGlimmerVisionConfig`` kwargs."""
|
||||
kwargs = {
|
||||
k: v
|
||||
for k, v in vision_config_dict.items()
|
||||
if k not in _HF_VISION_KEYS_TRANSLATED
|
||||
}
|
||||
kwargs["rope_theta"] = vision_config_dict["rope_parameters"]["rope_theta"]
|
||||
kwargs["attention_types"] = vision_config_dict["layer_types"]
|
||||
return kwargs
|
||||
|
||||
|
||||
def _f(v):
|
||||
return None if v is None else float(v)
|
||||
|
||||
|
||||
def _i(v):
|
||||
return None if v is None else int(v)
|
||||
|
||||
|
||||
def _mul_sqrt(v, head_dim):
|
||||
return None if v is None else float(v) * math.sqrt(head_dim)
|
||||
|
||||
|
||||
def muse_glimmer_config_kwargs_from_gguf(gguf_path: str) -> Dict[str, Any]:
|
||||
from gguf import GGUFReader
|
||||
|
||||
reader = GGUFReader(gguf_path)
|
||||
meta = {key: field.contents() for key, field in reader.fields.items()}
|
||||
shapes = {t.name: tuple(int(x) for x in t.shape) for t in reader.tensors}
|
||||
tensor_names = set(shapes)
|
||||
|
||||
def get(suffix):
|
||||
return meta[f"{_ARCH}.{suffix}"]
|
||||
|
||||
def opt(suffix):
|
||||
"""None when this converter generation did not emit the key."""
|
||||
return meta.get(f"{_ARCH}.{suffix}")
|
||||
|
||||
head_dim = int(get("attention.key_length"))
|
||||
swa_pattern = [bool(x) for x in get("attention.sliding_window_pattern")]
|
||||
|
||||
return dict(
|
||||
# token_embd is stored [n_embd, n_vocab] in ggml's reversed order.
|
||||
vocab_size=shapes["token_embd.weight"][1],
|
||||
hidden_size=int(get("embedding_length")),
|
||||
intermediate_size=int(get("feed_forward_length")),
|
||||
num_hidden_layers=int(get("block_count")),
|
||||
num_attention_heads=int(get("attention.head_count")),
|
||||
num_key_value_heads=int(get("attention.head_count_kv")),
|
||||
head_dim=head_dim,
|
||||
max_position_embeddings=int(get("context_length")),
|
||||
rms_norm_eps=float(get("attention.layer_norm_rms_epsilon")),
|
||||
rope_theta=float(get("rope.freq_base")),
|
||||
sliding_window=int(get("attention.sliding_window")),
|
||||
layer_types=[
|
||||
"sliding_attention" if s else "full_attention" for s in swa_pattern
|
||||
],
|
||||
no_rope_layers=[1 if s else 0 for s in swa_pattern],
|
||||
use_qk_norm=any(n.endswith("attn_q_norm.weight") for n in tensor_names),
|
||||
use_attn_output_gate=any(n.endswith("attn_gate.weight") for n in tensor_names),
|
||||
tie_word_embeddings="output.weight" not in tensor_names,
|
||||
architectures=["MuseGlimmerForCausalLM"],
|
||||
dtype="bfloat16",
|
||||
# Converter generations differ in which of these they emit, and every one
|
||||
# is an architecture constant that MuseGlimmerConfig already defaults to,
|
||||
# so an absent key falls back rather than raising. attention.scale is
|
||||
# stored pre-divided by sqrt(head_dim); the class stores it before that.
|
||||
**{
|
||||
k: v
|
||||
for k, v in (
|
||||
("post_norm_eps", _f(opt("attention.post_norm_rms_epsilon"))),
|
||||
("qk_scale_factor", _mul_sqrt(opt("attention.scale"), head_dim)),
|
||||
("output_multiplier", _f(opt("logit_scale"))),
|
||||
("output_soft_cap_temp", _f(opt("final_logit_softcapping"))),
|
||||
("bos_token_id", _i(meta.get("tokenizer.ggml.bos_token_id"))),
|
||||
("eos_token_id", _i(meta.get("tokenizer.ggml.eos_token_id"))),
|
||||
)
|
||||
if v is not None
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,302 @@
|
||||
# Copyright 2023-2026 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
|
||||
import itertools
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from transformers import AutoTokenizer
|
||||
from transformers.image_processing_backends import TorchvisionBackend
|
||||
from transformers.image_processing_utils import BatchFeature
|
||||
from transformers.image_transforms import group_images_by_shape, reorder_images
|
||||
from transformers.image_utils import PILImageResampling, SizeDict
|
||||
from transformers.processing_utils import ImagesKwargs, MultiModalData, ProcessorMixin
|
||||
from transformers.utils import TensorType
|
||||
from transformers.utils.constants import IMAGENET_STANDARD_MEAN, IMAGENET_STANDARD_STD
|
||||
from transformers.utils.hub import cached_file
|
||||
|
||||
PROCESSOR_CONFIG_NAME = "processor_config.json"
|
||||
|
||||
|
||||
def get_aspect_ratio_preserving_size(
|
||||
height: int,
|
||||
width: int,
|
||||
patch_size: int,
|
||||
max_tokens: int,
|
||||
) -> tuple[int, int]:
|
||||
"""Patch grid closest to the aspect ratio; returns (height, width) in pixels."""
|
||||
ideal_patches_height = height / patch_size
|
||||
ideal_patches_width = width / patch_size
|
||||
ratio = (
|
||||
ideal_patches_width / ideal_patches_height if ideal_patches_height > 0 else 1.0
|
||||
)
|
||||
if ideal_patches_height * ideal_patches_width > max_tokens:
|
||||
ideal_patches_height = (max_tokens / ratio) ** 0.5
|
||||
ideal_patches_width = ideal_patches_height * ratio
|
||||
candidates = list(
|
||||
set(
|
||||
itertools.product(
|
||||
[math.floor(ideal_patches_height), math.ceil(ideal_patches_height)],
|
||||
[math.floor(ideal_patches_width), math.ceil(ideal_patches_width)],
|
||||
)
|
||||
)
|
||||
)
|
||||
candidates = [
|
||||
(patches_height, patches_width)
|
||||
for patches_height, patches_width in candidates
|
||||
if patches_height >= 1
|
||||
and patches_width >= 1
|
||||
and patches_height * patches_width <= max_tokens
|
||||
]
|
||||
if not candidates:
|
||||
candidates = [
|
||||
(max(1, round(ideal_patches_height)), max(1, round(ideal_patches_width)))
|
||||
]
|
||||
patches_height, patches_width = min(
|
||||
candidates, key=lambda grid: abs(grid[0] / grid[1] - height / width)
|
||||
)
|
||||
return patches_height * patch_size, patches_width * patch_size
|
||||
|
||||
|
||||
class MuseGlimmerImageProcessorKwargs(ImagesKwargs, total=False):
|
||||
patch_size: int
|
||||
temporal_patch_size: int
|
||||
merge_size: int
|
||||
max_image_tokens: int
|
||||
|
||||
|
||||
class MuseGlimmerImageProcessor(TorchvisionBackend):
|
||||
do_resize = True
|
||||
resample = PILImageResampling.LANCZOS
|
||||
size = None
|
||||
default_to_square = False
|
||||
do_rescale = True
|
||||
rescale_factor = 1 / 255
|
||||
do_normalize = True
|
||||
image_mean = IMAGENET_STANDARD_MEAN
|
||||
image_std = IMAGENET_STANDARD_STD
|
||||
do_convert_rgb = True
|
||||
patch_size = 14
|
||||
temporal_patch_size = 2
|
||||
merge_size = 2
|
||||
max_image_tokens = 4096
|
||||
valid_kwargs = MuseGlimmerImageProcessorKwargs
|
||||
model_input_names = ["pixel_values", "image_grid_thw"]
|
||||
|
||||
def _preprocess(
|
||||
self,
|
||||
images: list[torch.Tensor],
|
||||
do_resize: bool,
|
||||
resample,
|
||||
do_rescale: bool,
|
||||
rescale_factor: float,
|
||||
do_normalize: bool,
|
||||
image_mean,
|
||||
image_std,
|
||||
return_tensors: Optional[TensorType],
|
||||
patch_size: int,
|
||||
temporal_patch_size: int,
|
||||
max_image_tokens: int,
|
||||
merge_size: int,
|
||||
disable_grouping: bool = False,
|
||||
**kwargs,
|
||||
) -> BatchFeature:
|
||||
if resample == PILImageResampling.LANCZOS:
|
||||
# BICUBIC stands in for LANCZOS, which is CPU-only.
|
||||
resample = PILImageResampling.BICUBIC
|
||||
|
||||
grouped_images, grouped_images_index = group_images_by_shape(
|
||||
images, disable_grouping=disable_grouping
|
||||
)
|
||||
resized_images_grouped = {}
|
||||
for shape, stacked_images in grouped_images.items():
|
||||
if do_resize:
|
||||
height, width = stacked_images.shape[-2:]
|
||||
resized_height, resized_width = get_aspect_ratio_preserving_size(
|
||||
height=height,
|
||||
width=width,
|
||||
patch_size=patch_size * merge_size,
|
||||
max_tokens=max_image_tokens,
|
||||
)
|
||||
stacked_images = self.resize(
|
||||
image=stacked_images,
|
||||
size=SizeDict(height=resized_height, width=resized_width),
|
||||
resample=resample,
|
||||
antialias=True,
|
||||
)
|
||||
resized_images_grouped[shape] = stacked_images
|
||||
resized_images = reorder_images(resized_images_grouped, grouped_images_index)
|
||||
|
||||
grouped_images, grouped_images_index = group_images_by_shape(
|
||||
resized_images, disable_grouping=disable_grouping
|
||||
)
|
||||
processed_images_grouped = {}
|
||||
processed_grids = {}
|
||||
for shape, stacked_images in grouped_images.items():
|
||||
resized_height, resized_width = stacked_images.shape[-2:]
|
||||
patches = self.rescale_and_normalize(
|
||||
stacked_images,
|
||||
do_rescale,
|
||||
rescale_factor,
|
||||
do_normalize,
|
||||
image_mean,
|
||||
image_std,
|
||||
)
|
||||
if patches.ndim == 4:
|
||||
patches = patches.unsqueeze(1)
|
||||
|
||||
if patches.shape[1] % temporal_patch_size != 0:
|
||||
repeats = patches[:, -1:].repeat(1, temporal_patch_size - 1, 1, 1, 1)
|
||||
patches = torch.cat([patches, repeats], dim=1)
|
||||
|
||||
batch_size, grid_t, channel = patches.shape[:3]
|
||||
grid_t = grid_t // temporal_patch_size
|
||||
grid_h, grid_w = resized_height // patch_size, resized_width // patch_size
|
||||
|
||||
patches = patches.view(
|
||||
batch_size,
|
||||
grid_t,
|
||||
temporal_patch_size,
|
||||
channel,
|
||||
grid_h,
|
||||
patch_size,
|
||||
grid_w,
|
||||
patch_size,
|
||||
)
|
||||
patches = patches.permute(0, 1, 4, 6, 2, 3, 5, 7)
|
||||
flatten_patches = patches.reshape(
|
||||
batch_size,
|
||||
grid_t * grid_h * grid_w,
|
||||
temporal_patch_size * channel * patch_size * patch_size,
|
||||
)
|
||||
|
||||
processed_images_grouped[shape] = flatten_patches
|
||||
processed_grids[shape] = [[grid_t, grid_h, grid_w]] * batch_size
|
||||
|
||||
processed_images = reorder_images(
|
||||
processed_images_grouped, grouped_images_index
|
||||
)
|
||||
processed_grids = reorder_images(processed_grids, grouped_images_index)
|
||||
pixel_values = torch.cat(processed_images, dim=0)
|
||||
image_grid_thw = torch.tensor(processed_grids)
|
||||
|
||||
return BatchFeature(
|
||||
data={"pixel_values": pixel_values, "image_grid_thw": image_grid_thw},
|
||||
tensor_type=return_tensors,
|
||||
)
|
||||
|
||||
def get_number_of_image_patches(self, height: int, width: int, images_kwargs=None):
|
||||
"""Patch rows a (height, width) image expands to."""
|
||||
images_kwargs = images_kwargs or {}
|
||||
patch_size = images_kwargs.get("patch_size", self.patch_size)
|
||||
merge_size = images_kwargs.get("merge_size", self.merge_size)
|
||||
max_image_tokens = images_kwargs.get("max_image_tokens", self.max_image_tokens)
|
||||
|
||||
resized_height, resized_width = get_aspect_ratio_preserving_size(
|
||||
height=height,
|
||||
width=width,
|
||||
patch_size=patch_size * merge_size,
|
||||
max_tokens=max_image_tokens,
|
||||
)
|
||||
return (resized_height // patch_size) * (resized_width // patch_size)
|
||||
|
||||
def _validate_preprocess_kwargs(self, **kwargs):
|
||||
kwargs["do_resize"] = False
|
||||
super()._validate_preprocess_kwargs(**kwargs)
|
||||
|
||||
|
||||
class MuseGlimmerProcessor(ProcessorMixin):
|
||||
"""Expands one ``<|patch|>`` placeholder into an image's patch-token run."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
image_processor=None,
|
||||
tokenizer=None,
|
||||
chat_template=None,
|
||||
**kwargs,
|
||||
):
|
||||
self.image_token = "<|patch|>"
|
||||
self.image_token_id = tokenizer.convert_tokens_to_ids(self.image_token)
|
||||
|
||||
super().__init__(
|
||||
image_processor=image_processor,
|
||||
tokenizer=tokenizer,
|
||||
chat_template=chat_template,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
|
||||
trust_remote_code = kwargs.pop("trust_remote_code", False)
|
||||
revision = kwargs.pop("revision", None)
|
||||
use_fast = kwargs.pop("use_fast", True)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
pretrained_model_name_or_path,
|
||||
trust_remote_code=trust_remote_code,
|
||||
revision=revision,
|
||||
use_fast=use_fast,
|
||||
)
|
||||
image_processor = MuseGlimmerImageProcessor(
|
||||
**_load_image_processor_kwargs(pretrained_model_name_or_path, revision)
|
||||
)
|
||||
return cls(
|
||||
image_processor=image_processor,
|
||||
tokenizer=tokenizer,
|
||||
chat_template=tokenizer.chat_template,
|
||||
)
|
||||
|
||||
def replace_image_token(self, image_inputs: dict, image_idx: int, **kwargs) -> str:
|
||||
merge_length = self.image_processor.merge_size**2
|
||||
num_image_tokens = (
|
||||
image_inputs["image_grid_thw"][image_idx].prod() // merge_length
|
||||
)
|
||||
return self.image_token * num_image_tokens
|
||||
|
||||
def _get_num_multimodal_tokens(self, image_sizes=None, **kwargs):
|
||||
"""Placeholder counts per image, without running the image processor."""
|
||||
vision_data = {}
|
||||
if image_sizes is not None:
|
||||
merge_size = self.image_processor.merge_size
|
||||
num_image_patches = [
|
||||
self.image_processor.get_number_of_image_patches(height, width, kwargs)
|
||||
for height, width in image_sizes
|
||||
]
|
||||
vision_data.update(
|
||||
num_image_tokens=[
|
||||
patches // merge_size**2 for patches in num_image_patches
|
||||
],
|
||||
num_image_patches=num_image_patches,
|
||||
)
|
||||
return MultiModalData(**vision_data)
|
||||
|
||||
|
||||
def _load_image_processor_kwargs(model_path: str, revision: Optional[str]) -> dict:
|
||||
"""Read the image_processor block from processor_config.json."""
|
||||
local = os.path.join(model_path, PROCESSOR_CONFIG_NAME)
|
||||
config_file = (
|
||||
local
|
||||
if os.path.isfile(local)
|
||||
else cached_file(model_path, PROCESSOR_CONFIG_NAME, revision=revision)
|
||||
)
|
||||
with open(config_file) as f:
|
||||
config = json.load(f)
|
||||
image_processor_config = config.get("image_processor", {})
|
||||
return {
|
||||
key: value
|
||||
for key, value in image_processor_config.items()
|
||||
if key != "image_processor_type"
|
||||
}
|
||||
@@ -2196,6 +2196,7 @@ def _execute_server_warmup(server_args: ServerArgs):
|
||||
is_vlm = (
|
||||
bool(model_info.get("has_image_understanding", False))
|
||||
and not server_args.language_only
|
||||
and not server_args.language_model_only
|
||||
and not is_mps()
|
||||
)
|
||||
if model_info["is_generation"]:
|
||||
|
||||
@@ -727,11 +727,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
remaining_logprobs = None
|
||||
|
||||
# Handle tool calls
|
||||
if (
|
||||
request.tool_choice != "none"
|
||||
and self._effective_tools(request)
|
||||
and self.tool_call_parser
|
||||
):
|
||||
if self._tool_call_parsing_active(request):
|
||||
async for chunk in self._process_tool_call_stream(
|
||||
index,
|
||||
delta,
|
||||
@@ -740,6 +736,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
request,
|
||||
has_tool_calls,
|
||||
continuous_usage_stats,
|
||||
flush=finish_reason_type is not None and finish_reason_type != "abort",
|
||||
):
|
||||
if chunk:
|
||||
yield chunk
|
||||
@@ -802,6 +799,18 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
usage=usage,
|
||||
)
|
||||
|
||||
def _tool_call_parsing_active(self, request: ChatCompletionRequest) -> bool:
|
||||
"""Whether this request's output runs through the tool-call parser.
|
||||
|
||||
The reasoning parser is told the same thing, so channel-framed formats
|
||||
keep their framing intact exactly when a tool-call parser consumes it.
|
||||
"""
|
||||
return bool(
|
||||
request.tool_choice != "none"
|
||||
and self._effective_tools(request)
|
||||
and self.tool_call_parser
|
||||
)
|
||||
|
||||
def _validate_request(self, request: ChatCompletionRequest) -> Optional[str]:
|
||||
"""Validate that the input is valid."""
|
||||
if not request.messages:
|
||||
@@ -1129,6 +1138,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
|
||||
result.tool_call_constraint = tool_call_constraint
|
||||
result.require_reasoning = thinking_mode
|
||||
result.skip_special_tokens = request.skip_special_tokens
|
||||
return result
|
||||
|
||||
def _apply_jinja_template(
|
||||
@@ -1797,6 +1807,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
force_reasoning=force_reasoning,
|
||||
request=request,
|
||||
tokenizer=self.tokenizer_manager.tokenizer,
|
||||
tool_call_parser_active=self._tool_call_parsing_active(request),
|
||||
)
|
||||
reasoning_text, text = parser.parse_non_stream(text)
|
||||
except Exception as e:
|
||||
@@ -1810,11 +1821,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
# Handle tool calls
|
||||
tool_calls = None
|
||||
effective_tools = self._effective_tools(request)
|
||||
if (
|
||||
request.tool_choice != "none"
|
||||
and effective_tools
|
||||
and self.tool_call_parser
|
||||
):
|
||||
if self._tool_call_parsing_active(request):
|
||||
history_tool_calls_cnt = self._get_history_tool_calls_cnt(request)
|
||||
tool_calls, text, finish_reason = self._process_tool_calls(
|
||||
text,
|
||||
@@ -2106,6 +2113,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
is_force_reasoning,
|
||||
request,
|
||||
tokenizer=self.tokenizer_manager.tokenizer,
|
||||
tool_call_parser_active=self._tool_call_parsing_active(request),
|
||||
)
|
||||
reasoning_parser = reasoning_parser_dict[index]
|
||||
reasoning_text, normal_text = reasoning_parser.parse_stream_chunk(delta)
|
||||
@@ -2158,6 +2166,8 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
request.skip_special_tokens = False
|
||||
elif self.reasoning_parser == "inkling":
|
||||
request.skip_special_tokens = False
|
||||
elif self.reasoning_parser == "muse":
|
||||
request.skip_special_tokens = False
|
||||
|
||||
def wrap_reasoning_history(self, reasoning_text: str) -> str:
|
||||
"""Wrap prior-turn reasoning in the detector's own start/end tokens.
|
||||
@@ -2357,8 +2367,13 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
request: ChatCompletionRequest,
|
||||
has_tool_calls: Dict[int, bool],
|
||||
continuous_usage_stats: bool = False,
|
||||
flush: bool = False,
|
||||
):
|
||||
"""Process tool calls in streaming response"""
|
||||
"""Process tool calls in streaming response.
|
||||
|
||||
With flush=True (the terminal delta), the parser also drains text it
|
||||
held back waiting for a marker that can no longer arrive.
|
||||
"""
|
||||
effective_tools = self._effective_tools(request)
|
||||
if index not in parser_dict:
|
||||
is_required = request.tool_choice == "required" or isinstance(
|
||||
@@ -2400,6 +2415,10 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
normal_text, calls = result.normal_text, result.calls
|
||||
else:
|
||||
normal_text, calls = parser.parse_stream_chunk(delta)
|
||||
if flush:
|
||||
end_text, end_calls = parser.parse_stream_end()
|
||||
normal_text = (normal_text or "") + end_text
|
||||
calls = list(calls) + end_calls
|
||||
|
||||
# Yield normal text
|
||||
if normal_text:
|
||||
|
||||
@@ -394,7 +394,6 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
# _process_messages set skip_special_tokens on a chat_request
|
||||
# we then discard, so re-apply it to the engine sampling dict.
|
||||
if processed_messages is not None and (
|
||||
@@ -592,7 +591,6 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
|
||||
is_multimodal = self.tokenizer_manager.model_config.is_multimodal
|
||||
processed_messages = self._process_messages(chat_request, is_multimodal)
|
||||
processed_messages.skip_special_tokens = chat_request.skip_special_tokens
|
||||
|
||||
if is_multimodal:
|
||||
request_prompts = [processed_messages.prompt]
|
||||
@@ -808,6 +806,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
*,
|
||||
require_reasoning: bool,
|
||||
):
|
||||
chat_tools = self._response_tools_to_chat_tools(request)
|
||||
if self.reasoning_parser:
|
||||
reasoning_parser = ReasoningParser(
|
||||
model_type=self.reasoning_parser,
|
||||
@@ -819,6 +818,11 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
),
|
||||
request=request,
|
||||
tokenizer=self.tokenizer_manager.tokenizer,
|
||||
tool_call_parser_active=bool(
|
||||
chat_tools
|
||||
and self.tool_call_parser
|
||||
and request.tool_choice != "none"
|
||||
),
|
||||
)
|
||||
reasoning_content, content = reasoning_parser.parse_non_stream(final_output)
|
||||
else:
|
||||
@@ -851,7 +855,6 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
)
|
||||
output_items.append(reasoning_item)
|
||||
|
||||
chat_tools = self._response_tools_to_chat_tools(request)
|
||||
is_required = request.tool_choice == "required"
|
||||
tool_call_items: list[ResponseFunctionToolCall] = []
|
||||
parsed_via_native = False
|
||||
@@ -1977,6 +1980,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
),
|
||||
request=request,
|
||||
tokenizer=self.tokenizer_manager.tokenizer,
|
||||
tool_call_parser_active=isinstance(tool_parser, FunctionCallParser),
|
||||
)
|
||||
|
||||
current_output_index = -1
|
||||
@@ -2003,6 +2007,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
total_tokens_meta = 0
|
||||
reasoning_tokens_meta = 0
|
||||
finish_reason: Optional[dict[str, Any]] = None
|
||||
flushed = False
|
||||
stream_offset = 0
|
||||
incremental = self.tokenizer_manager.server_args.incremental_streaming_output
|
||||
|
||||
@@ -2210,11 +2215,26 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
stream_offset = len(text)
|
||||
if not delta and finish_reason is None:
|
||||
continue
|
||||
# finish_reason is sticky, so it would otherwise re-flush.
|
||||
flush = (
|
||||
not flushed
|
||||
and finish_reason is not None
|
||||
and finish_reason.get("type") != "abort"
|
||||
)
|
||||
flushed = flushed or flush
|
||||
|
||||
if reasoning_parser_obj is not None:
|
||||
reasoning_chunk, delta = reasoning_parser_obj.parse_stream_chunk(
|
||||
delta
|
||||
)
|
||||
if flush:
|
||||
end_reasoning, end_normal = (
|
||||
reasoning_parser_obj.parse_stream_end()
|
||||
)
|
||||
if end_reasoning:
|
||||
reasoning_chunk = (reasoning_chunk or "") + end_reasoning
|
||||
if end_normal:
|
||||
delta = (delta or "") + end_normal
|
||||
else:
|
||||
reasoning_chunk = None
|
||||
|
||||
@@ -2278,7 +2298,7 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
)
|
||||
)
|
||||
|
||||
if not delta:
|
||||
if not delta and not flush:
|
||||
continue
|
||||
|
||||
if isinstance(tool_parser, JsonArrayParser):
|
||||
@@ -2286,6 +2306,10 @@ class OpenAIServingResponses(OpenAIServingChat):
|
||||
normal_text, tool_calls = sp.normal_text or "", sp.calls
|
||||
elif tool_parser is not None:
|
||||
normal_text, tool_calls = tool_parser.parse_stream_chunk(delta)
|
||||
if flush:
|
||||
end_text, end_calls = tool_parser.parse_stream_end()
|
||||
normal_text = (normal_text or "") + end_text
|
||||
tool_calls = list(tool_calls) + end_calls
|
||||
else:
|
||||
normal_text, tool_calls = delta, []
|
||||
|
||||
|
||||
@@ -652,6 +652,8 @@ class Envs:
|
||||
# Number of decode steps between periodic mx.clear_cache() calls.
|
||||
# Set to 0 to disable cache clearing entirely.
|
||||
SGLANG_MLX_CLEAR_CACHE_STEPS = EnvInt(256)
|
||||
# MLX buffer-cache cap in GB.
|
||||
SGLANG_MLX_CACHE_LIMIT_GB = EnvFloat(None)
|
||||
|
||||
# NPU
|
||||
SGLANG_NPU_DISABLE_ACL_FORMAT_WEIGHT = EnvBool(False)
|
||||
|
||||
@@ -350,6 +350,14 @@ class BaseFormatDetector(ABC):
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
def finish(self, tools: List[Tool]) -> StreamingParseResult:
|
||||
"""Called once when the stream ends; flush any buffered state.
|
||||
|
||||
Detectors that hold text back while waiting for a marker that can no
|
||||
longer arrive (the stream is over) override this to release it.
|
||||
"""
|
||||
return StreamingParseResult()
|
||||
|
||||
def supports_structural_tag(self) -> bool:
|
||||
"""Return True if this detector supports structural tag format."""
|
||||
return True
|
||||
|
||||
@@ -37,6 +37,7 @@ from sglang.srt.function_call.minicpm5_detector import MiniCPM5Detector
|
||||
from sglang.srt.function_call.minimax_m2 import MinimaxM2Detector
|
||||
from sglang.srt.function_call.minimax_m3 import MinimaxM3Detector
|
||||
from sglang.srt.function_call.mistral_detector import MistralDetector
|
||||
from sglang.srt.function_call.muse_glimmer_detector import MuseGlimmerDetector
|
||||
from sglang.srt.function_call.poolside_v1_detector import PoolsideV1Detector
|
||||
from sglang.srt.function_call.pythonic_detector import PythonicDetector
|
||||
from sglang.srt.function_call.qwen3_coder_detector import Qwen3CoderDetector
|
||||
@@ -78,6 +79,7 @@ class FunctionCallParser:
|
||||
"mimo": MiMoDetector,
|
||||
"minicpm5": MiniCPM5Detector,
|
||||
"mistral": MistralDetector,
|
||||
"muse": MuseGlimmerDetector,
|
||||
"poolside_v1": PoolsideV1Detector,
|
||||
"pythonic": PythonicDetector,
|
||||
"qwen": Qwen25Detector,
|
||||
@@ -175,6 +177,17 @@ class FunctionCallParser:
|
||||
|
||||
return final_normal_text, final_calls
|
||||
|
||||
def parse_stream_end(self) -> Tuple[str, list[ToolCallItem]]:
|
||||
"""Flush detector state once the stream ends.
|
||||
|
||||
Text a detector held back waiting for a marker (which can no longer
|
||||
arrive) is released as normal text; see BaseFormatDetector.finish().
|
||||
"""
|
||||
if not self.tools:
|
||||
return "", []
|
||||
sp_result = self.detector.finish(self.tools)
|
||||
return sp_result.normal_text, sp_result.calls
|
||||
|
||||
def get_legacy_structural_tag(
|
||||
self, at_least_one: bool = False
|
||||
) -> StructuralTagResponseFormat:
|
||||
|
||||
@@ -0,0 +1,259 @@
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from typing import Dict, List, Optional, Set
|
||||
|
||||
from sglang.srt.entrypoints.openai.protocol import Tool
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.function_call.base_format_detector import BaseFormatDetector
|
||||
from sglang.srt.function_call.core_types import (
|
||||
StreamingParseResult,
|
||||
StructureInfo,
|
||||
ToolCallItem,
|
||||
_GetInfoFunc,
|
||||
)
|
||||
from sglang.srt.function_call.muse_glimmer_format import (
|
||||
EOM,
|
||||
EOT,
|
||||
FUNCTION_CALLS_CLOSE,
|
||||
FUNCTION_CALLS_OPEN,
|
||||
INVOKE_CLOSE,
|
||||
INVOKE_OPEN,
|
||||
MAX_MARKER,
|
||||
MESSAGE,
|
||||
RECIPIENT_RE,
|
||||
START,
|
||||
could_start_header,
|
||||
has_atem_markers,
|
||||
partial_marker_len,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_INVOKE_OPEN_RE = re.compile(r'<atem:invoke\b[^>]*?\bname="(?P<name>[^"]+)"[^>]*?>')
|
||||
_PARAM_RE = re.compile(
|
||||
r'<atem:parameter\b[^>]*?\bname="(?P<key>[^"]+)"[^>]*?>(?P<value>.*?)'
|
||||
r"</atem:parameter>",
|
||||
re.DOTALL,
|
||||
)
|
||||
|
||||
# Recipients whose bodies are prose, never tool calls.
|
||||
_NON_TOOL_RECIPIENTS = frozenset({"self", "user"})
|
||||
|
||||
|
||||
def _is_tool_channel(recipient: Optional[str]) -> bool:
|
||||
"""True when this channel routes to a tool."""
|
||||
return recipient is not None and recipient not in _NON_TOOL_RECIPIENTS
|
||||
|
||||
|
||||
def _decode_value(raw: str):
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return raw
|
||||
|
||||
|
||||
def _normalize_name(emitted: str, registered: Set[str]) -> str:
|
||||
"""Strip the chat template's doubled namespace."""
|
||||
if not registered or emitted in registered:
|
||||
return emitted
|
||||
if "." not in emitted:
|
||||
return emitted
|
||||
head, _, tail = emitted.partition(".")
|
||||
if head == tail and head in registered:
|
||||
return head
|
||||
leaf = emitted.rsplit(".", 1)[-1]
|
||||
matches = [n for n in registered if n.rsplit(".", 1)[-1] == leaf]
|
||||
if len(matches) == 1:
|
||||
return matches[0]
|
||||
return emitted
|
||||
|
||||
|
||||
class MuseGlimmerDetector(BaseFormatDetector):
|
||||
"""Format detector for Muse Glimmer's ATEM tool-call blocks."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
# Streaming channel state.
|
||||
self._recipient: Optional[str] = None
|
||||
self._in_body = False
|
||||
self._at_stream_start = True
|
||||
# Name of the invoke whose arguments are still arriving, if any.
|
||||
self._open_invoke: Optional[str] = None
|
||||
|
||||
def has_tool_call(self, text: str) -> bool:
|
||||
return has_atem_markers(text)
|
||||
|
||||
def _registered_names(self, tools: Optional[List[Tool]]) -> Set[str]:
|
||||
return {t.function.name for t in tools or [] if t.function and t.function.name}
|
||||
|
||||
def _emit_call(
|
||||
self, name: str, args: Dict, registered: Set[str]
|
||||
) -> Optional[ToolCallItem]:
|
||||
"""Build one ToolCallItem, honoring the unknown-tool policy."""
|
||||
name = _normalize_name(name, registered)
|
||||
if name not in registered:
|
||||
logger.warning("Model attempted to call undefined function: %s", name)
|
||||
if not envs.SGLANG_FORWARD_UNKNOWN_TOOLS.get():
|
||||
return None
|
||||
self.current_tool_id += 1
|
||||
parameters = json.dumps(args, ensure_ascii=False)
|
||||
self.prev_tool_call_arr.append({"name": name, "arguments": args})
|
||||
self.streamed_args_for_tool.append(parameters)
|
||||
return ToolCallItem(
|
||||
tool_index=self.current_tool_id, name=name, parameters=parameters
|
||||
)
|
||||
|
||||
def detect_and_parse(self, text: str, tools: List[Tool]) -> StreamingParseResult:
|
||||
result = self.parse_streaming_increment(text, tools)
|
||||
end = self.finish(tools)
|
||||
return StreamingParseResult(
|
||||
normal_text=result.normal_text + end.normal_text,
|
||||
calls=result.calls + end.calls,
|
||||
)
|
||||
|
||||
def finish(self, tools: List[Tool]) -> StreamingParseResult:
|
||||
registered = self._registered_names(tools)
|
||||
calls: List[ToolCallItem] = []
|
||||
normal_parts: List[str] = []
|
||||
|
||||
if self._buffer:
|
||||
if self._in_body:
|
||||
self._consume_body(
|
||||
self._buffer,
|
||||
registered,
|
||||
calls,
|
||||
normal_parts,
|
||||
final=True,
|
||||
)
|
||||
else:
|
||||
# Truncated header: no body ever arrived, keep it as text.
|
||||
normal_parts.append(self._buffer)
|
||||
self._buffer = ""
|
||||
|
||||
return StreamingParseResult(normal_text="".join(normal_parts), calls=calls)
|
||||
|
||||
def _held_back(self, text: str) -> int:
|
||||
markers = (MESSAGE, EOM, EOT, START, FUNCTION_CALLS_OPEN, INVOKE_OPEN)
|
||||
return partial_marker_len(text, markers, MAX_MARKER)
|
||||
|
||||
def parse_streaming_increment(
|
||||
self, new_text: str, tools: List[Tool]
|
||||
) -> StreamingParseResult:
|
||||
self._buffer += new_text
|
||||
registered = self._registered_names(tools)
|
||||
calls: List[ToolCallItem] = []
|
||||
normal_parts: List[str] = []
|
||||
|
||||
while self._buffer:
|
||||
if not self._in_body:
|
||||
# Resolve the channel header before anything can be emitted.
|
||||
if self._at_stream_start and could_start_header(self._buffer):
|
||||
pass
|
||||
else:
|
||||
ws = len(self._buffer) - len(self._buffer.lstrip())
|
||||
head = self._buffer[ws : ws + len(START)]
|
||||
if not START.startswith(head):
|
||||
# Unframed prose: no header is coming.
|
||||
self._in_body = True
|
||||
self._recipient = None
|
||||
self._at_stream_start = False
|
||||
continue
|
||||
if ws:
|
||||
normal_parts.append(self._buffer[:ws])
|
||||
self._buffer = self._buffer[ws:]
|
||||
if len(head) < len(START):
|
||||
break # Partial "<|start|>", wait for the rest.
|
||||
|
||||
idx = self._buffer.find(MESSAGE)
|
||||
if idx == -1:
|
||||
break
|
||||
header = self._buffer[:idx]
|
||||
m = RECIPIENT_RE.search(header)
|
||||
self._recipient = m.group(1) if m else "user"
|
||||
self._buffer = self._buffer[idx + len(MESSAGE) :]
|
||||
self._in_body = True
|
||||
self._at_stream_start = False
|
||||
continue
|
||||
|
||||
# Inside a body: find the terminator, if it has arrived.
|
||||
end_at, end_len = -1, 0
|
||||
for tok in (EOM, EOT):
|
||||
i = self._buffer.find(tok)
|
||||
if i != -1 and (end_at == -1 or i < end_at):
|
||||
end_at, end_len = i, len(tok)
|
||||
|
||||
if end_at == -1:
|
||||
keep = self._held_back(self._buffer)
|
||||
chunk = self._buffer[: len(self._buffer) - keep]
|
||||
if not chunk:
|
||||
break
|
||||
consumed = self._consume_body(
|
||||
chunk, registered, calls, normal_parts, final=False
|
||||
)
|
||||
if consumed == 0:
|
||||
break
|
||||
self._buffer = self._buffer[consumed:]
|
||||
continue
|
||||
|
||||
self._consume_body(
|
||||
self._buffer[:end_at],
|
||||
registered,
|
||||
calls,
|
||||
normal_parts,
|
||||
final=True,
|
||||
)
|
||||
self._buffer = self._buffer[end_at + end_len :]
|
||||
self._in_body = False
|
||||
self._recipient = None
|
||||
self._open_invoke = None
|
||||
|
||||
return StreamingParseResult(normal_text="".join(normal_parts), calls=calls)
|
||||
|
||||
def _consume_body(
|
||||
self,
|
||||
chunk: str,
|
||||
registered: Set[str],
|
||||
calls: List[ToolCallItem],
|
||||
normal_parts: List[str],
|
||||
final: bool,
|
||||
) -> int:
|
||||
if not _is_tool_channel(self._recipient):
|
||||
normal_parts.append(chunk)
|
||||
return len(chunk)
|
||||
|
||||
pos = 0
|
||||
while pos < len(chunk):
|
||||
if self._open_invoke is None:
|
||||
m = _INVOKE_OPEN_RE.search(chunk, pos)
|
||||
if m is None:
|
||||
return pos if not final else len(chunk)
|
||||
self._open_invoke = m.group("name")
|
||||
pos = m.end()
|
||||
continue
|
||||
|
||||
close_at = chunk.find(INVOKE_CLOSE, pos)
|
||||
if close_at == -1:
|
||||
return pos if not final else len(chunk)
|
||||
body = chunk[pos:close_at]
|
||||
args = {
|
||||
pm.group("key"): _decode_value(pm.group("value"))
|
||||
for pm in _PARAM_RE.finditer(body)
|
||||
}
|
||||
item = self._emit_call(self._open_invoke, args, registered)
|
||||
if item is not None:
|
||||
calls.append(item)
|
||||
self._open_invoke = None
|
||||
pos = close_at + len(INVOKE_CLOSE)
|
||||
|
||||
return len(chunk)
|
||||
|
||||
def supports_structural_tag(self) -> bool:
|
||||
return False
|
||||
|
||||
def structure_info(self) -> _GetInfoFunc:
|
||||
return lambda name: StructureInfo(
|
||||
begin=f'{FUNCTION_CALLS_OPEN}\n{INVOKE_OPEN} name="{name}">',
|
||||
end=f"{INVOKE_CLOSE}\n{FUNCTION_CALLS_CLOSE}",
|
||||
trigger=FUNCTION_CALLS_OPEN,
|
||||
)
|
||||
@@ -0,0 +1,54 @@
|
||||
"""Muse Glimmer wire format, shared by its reasoning and function-call detectors."""
|
||||
|
||||
import re
|
||||
from typing import Sequence
|
||||
|
||||
# Channel framing.
|
||||
MESSAGE = "<|message|>"
|
||||
EOM = "<|eom|>"
|
||||
EOT = "<|eot|>"
|
||||
START = "<|start|>"
|
||||
|
||||
# ATEM payload markers.
|
||||
FUNCTION_CALLS_OPEN = "<atem:function_calls>"
|
||||
FUNCTION_CALLS_CLOSE = "</atem:function_calls>"
|
||||
INVOKE_OPEN = "<atem:invoke"
|
||||
INVOKE_CLOSE = "</atem:invoke>"
|
||||
|
||||
RECIPIENT_RE = re.compile(r"to=([^\s<]+)")
|
||||
|
||||
# Longest marker that could straddle a chunk boundary while streaming.
|
||||
MAX_CHANNEL_MARKER = max(len(m) for m in (MESSAGE, EOM, EOT, START))
|
||||
MAX_MARKER = max(MAX_CHANNEL_MARKER, len(FUNCTION_CALLS_OPEN))
|
||||
|
||||
|
||||
def could_start_header(text: str) -> bool:
|
||||
"""Whether the tail could still grow into a header."""
|
||||
stripped = text.lstrip()
|
||||
if not stripped:
|
||||
return True
|
||||
if not (stripped.startswith("to=") or "to=".startswith(stripped[:3])):
|
||||
return False
|
||||
if MESSAGE in stripped:
|
||||
return True
|
||||
recipient, angle, marker = stripped[3:].partition("<")
|
||||
if any(c.isspace() for c in recipient):
|
||||
return False
|
||||
return not angle or MESSAGE.startswith("<" + marker)
|
||||
|
||||
|
||||
def has_atem_markers(text: str) -> bool:
|
||||
return INVOKE_OPEN in text or FUNCTION_CALLS_OPEN in text
|
||||
|
||||
|
||||
def partial_marker_len(text: str, markers: Sequence[str], max_len: int) -> int:
|
||||
"""Length of the longest suffix of ``text`` that could still become a marker.
|
||||
|
||||
Returns 0 when nothing is held back, so ordinary text streams out immediately
|
||||
instead of waiting for a terminator that may never arrive.
|
||||
"""
|
||||
for k in range(min(len(text), max_len - 1), 0, -1):
|
||||
tail = text[-k:]
|
||||
if any(m.startswith(tail) for m in markers):
|
||||
return k
|
||||
return 0
|
||||
@@ -58,6 +58,10 @@ from sglang.srt.hardware_backend.mlx.kv_cache import (
|
||||
set_context,
|
||||
uses_sliding_window_attention,
|
||||
)
|
||||
from sglang.srt.hardware_backend.mlx.remote_code_gate import (
|
||||
ensure_remote_code_allowed,
|
||||
resolve_model_directory,
|
||||
)
|
||||
from sglang.srt.hardware_backend.mlx.sampling import (
|
||||
GREEDY_PARAMS,
|
||||
MlxLazyLogprobs,
|
||||
@@ -166,12 +170,14 @@ class MlxModelRunner:
|
||||
pool_size: int | None = None,
|
||||
mem_fraction_static: float = 0.8,
|
||||
quantization: str | None = None,
|
||||
revision: str | None = None,
|
||||
enable_sampling: bool = False,
|
||||
sampling_rng_seed: int = 0,
|
||||
deterministic_seeding: bool = False,
|
||||
):
|
||||
self.model_path = model_path
|
||||
self.trust_remote_code = trust_remote_code
|
||||
self.revision = revision
|
||||
self.model = None
|
||||
self.disable_radix_cache = disable_radix_cache
|
||||
self._mem_fraction_static = mem_fraction_static
|
||||
@@ -196,6 +202,19 @@ class MlxModelRunner:
|
||||
# modules directly.
|
||||
self._quantization: str | None = quantization
|
||||
|
||||
# Optionally cap the buffer cache (recycled GPU buffers). MLX never
|
||||
# returns freed buffers to the OS, so without a cap the process
|
||||
# footprint ratchets up to the worst transient — which is model
|
||||
# load/quantization itself, so the cap must be in place before it.
|
||||
cache_limit_gb = envs.SGLANG_MLX_CACHE_LIMIT_GB.get()
|
||||
if cache_limit_gb is not None:
|
||||
if cache_limit_gb < 0:
|
||||
raise ValueError(
|
||||
f"SGLANG_MLX_CACHE_LIMIT_GB must be >= 0, got {cache_limit_gb}"
|
||||
)
|
||||
mx.set_cache_limit(int(cache_limit_gb * (1024**3)))
|
||||
logger.info(f"MLX buffer cache limit set to {cache_limit_gb:.1f} GB")
|
||||
|
||||
self._load_model()
|
||||
|
||||
# Pin MLX allocations to prevent OS paging
|
||||
@@ -481,10 +500,18 @@ class MlxModelRunner:
|
||||
logger.info(f"Loading MLX model: {self.model_path}")
|
||||
start_time = time.time()
|
||||
|
||||
# Resolve the checkpoint directory once and inspect that exact
|
||||
# directory before mlx-lm can execute any checkpoint-shipped
|
||||
# model_file; the same directory is then handed to mlx_lm_load
|
||||
# (identity resolution for local dirs), so the inspected and
|
||||
# executed snapshots cannot diverge.
|
||||
model_dir = resolve_model_directory(self.model_path, revision=self.revision)
|
||||
ensure_remote_code_allowed(model_dir, self.trust_remote_code)
|
||||
|
||||
# We need the config dict to pass into quantize_model so it knows tied/embedding
|
||||
# layout. return_config=True is cheap and ignored when no quantization is requested.
|
||||
loaded = mlx_lm_load(
|
||||
self.model_path,
|
||||
str(model_dir),
|
||||
tokenizer_config={"trust_remote_code": self.trust_remote_code},
|
||||
return_config=True,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,740 @@
|
||||
# Copyright 2023-2026 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Muse Glimmer (dense, text-only) for mlx-lm.
|
||||
|
||||
Loaded via mlx-lm's custom-architecture path: ship this file in the checkpoint
|
||||
directory as ``muse_glimmer_mlx.py``, set ``"model_file": "muse_glimmer_mlx.py"`` in
|
||||
``config.json``. This copy under ``sglang/srt/hardware_backend/mlx/models/``
|
||||
is the maintained source; artifacts ship a byte-identical copy. It must stay
|
||||
importable standalone (mlx / mlx-lm imports only — no sglang imports),
|
||||
because mlx-lm executes it from the checkpoint directory.
|
||||
|
||||
Ported from the vendor reference implementation (and cross-checked against the
|
||||
SGLang CUDA port in ``python/sglang/srt/models/muse_glimmer.py``). Deviations from a
|
||||
llama-style decoder, and how each is mapped:
|
||||
|
||||
* **Sandwich norms.** ``h = h + post_norm(branch(pre_norm(h)))`` — the post
|
||||
norm applies to the branch output. Pre-norms use ``rms_norm_eps`` (1e-5),
|
||||
post-norms ``post_norm_eps`` (1e-8).
|
||||
* **Norm weights are offsets from 1.0.** The four per-layer norms compute
|
||||
``rms_norm(x, weight + 1.0)``; ``sanitize`` folds the +1 in at load time so
|
||||
plain ``nn.RMSNorm`` is exact. The final ``model.norm`` uses its weight
|
||||
directly and is NOT offset.
|
||||
* **Non-parametric QK-norm** over ``head_dim`` (no learnable scale), applied
|
||||
BEFORE RoPE. Exposed as ``q_norm``/``k_norm`` so the SGLang MLX batched
|
||||
decode wrapper applies them at the same point.
|
||||
* **Folded attention scale.** The reference multiplies q by
|
||||
``qk_scale_factor / sqrt(head_dim)`` after the QK-norm and SDPA then applies
|
||||
its default ``1/sqrt(head_dim)``. RoPE is orthogonal and softmax scale is
|
||||
linear in q, so both fold into ``scale = qk_scale_factor / head_dim``.
|
||||
* **Attention output gate.** ``sigmoid(output_gate_proj(pre_normed_x))`` is
|
||||
applied elementwise to the attention output before ``o_proj``. The gate
|
||||
reads the same input as ``q_proj``, so ``sanitize`` fuses it into ``q_proj``
|
||||
per-head-interleaved (``[q_head; gate_head]``) — the exact layout the SGLang
|
||||
``MLXAttentionWrapper`` gate path splits back out during batched decode.
|
||||
* **iRoPE.** ``no_rope_layers[i] == 0`` marks NoPE layers (also the
|
||||
``full_attention`` layers); they get a ``NoPE`` identity that still
|
||||
satisfies the wrapper's ``rope(x, offset=...)`` call. RoPE layers use the
|
||||
interleaved GPT-J convention (``nn.RoPE(traditional=True)``), which the AOT
|
||||
Metal RoPE kernel does not support — Muse Glimmer always takes the
|
||||
``mx.fast.rope`` fallback.
|
||||
* **Sliding window.** ``layer_types`` marks the non-NoPE layers as
|
||||
``sliding_attention`` (window 2048, including the query position — the same
|
||||
band as HF's ``create_sliding_window_causal_mask``, so no off-by-one). The
|
||||
container exposes ``layer_types`` + ``sliding_window`` per the gpt-oss
|
||||
convention that both mlx-lm and the SGLang MLX backend read; windowing is
|
||||
done by banded masks over full-history KV, never a per-module
|
||||
``is_sliding`` flag.
|
||||
* **Full-history caches.** ``make_cache`` returns a plain ``KVCache`` for
|
||||
every layer, including sliding ones (unlike gpt-oss's ``RotatingKVCache``).
|
||||
Banded masks provide the window semantics; keeping full history makes
|
||||
greedy output exactly reproducible across prefill chunkings and matches how
|
||||
the SGLang MLX KV pool stores history.
|
||||
* **Embedding norm** (scaleless RMS) when ``normalize_tok_embeddings``.
|
||||
* **Logit head.** ``cap * tanh(lm_head(h) * output_multiplier / cap)``,
|
||||
computed in float32 like the reference (``None`` cap leaves just the
|
||||
multiplier).
|
||||
|
||||
Checkpoint formats. ``sanitize`` accepts exactly three weight layouts and
|
||||
rejects everything else with an actionable error:
|
||||
|
||||
* **Raw HF export** (output of the vendor's HF converter): carries
|
||||
``output_gate_proj`` and offset-form norm weights. Recognized by the
|
||||
complete raw key schema; transformed on load.
|
||||
* **RC multimodal export** (``transformers >= 5.15``
|
||||
vendor schema): text weights under ``model.language_model.`` with
|
||||
HF-canonical names. Normalized to the raw schema first (see the rename
|
||||
table at ``_RC_SUFFIX_RENAMES`` — the norm renames are POSITIONAL: the
|
||||
RC ``post_attention_layernorm`` is the post-attn sandwich norm, i.e. the
|
||||
raw ``post_attn_norm``, while the raw ``post_attention_layernorm`` is the
|
||||
pre-MLP norm, i.e. the RC ``pre_feedforward_layernorm``), vision tower /
|
||||
adapter / projection dropped, then transformed like a raw export. Two
|
||||
converter generations ship this layout: the older one bakes the scaleless
|
||||
embedding RMS-norm into ``embed_tokens.weight`` and keeps q/k in the native
|
||||
interleaved layout; the current one ships the raw table and permutes q/k
|
||||
into the NeoX rotary layout. The RC path reads the current
|
||||
conventions (``rope_is_neox_style`` pinned True) and leaves
|
||||
``normalize_tok_embeddings`` at its default True — the norm is
|
||||
idempotent on a baked table, so always-on covers both generations,
|
||||
but an older-generation export served through this path gets the wrong
|
||||
rope layout (their configs are byte-identical; prefer repackaging).
|
||||
* **Packaged MLX artifact**: already fused/folded, marked by
|
||||
``"muse_glimmer_mlx_format": 1`` in ``config.json`` (stamped at packaging time
|
||||
only, never present on raw HF exports). Passed through untouched.
|
||||
|
||||
Config schemas. ``ModelArgs.from_dict`` accepts the flat schema written at
|
||||
packaging time and the RC nested schema (``text_config`` present).
|
||||
The RC schema differs in two conventions beyond field names:
|
||||
``qk_scale_factor`` is expressed against SDPA's standard ``1/sqrt(head_dim)``
|
||||
(flat-schema value = RC value * sqrt(head_dim); both fold to the same
|
||||
``scale = flat_qk_scale / head_dim``), and NoPE layers are marked by zeros
|
||||
in ``layer_rope_theta`` rather than ``no_rope_layers``.
|
||||
"""
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, List, Optional
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
from mlx_lm.models.base import (
|
||||
BaseModelArgs,
|
||||
create_attention_mask,
|
||||
scaled_dot_product_attention,
|
||||
)
|
||||
from mlx_lm.models.cache import KVCache
|
||||
|
||||
# Version of the packaged (fused/folded) weight layout this file understands.
|
||||
MUSE_GLIMMER_MLX_FORMAT_VERSION = 1
|
||||
|
||||
# The four per-layer norms whose checkpoint weight is an offset from 1.0.
|
||||
# model.norm (MuseGlimmerFinalRMSNorm) is NOT in this list and must not be offset.
|
||||
_OFFSET_NORM_SUFFIXES = (
|
||||
"input_layernorm.weight",
|
||||
"post_attn_norm.weight",
|
||||
"post_attention_layernorm.weight",
|
||||
"post_ffn_norm.weight",
|
||||
)
|
||||
|
||||
# Text-only port: the vision tower/projector are not built.
|
||||
_VISION_KEY_MARKERS = (
|
||||
"vision_encoder",
|
||||
"vision_adapter",
|
||||
"vision_projection",
|
||||
"vision_tower",
|
||||
"perception_emb_norm",
|
||||
)
|
||||
|
||||
# Keys that only appear in a raw HF export, never in a packaged artifact.
|
||||
# "language_model" catches RC-layout strays (text weights live under
|
||||
# model.language_model. there).
|
||||
_RAW_ONLY_KEY_MARKERS = (
|
||||
"output_gate_proj",
|
||||
"rotary_emb",
|
||||
"language_model",
|
||||
) + _VISION_KEY_MARKERS
|
||||
|
||||
# RC (transformers >= 5.15 vendor schema) -> raw-schema key renames, applied
|
||||
# per key after stripping the "model.language_model." prefix. The norm
|
||||
# renames are positional, not textual: RC's post_attention_layernorm is the
|
||||
# post-attn sandwich norm (raw post_attn_norm, eps=post_norm_eps) and RC's
|
||||
# pre_feedforward_layernorm is the pre-MLP norm (raw post_attention_layernorm,
|
||||
# eps=rms_norm_eps). self_attn.gate_proj is the attention output gate
|
||||
# (mlp.gate_proj is untouched: the suffixes below carry the self_attn./
|
||||
# module context).
|
||||
_RC_SUFFIX_RENAMES = (
|
||||
("self_attn.gate_proj.weight", "self_attn.output_gate_proj.weight"),
|
||||
("post_attention_layernorm.weight", "post_attn_norm.weight"),
|
||||
("pre_feedforward_layernorm.weight", "post_attention_layernorm.weight"),
|
||||
("post_feedforward_layernorm.weight", "post_ffn_norm.weight"),
|
||||
)
|
||||
|
||||
_RC_PREFIX = "model.language_model."
|
||||
|
||||
|
||||
def flatten_rc_config(config: dict) -> dict:
|
||||
"""Translate the RC nested config schema into this file's flat schema.
|
||||
|
||||
Field mapping plus three convention conversions (see module docstring):
|
||||
qk_scale_factor gains the sqrt(head_dim) that the RC schema leaves to
|
||||
SDPA, NoPE layers come from zeros in layer_rope_theta, and the vendor
|
||||
export permutes q/k into the NeoX rotary layout (``_permute_for_rope``)
|
||||
so rope_is_neox_style is pinned True -- ``nn.RoPE(traditional=True)``
|
||||
on those weights emits garbled text rather than raising.
|
||||
|
||||
normalize_tok_embeddings is left at its default. Older vendor exports baked
|
||||
the embedding norm into embed_tokens.weight and needed it off; the current
|
||||
export ships the native table instead.
|
||||
"""
|
||||
text = config["text_config"]
|
||||
|
||||
activation = text.get("hidden_activation", "silu")
|
||||
if activation != "silu":
|
||||
raise ValueError(
|
||||
f"RC config has hidden_activation={activation!r}; this port "
|
||||
"hardcodes silu"
|
||||
)
|
||||
|
||||
head_dim = int(text.get("head_dim", 128))
|
||||
rope_params = text.get("rope_parameters") or {}
|
||||
layer_rope_theta = text.get("layer_rope_theta")
|
||||
|
||||
flat = {
|
||||
"model_type": "muse_glimmer",
|
||||
"hidden_size": text["hidden_size"],
|
||||
"num_hidden_layers": text["num_hidden_layers"],
|
||||
"num_attention_heads": text["num_attention_heads"],
|
||||
"num_key_value_heads": text["num_key_value_heads"],
|
||||
"head_dim": head_dim,
|
||||
"intermediate_size": text["intermediate_size"],
|
||||
"vocab_size": text["vocab_size"],
|
||||
"rms_norm_eps": text["rms_norm_eps"],
|
||||
"post_norm_eps": text["post_norm_eps"],
|
||||
"rope_theta": rope_params.get("rope_theta", text.get("rope_theta", 500_000.0)),
|
||||
"max_position_embeddings": text["max_position_embeddings"],
|
||||
"qk_scale_factor": text["qk_scale_factor"] * math.sqrt(head_dim),
|
||||
"output_multiplier": text["output_multiplier"],
|
||||
"output_soft_cap_temp": text.get("final_logit_softcapping"),
|
||||
"rope_is_neox_style": True,
|
||||
"sliding_window": text["sliding_window"],
|
||||
}
|
||||
if "layer_types" in text:
|
||||
flat["layer_types"] = list(text["layer_types"])
|
||||
if layer_rope_theta is not None:
|
||||
flat["no_rope_layers"] = [0 if not theta else 1 for theta in layer_rope_theta]
|
||||
return flat
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelArgs(BaseModelArgs):
|
||||
model_type: str = "muse_glimmer"
|
||||
hidden_size: int = 6656
|
||||
num_hidden_layers: int = 52
|
||||
num_attention_heads: int = 32
|
||||
num_key_value_heads: int = 2
|
||||
head_dim: int = 128
|
||||
intermediate_size: int = 19968
|
||||
vocab_size: int = 202048
|
||||
rms_norm_eps: float = 1e-5
|
||||
post_norm_eps: float = 1e-8
|
||||
rope_theta: float = 500_000.0
|
||||
max_position_embeddings: int = 16384
|
||||
use_qk_norm: bool = True
|
||||
qk_scale_factor: float = 43.7840518911
|
||||
use_attn_output_gate: bool = True
|
||||
output_multiplier: float = 0.19611613513818404
|
||||
output_soft_cap_temp: Optional[float] = 20.0
|
||||
rope_is_neox_style: bool = False
|
||||
normalize_tok_embeddings: bool = True
|
||||
sliding_window: int = 2048
|
||||
every_n_layers_nope: int = 4
|
||||
no_rope_layers: Optional[List[int]] = None
|
||||
layer_types: Optional[List[str]] = None
|
||||
# Set on saved MLX artifacts at packaging time (never on raw HF
|
||||
# exports); tells sanitize() the weights are already fused/folded.
|
||||
muse_glimmer_mlx_format: Optional[int] = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, params):
|
||||
# RC multimodal schema: text fields nested under text_config, with
|
||||
# convention differences handled by flatten_rc_config.
|
||||
if "text_config" in params:
|
||||
params = flatten_rc_config(params)
|
||||
return super().from_dict(params)
|
||||
|
||||
def __post_init__(self):
|
||||
# Mirror the vendor config's derivations so a config.json that
|
||||
# omits the explicit lists still builds the right architecture.
|
||||
if self.every_n_layers_nope <= 0:
|
||||
raise ValueError(
|
||||
f"every_n_layers_nope must be positive, got {self.every_n_layers_nope}"
|
||||
)
|
||||
if self.num_attention_heads % self.num_key_value_heads != 0:
|
||||
raise ValueError(
|
||||
f"num_attention_heads ({self.num_attention_heads}) must be a "
|
||||
f"multiple of num_key_value_heads ({self.num_key_value_heads})"
|
||||
)
|
||||
|
||||
derived_no_rope = [
|
||||
0 if (self.num_hidden_layers - i - 1) % self.every_n_layers_nope == 0 else 1
|
||||
for i in range(self.num_hidden_layers)
|
||||
]
|
||||
if self.no_rope_layers is None:
|
||||
self.no_rope_layers = derived_no_rope
|
||||
else:
|
||||
if len(self.no_rope_layers) != self.num_hidden_layers:
|
||||
raise ValueError(
|
||||
f"no_rope_layers has {len(self.no_rope_layers)} entries but "
|
||||
f"num_hidden_layers is {self.num_hidden_layers}"
|
||||
)
|
||||
bad_flags = sorted(set(self.no_rope_layers) - {0, 1})
|
||||
if bad_flags:
|
||||
raise ValueError(
|
||||
f"no_rope_layers contains non-binary entries {bad_flags}; "
|
||||
"each entry must be 0 (NoPE) or 1 (RoPE)"
|
||||
)
|
||||
|
||||
# NoPE layers are the full-attention layers; the rest slide.
|
||||
derived_layer_types = [
|
||||
"full_attention" if rope_flag == 0 else "sliding_attention"
|
||||
for rope_flag in self.no_rope_layers
|
||||
]
|
||||
if self.layer_types is None:
|
||||
self.layer_types = derived_layer_types
|
||||
else:
|
||||
if len(self.layer_types) != self.num_hidden_layers:
|
||||
raise ValueError(
|
||||
f"layer_types has {len(self.layer_types)} entries but "
|
||||
f"num_hidden_layers is {self.num_hidden_layers}"
|
||||
)
|
||||
bad = sorted(
|
||||
set(self.layer_types) - {"full_attention", "sliding_attention"}
|
||||
)
|
||||
if bad:
|
||||
raise ValueError(
|
||||
f"layer_types contains unknown entries {bad}; expected only "
|
||||
"'full_attention' or 'sliding_attention'"
|
||||
)
|
||||
if self.layer_types != derived_layer_types:
|
||||
mismatches = [
|
||||
i
|
||||
for i, (got, want) in enumerate(
|
||||
zip(self.layer_types, derived_layer_types)
|
||||
)
|
||||
if got != want
|
||||
]
|
||||
raise ValueError(
|
||||
"layer_types disagrees with no_rope_layers (NoPE layers "
|
||||
"must be the full_attention layers) at layer indices "
|
||||
f"{mismatches}"
|
||||
)
|
||||
|
||||
if self.muse_glimmer_mlx_format is not None and (
|
||||
self.muse_glimmer_mlx_format != MUSE_GLIMMER_MLX_FORMAT_VERSION
|
||||
):
|
||||
raise ValueError(
|
||||
f"muse_glimmer_mlx_format {self.muse_glimmer_mlx_format} is not supported by "
|
||||
f"this model file (expected {MUSE_GLIMMER_MLX_FORMAT_VERSION}); "
|
||||
"regenerate the artifact with a matching packager"
|
||||
)
|
||||
|
||||
|
||||
class ScalelessRMSNorm(nn.Module):
|
||||
"""RMS norm with no learnable scale (reference MuseGlimmerScalelessRMSNorm)."""
|
||||
|
||||
def __init__(self, dims: int, eps: float):
|
||||
super().__init__()
|
||||
self.dims = dims
|
||||
self.eps = eps
|
||||
|
||||
def __call__(self, x: mx.array) -> mx.array:
|
||||
return mx.fast.rms_norm(x, None, self.eps)
|
||||
|
||||
|
||||
class NoPE(nn.Module):
|
||||
"""Identity standing in for RoPE on NoPE layers.
|
||||
|
||||
Accepts the ``offset`` kwarg so both this file's forward and the SGLang
|
||||
``MLXAttentionWrapper`` (which calls ``rope(x, offset=offsets)``
|
||||
unconditionally) can treat every layer uniformly. ``dims = 0`` keeps the
|
||||
AOT Metal RoPE kernel gating disabled for these layers.
|
||||
"""
|
||||
|
||||
dims = 0
|
||||
traditional = True
|
||||
|
||||
def __call__(self, x: mx.array, offset: Any = 0) -> mx.array:
|
||||
return x
|
||||
|
||||
|
||||
class MuseGlimmerAttention(nn.Module):
|
||||
def __init__(self, args: ModelArgs, layer_idx: int):
|
||||
super().__init__()
|
||||
self.num_attention_heads = args.num_attention_heads
|
||||
self.num_key_value_heads = args.num_key_value_heads
|
||||
self.head_dim = args.head_dim
|
||||
self.use_attn_output_gate = args.use_attn_output_gate
|
||||
|
||||
q_dim = args.num_attention_heads * args.head_dim
|
||||
kv_dim = args.num_key_value_heads * args.head_dim
|
||||
|
||||
# With the output gate, q_proj holds the per-head-interleaved
|
||||
# [q_head; gate_head] fusion produced by sanitize(): output width
|
||||
# 2 * q_dim, split back out in the forward pass.
|
||||
self.q_proj = nn.Linear(
|
||||
args.hidden_size,
|
||||
2 * q_dim if self.use_attn_output_gate else q_dim,
|
||||
bias=False,
|
||||
)
|
||||
self.k_proj = nn.Linear(args.hidden_size, kv_dim, bias=False)
|
||||
self.v_proj = nn.Linear(args.hidden_size, kv_dim, bias=False)
|
||||
self.o_proj = nn.Linear(q_dim, args.hidden_size, bias=False)
|
||||
|
||||
# Bool flag for the forward-pass branch; q_norm/k_norm stay ABSENT
|
||||
# (not None) when unused — the SGLang batched-decode wrapper
|
||||
# duck-types them via hasattr.
|
||||
self.use_qk_norm = args.use_qk_norm
|
||||
if args.use_qk_norm:
|
||||
self.q_norm = ScalelessRMSNorm(args.head_dim, args.rms_norm_eps)
|
||||
self.k_norm = ScalelessRMSNorm(args.head_dim, args.rms_norm_eps)
|
||||
# Reference: q *= qk_scale_factor / sqrt(head_dim) after the
|
||||
# QK-norm, then SDPA scales by 1/sqrt(head_dim); folded here.
|
||||
self.scale = args.qk_scale_factor / args.head_dim
|
||||
else:
|
||||
self.scale = args.head_dim**-0.5
|
||||
|
||||
use_rope = args.no_rope_layers[layer_idx] == 1
|
||||
self.rope = (
|
||||
nn.RoPE(
|
||||
args.head_dim,
|
||||
traditional=not args.rope_is_neox_style,
|
||||
base=args.rope_theta,
|
||||
)
|
||||
if use_rope
|
||||
else NoPE()
|
||||
)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
x: mx.array,
|
||||
mask: Optional[mx.array] = None,
|
||||
cache: Optional[Any] = None,
|
||||
) -> mx.array:
|
||||
B, L, _ = x.shape
|
||||
H, Hk, D = self.num_attention_heads, self.num_key_value_heads, self.head_dim
|
||||
|
||||
q = self.q_proj(x)
|
||||
gate = None
|
||||
if self.use_attn_output_gate:
|
||||
# Per-head layout [q_head; gate_head]: same split the SGLang MLX
|
||||
# batched-decode wrapper performs.
|
||||
q, gate = mx.split(q.reshape(B, L, H, 2 * D), 2, axis=-1)
|
||||
else:
|
||||
q = q.reshape(B, L, H, D)
|
||||
k = self.k_proj(x).reshape(B, L, Hk, D)
|
||||
v = self.v_proj(x).reshape(B, L, Hk, D)
|
||||
|
||||
# QK-norm BEFORE RoPE, matching the reference.
|
||||
if self.use_qk_norm:
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
q = q.transpose(0, 2, 1, 3)
|
||||
k = k.transpose(0, 2, 1, 3)
|
||||
v = v.transpose(0, 2, 1, 3)
|
||||
|
||||
if cache is not None:
|
||||
q = self.rope(q, offset=cache.offset)
|
||||
k = self.rope(k, offset=cache.offset)
|
||||
k, v = cache.update_and_fetch(k, v)
|
||||
else:
|
||||
q = self.rope(q)
|
||||
k = self.rope(k)
|
||||
|
||||
out = scaled_dot_product_attention(q, k, v, cache, scale=self.scale, mask=mask)
|
||||
|
||||
out = out.transpose(0, 2, 1, 3)
|
||||
if gate is not None:
|
||||
out = mx.sigmoid(gate) * out
|
||||
return self.o_proj(out.reshape(B, L, -1))
|
||||
|
||||
|
||||
class MuseGlimmerMLP(nn.Module):
|
||||
def __init__(self, args: ModelArgs):
|
||||
super().__init__()
|
||||
self.gate_proj = nn.Linear(args.hidden_size, args.intermediate_size, bias=False)
|
||||
self.up_proj = nn.Linear(args.hidden_size, args.intermediate_size, bias=False)
|
||||
self.down_proj = nn.Linear(args.intermediate_size, args.hidden_size, bias=False)
|
||||
|
||||
def __call__(self, x: mx.array) -> mx.array:
|
||||
return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x))
|
||||
|
||||
|
||||
class MuseGlimmerDecoderLayer(nn.Module):
|
||||
def __init__(self, args: ModelArgs, layer_idx: int):
|
||||
super().__init__()
|
||||
self.input_layernorm = nn.RMSNorm(args.hidden_size, args.rms_norm_eps)
|
||||
self.self_attn = MuseGlimmerAttention(args, layer_idx)
|
||||
self.post_attn_norm = nn.RMSNorm(args.hidden_size, args.post_norm_eps)
|
||||
self.post_attention_layernorm = nn.RMSNorm(args.hidden_size, args.rms_norm_eps)
|
||||
self.mlp = MuseGlimmerMLP(args)
|
||||
self.post_ffn_norm = nn.RMSNorm(args.hidden_size, args.post_norm_eps)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
x: mx.array,
|
||||
mask: Optional[mx.array] = None,
|
||||
cache: Optional[Any] = None,
|
||||
) -> mx.array:
|
||||
# Sandwich norms: the post-norm normalizes the branch output before
|
||||
# the residual add.
|
||||
x = x + self.post_attn_norm(
|
||||
self.self_attn(self.input_layernorm(x), mask, cache)
|
||||
)
|
||||
return x + self.post_ffn_norm(self.mlp(self.post_attention_layernorm(x)))
|
||||
|
||||
|
||||
class MuseGlimmerModel(nn.Module):
|
||||
def __init__(self, args: ModelArgs):
|
||||
super().__init__()
|
||||
self.embed_tokens = nn.Embedding(args.vocab_size, args.hidden_size)
|
||||
self.embed_norm = (
|
||||
ScalelessRMSNorm(args.hidden_size, args.rms_norm_eps)
|
||||
if args.normalize_tok_embeddings
|
||||
else None
|
||||
)
|
||||
self.layers = [
|
||||
MuseGlimmerDecoderLayer(args, i) for i in range(args.num_hidden_layers)
|
||||
]
|
||||
# Reference MuseGlimmerFinalRMSNorm: weight is the scale, not an offset.
|
||||
self.norm = nn.RMSNorm(args.hidden_size, args.rms_norm_eps)
|
||||
|
||||
# Container-level window declaration (gpt-oss convention), read by
|
||||
# both this forward and the SGLang MLX backend's
|
||||
# get_layer_window_sizes(); per-module ``is_sliding`` flags would
|
||||
# instead trip the backend's uniform-KV-pool check.
|
||||
self.layer_types = list(args.layer_types)
|
||||
self.sliding_window = args.sliding_window
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
inputs: mx.array,
|
||||
cache: Optional[Any] = None,
|
||||
input_embeddings: Optional[mx.array] = None,
|
||||
) -> mx.array:
|
||||
x = (
|
||||
input_embeddings
|
||||
if input_embeddings is not None
|
||||
else self.embed_tokens(inputs)
|
||||
)
|
||||
if self.embed_norm is not None:
|
||||
x = self.embed_norm(x)
|
||||
|
||||
if cache is None:
|
||||
cache = [None] * len(self.layers)
|
||||
|
||||
# One mask per layer type present, anchored to the first cache of
|
||||
# that type (all caches of a type share the same offset).
|
||||
masks = {}
|
||||
for layer_type in ("full_attention", "sliding_attention"):
|
||||
try:
|
||||
idx = self.layer_types.index(layer_type)
|
||||
except ValueError:
|
||||
continue
|
||||
window = self.sliding_window if layer_type == "sliding_attention" else None
|
||||
if window is not None:
|
||||
masks[layer_type] = create_attention_mask(
|
||||
x, cache[idx], window_size=window
|
||||
)
|
||||
else:
|
||||
masks[layer_type] = create_attention_mask(x, cache[idx])
|
||||
|
||||
for layer, c, layer_type in zip(self.layers, cache, self.layer_types):
|
||||
x = layer(x, masks[layer_type], c)
|
||||
|
||||
return self.norm(x)
|
||||
|
||||
|
||||
class Model(nn.Module):
|
||||
def __init__(self, args: ModelArgs):
|
||||
super().__init__()
|
||||
self.args = args
|
||||
self.model_type = args.model_type
|
||||
self.model = MuseGlimmerModel(args)
|
||||
self.lm_head = nn.Linear(args.hidden_size, args.vocab_size, bias=False)
|
||||
|
||||
@property
|
||||
def layers(self):
|
||||
return self.model.layers
|
||||
|
||||
def make_cache(self) -> List[Any]:
|
||||
# Full-history caches for every layer, sliding ones included: banded
|
||||
# masks provide the window, and full history keeps greedy output
|
||||
# exactly reproducible across prefill chunkings (a RotatingKVCache
|
||||
# would diverge once the prompt exceeds the window).
|
||||
return [KVCache() for _ in range(len(self.model.layers))]
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
inputs: mx.array,
|
||||
cache: Optional[Any] = None,
|
||||
input_embeddings: Optional[mx.array] = None,
|
||||
) -> mx.array:
|
||||
hidden = self.model(inputs, cache, input_embeddings)
|
||||
# Reference computes the logit head in float32.
|
||||
logits = self.lm_head(hidden).astype(mx.float32)
|
||||
if self.args.output_soft_cap_temp is not None:
|
||||
cap = self.args.output_soft_cap_temp
|
||||
logits = cap * mx.tanh(logits * self.args.output_multiplier / cap)
|
||||
else:
|
||||
logits = logits * self.args.output_multiplier
|
||||
return logits
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Weight loading
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _expected_raw_keys(self) -> set:
|
||||
"""The complete key schema of a raw HF export (text path only)."""
|
||||
keys = {"model.embed_tokens.weight", "model.norm.weight", "lm_head.weight"}
|
||||
for i in range(self.args.num_hidden_layers):
|
||||
prefix = f"model.layers.{i}."
|
||||
keys.update(
|
||||
prefix + suffix
|
||||
for suffix in (
|
||||
"self_attn.q_proj.weight",
|
||||
"self_attn.k_proj.weight",
|
||||
"self_attn.v_proj.weight",
|
||||
"self_attn.o_proj.weight",
|
||||
"input_layernorm.weight",
|
||||
"post_attn_norm.weight",
|
||||
"post_attention_layernorm.weight",
|
||||
"post_ffn_norm.weight",
|
||||
"mlp.gate_proj.weight",
|
||||
"mlp.up_proj.weight",
|
||||
"mlp.down_proj.weight",
|
||||
)
|
||||
)
|
||||
if self.args.use_attn_output_gate:
|
||||
keys.add(prefix + "self_attn.output_gate_proj.weight")
|
||||
return keys
|
||||
|
||||
def sanitize(self, weights: dict) -> dict:
|
||||
if self.args.muse_glimmer_mlx_format == MUSE_GLIMMER_MLX_FORMAT_VERSION:
|
||||
# Packaged artifact: weights are already fused/folded. A raw-only
|
||||
# key here means the marker was stamped on the wrong directory.
|
||||
stray = sorted(
|
||||
k
|
||||
for k in weights
|
||||
if any(marker in k for marker in _RAW_ONLY_KEY_MARKERS)
|
||||
)
|
||||
if stray:
|
||||
raise ValueError(
|
||||
"config.json claims a packaged Muse Glimmer MLX artifact "
|
||||
f"(muse_glimmer_mlx_format={MUSE_GLIMMER_MLX_FORMAT_VERSION}) but the "
|
||||
f"weights contain raw-checkpoint keys {stray[:4]}"
|
||||
f"{'...' if len(stray) > 4 else ''}; the marker belongs "
|
||||
"on packaged artifacts only — repackage from the raw HF export"
|
||||
)
|
||||
return weights
|
||||
|
||||
# No marker: a raw HF export, possibly in the RC multimodal layout —
|
||||
# normalize that to the raw schema first.
|
||||
if any(k.startswith(_RC_PREFIX) for k in weights):
|
||||
weights = _normalize_rc_layout(weights)
|
||||
|
||||
text_keys = {
|
||||
k
|
||||
for k in weights
|
||||
if not any(marker in k for marker in _VISION_KEY_MARKERS)
|
||||
and "rotary_emb" not in k
|
||||
}
|
||||
expected = self._expected_raw_keys()
|
||||
missing = sorted(expected - text_keys)
|
||||
unexpected = sorted(text_keys - expected)
|
||||
if missing or unexpected:
|
||||
hint = ""
|
||||
gate_missing = all("output_gate_proj" in k for k in missing) and missing
|
||||
if gate_missing and not unexpected:
|
||||
hint = (
|
||||
" (weights look already fused: if this is a packaged "
|
||||
'artifact, its config.json must carry "muse_glimmer_mlx_format": '
|
||||
f"{MUSE_GLIMMER_MLX_FORMAT_VERSION})"
|
||||
)
|
||||
raise ValueError(
|
||||
"not a complete raw Muse Glimmer HF checkpoint: "
|
||||
f"{len(missing)} missing keys {missing[:4]}"
|
||||
f"{'...' if len(missing) > 4 else ''}, "
|
||||
f"{len(unexpected)} unexpected keys {unexpected[:4]}"
|
||||
f"{'...' if len(unexpected) > 4 else ''}{hint}"
|
||||
)
|
||||
|
||||
H = self.args.num_attention_heads
|
||||
D = self.args.head_dim
|
||||
hidden = self.args.hidden_size
|
||||
|
||||
embed_shape = tuple(weights["model.embed_tokens.weight"].shape)
|
||||
if embed_shape != (self.args.vocab_size, hidden):
|
||||
raise ValueError(
|
||||
f"embed_tokens.weight has shape {embed_shape} but config says "
|
||||
f"(vocab_size, hidden_size) = ({self.args.vocab_size}, {hidden})"
|
||||
)
|
||||
raw_q_shape = tuple(weights["model.layers.0.self_attn.q_proj.weight"].shape)
|
||||
if raw_q_shape != (H * D, hidden):
|
||||
raise ValueError(
|
||||
f"raw q_proj.weight has shape {raw_q_shape}, expected "
|
||||
f"({H * D}, {hidden}); a width of {2 * H * D} means the gate "
|
||||
"is already fused — such artifacts must carry "
|
||||
f'"muse_glimmer_mlx_format": {MUSE_GLIMMER_MLX_FORMAT_VERSION} in config.json'
|
||||
)
|
||||
|
||||
new_weights = {}
|
||||
for name, w in weights.items():
|
||||
# mlx derives RoPE itself; drop cached buffers.
|
||||
if "rotary_emb" in name:
|
||||
continue
|
||||
if any(marker in name for marker in _VISION_KEY_MARKERS):
|
||||
continue
|
||||
# Consumed below when its q_proj comes up.
|
||||
if name.endswith("output_gate_proj.weight"):
|
||||
continue
|
||||
|
||||
# The reference computes rms_norm(x, weight + 1.0) for these four
|
||||
# norms; fold the +1 so plain nn.RMSNorm is exact. model.norm
|
||||
# (MuseGlimmerFinalRMSNorm) is deliberately not offset.
|
||||
if name.endswith(_OFFSET_NORM_SUFFIXES):
|
||||
w = w + 1.0
|
||||
|
||||
if name.endswith("q_proj.weight") and self.args.use_attn_output_gate:
|
||||
gate_name = name.replace("q_proj.weight", "output_gate_proj.weight")
|
||||
g = weights[gate_name]
|
||||
if tuple(g.shape) != (H * D, hidden):
|
||||
raise ValueError(
|
||||
f"{gate_name} has shape {tuple(g.shape)}, expected "
|
||||
f"({H * D}, {hidden})"
|
||||
)
|
||||
# Per-head interleave [q_head; gate_head]: (H*D, hidden) x2
|
||||
# -> (H, 2D, hidden) -> (2*H*D, hidden).
|
||||
w = mx.concatenate(
|
||||
[w.reshape(H, D, hidden), g.reshape(H, D, hidden)], axis=1
|
||||
).reshape(2 * H * D, hidden)
|
||||
|
||||
new_weights[name] = w
|
||||
|
||||
return new_weights
|
||||
|
||||
|
||||
def _normalize_rc_layout(weights: dict) -> dict:
|
||||
"""Rewrite RC multimodal keys to the raw text-only schema.
|
||||
|
||||
Drops the vision tower/adapter/projection, strips the
|
||||
``model.language_model.`` prefix, and applies the positional norm/gate
|
||||
renames from ``_RC_SUFFIX_RENAMES``. Suffix matching happens per key in
|
||||
one pass, so the post_attention_layernorm name swap cannot cascade.
|
||||
"""
|
||||
out = {}
|
||||
for name, w in weights.items():
|
||||
if any(marker in name for marker in _VISION_KEY_MARKERS):
|
||||
continue
|
||||
if name.startswith(_RC_PREFIX):
|
||||
name = "model." + name[len(_RC_PREFIX) :]
|
||||
for rc_suffix, raw_suffix in _RC_SUFFIX_RENAMES:
|
||||
if name.endswith(rc_suffix):
|
||||
name = name[: -len(rc_suffix)] + raw_suffix
|
||||
break
|
||||
out[name] = w
|
||||
return out
|
||||
|
||||
|
||||
EntryClass = Model
|
||||
@@ -0,0 +1,126 @@
|
||||
# Copyright 2023-2026 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Pre-execution gate for checkpoint-shipped model code on the MLX backend.
|
||||
|
||||
mlx-lm's loader executes ``config.json``'s ``model_file`` unconditionally:
|
||||
``mlx_lm.utils.load_model`` imports that Python file straight out of the
|
||||
checkpoint directory, and ``mlx_lm.load()`` exposes no ``trust_remote_code``
|
||||
parameter to refuse it. The gate therefore lives on the SGLang side:
|
||||
|
||||
1. Resolve the model path (local directory or HF repo id + revision) to a
|
||||
local directory exactly once, with mlx-lm's own resolver.
|
||||
2. Inspect THAT directory's ``config.json``. If it declares ``model_file``
|
||||
and the server was not started with ``--trust-remote-code``, refuse
|
||||
before any checkpoint Python can execute.
|
||||
3. Hand the same resolved directory to ``mlx_lm.load`` (for which an
|
||||
existing local directory is a no-op resolution), so the inspected and
|
||||
executed snapshots cannot diverge.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class RemoteCodeGateError(RuntimeError):
|
||||
"""A checkpoint failed the remote-code gate (refusal or bad metadata)."""
|
||||
|
||||
|
||||
def resolve_model_directory(model_path: str, revision: Optional[str] = None) -> Path:
|
||||
"""Resolve a model path or HF repo id to a local snapshot directory.
|
||||
|
||||
Uses mlx-lm's resolver so the directory is byte-identical to what a
|
||||
direct ``mlx_lm.load`` call would consume; existing local paths are
|
||||
returned as-is (no network access). mlx-lm 0.31.x exposes this as
|
||||
``mlx_lm.utils._download`` (formerly ``get_model_path``); mlx-lm is
|
||||
unpinned, so accept either name.
|
||||
"""
|
||||
from mlx_lm import utils as mlx_lm_utils
|
||||
|
||||
resolver = getattr(mlx_lm_utils, "_download", None) or getattr(
|
||||
mlx_lm_utils, "get_model_path", None
|
||||
)
|
||||
if resolver is None:
|
||||
raise RemoteCodeGateError(
|
||||
"this mlx-lm exposes neither mlx_lm.utils._download nor "
|
||||
"mlx_lm.utils.get_model_path, so the checkpoint directory cannot "
|
||||
"be resolved for inspection before mlx-lm loads it"
|
||||
)
|
||||
resolved = resolver(model_path, revision=revision)
|
||||
# get_model_path returned (path, config) in some releases.
|
||||
if isinstance(resolved, tuple):
|
||||
resolved = resolved[0]
|
||||
return Path(resolved)
|
||||
|
||||
|
||||
def ensure_remote_code_allowed(model_dir: Path, trust_remote_code: bool) -> None:
|
||||
"""Refuse ``model_file`` checkpoints unless remote code is trusted.
|
||||
|
||||
Must be called with the SAME resolved directory that is subsequently
|
||||
passed to ``mlx_lm.load``. Raises :class:`RemoteCodeGateError` before
|
||||
any checkpoint Python executes when the checkpoint declares
|
||||
``model_file`` without ``--trust-remote-code``, when its config is
|
||||
unreadable, or when the ``model_file`` value is malformed.
|
||||
"""
|
||||
config_path = model_dir / "config.json"
|
||||
try:
|
||||
config = json.loads(config_path.read_text())
|
||||
except FileNotFoundError:
|
||||
raise RemoteCodeGateError(
|
||||
f"no config.json in resolved model directory {model_dir}; "
|
||||
"not a loadable MLX checkpoint"
|
||||
) from None
|
||||
except json.JSONDecodeError as e:
|
||||
raise RemoteCodeGateError(
|
||||
f"config.json in {model_dir} is not valid JSON ({e}); refusing "
|
||||
"to load a checkpoint whose metadata cannot be inspected"
|
||||
) from None
|
||||
if not isinstance(config, dict):
|
||||
raise RemoteCodeGateError(
|
||||
f"config.json in {model_dir} must contain a JSON object, "
|
||||
f"found {type(config).__name__}"
|
||||
)
|
||||
|
||||
model_file = config.get("model_file")
|
||||
if model_file is None:
|
||||
return
|
||||
|
||||
if not isinstance(model_file, str) or not model_file:
|
||||
raise RemoteCodeGateError(
|
||||
f"config.json in {model_dir} has a non-string or empty "
|
||||
f"model_file entry ({model_file!r})"
|
||||
)
|
||||
candidate = Path(model_file)
|
||||
if candidate.is_absolute() or ".." in candidate.parts:
|
||||
raise RemoteCodeGateError(
|
||||
f"model_file {model_file!r} in {model_dir} must be a relative "
|
||||
"path inside the checkpoint directory (no absolute paths, no "
|
||||
"'..' traversal)"
|
||||
)
|
||||
if not (model_dir / candidate).is_file():
|
||||
raise RemoteCodeGateError(
|
||||
f"config.json in {model_dir} declares model_file "
|
||||
f"{model_file!r} but that file does not exist in the "
|
||||
"checkpoint directory"
|
||||
)
|
||||
|
||||
if not trust_remote_code:
|
||||
raise RemoteCodeGateError(
|
||||
f"checkpoint {model_dir} ships custom model code "
|
||||
f"(model_file={model_file!r} in config.json), which mlx-lm "
|
||||
"would execute at load time. Refusing to load it: restart the "
|
||||
"server with --trust-remote-code if you trust this checkpoint."
|
||||
)
|
||||
@@ -90,6 +90,7 @@ class MlxTpModelWorker(TpModelWorker):
|
||||
disable_radix_cache=get_memory().disable_radix_cache,
|
||||
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||
quantization=get_model().quantization,
|
||||
revision=get_model().revision,
|
||||
enable_sampling=get_device().mlx_enable_sampling,
|
||||
sampling_rng_seed=get_device().random_seed,
|
||||
deterministic_seeding=(
|
||||
|
||||
@@ -1250,7 +1250,8 @@ def flashinfer_mxfp8_blockscaled_linear(
|
||||
# At small M the persistent CUTLASS kernel is 2-5x slower than the
|
||||
# CuTe-DSL swap-AB/split-K kernels (both consume the same swizzled
|
||||
# 1D scales).
|
||||
if backend == "cutlass" and q_input.shape[0] <= 64:
|
||||
# CuTe-DSL has no mm_mxfp8 kernel on SM120, so the swap is SM100-only there.
|
||||
if backend == "cutlass" and q_input.shape[0] <= 64 and _is_sm100_supported:
|
||||
backend = "cute-dsl"
|
||||
|
||||
if backend == "trtllm":
|
||||
|
||||
@@ -829,7 +829,11 @@ class Scheduler(
|
||||
|
||||
# Load multimodal processor for M-RoPE fallback computation.
|
||||
self._mm_processor = None
|
||||
if self.model_config.is_multimodal and self.processor is not None:
|
||||
if (
|
||||
self.model_config.is_multimodal
|
||||
and self.processor is not None
|
||||
and not server_args.language_model_only
|
||||
):
|
||||
try:
|
||||
import_processors("sglang.srt.multimodal.processors")
|
||||
self._mm_processor = get_mm_processor(
|
||||
|
||||
@@ -467,7 +467,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
server_args = self.server_args
|
||||
|
||||
# Initialize tokenizer and processor
|
||||
if self.model_config.is_multimodal:
|
||||
if self.model_config.is_multimodal and not server_args.language_model_only:
|
||||
import_processors("sglang.srt.multimodal.processors")
|
||||
if mm_process_pkg := envs.SGLANG_EXTERNAL_MM_PROCESSOR_PACKAGE.get():
|
||||
import_processors(mm_process_pkg, overwrite=True)
|
||||
@@ -1024,6 +1024,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
)
|
||||
|
||||
contains_mm_input = obj.contains_mm_input()
|
||||
if contains_mm_input and self.server_args.language_model_only:
|
||||
raise ValueError(
|
||||
"Multimodal inputs are not supported when --language-model-only "
|
||||
"is set; the encoder is not loaded. Restart without the flag."
|
||||
)
|
||||
is_mossvl = (
|
||||
"MossVLForConditionalGeneration"
|
||||
in self.model_config.hf_config.architectures
|
||||
|
||||
@@ -27,8 +27,13 @@ def configure_kv_cache_dtype(
|
||||
is_draft_worker: bool,
|
||||
is_dflash: bool,
|
||||
speculative_draft_attention_backend: str,
|
||||
speculative_draft_kv_cache_dtype: Optional[str] = None,
|
||||
) -> tuple[Optional[str], torch.dtype]:
|
||||
resolved_kv_cache_dtype: Optional[str] = None
|
||||
if is_draft_worker and speculative_draft_kv_cache_dtype is not None:
|
||||
server_args_kv_cache_dtype = speculative_draft_kv_cache_dtype
|
||||
if server_args_kv_cache_dtype != "auto":
|
||||
resolved_kv_cache_dtype = server_args_kv_cache_dtype
|
||||
if server_args_kv_cache_dtype == "auto":
|
||||
quant_config = getattr(model, "quant_config", None)
|
||||
kv_cache_quant_algo = getattr(quant_config, "kv_cache_quant_algo", None)
|
||||
@@ -90,5 +95,7 @@ def configure_kv_cache_dtype(
|
||||
model_dtype,
|
||||
)
|
||||
kv_cache_dtype = model_dtype
|
||||
# "auto" is the tag for an unquantized pool; backends gate descale on it.
|
||||
resolved_kv_cache_dtype = "auto"
|
||||
|
||||
return resolved_kv_cache_dtype, kv_cache_dtype
|
||||
|
||||
@@ -722,6 +722,11 @@ class ModelRunner:
|
||||
ElasticEPStateManager.init(self.server_args)
|
||||
|
||||
def init_token_oracle(self):
|
||||
# The oracle sampler is process-wide, so a draft would overwrite the
|
||||
# target's with its own vocab -- which a DFlash draft does not have.
|
||||
if self.is_draft_worker:
|
||||
self._token_oracle_manager = None
|
||||
return
|
||||
self._token_oracle_manager = install_token_oracle_from_env(
|
||||
server_args=self.server_args,
|
||||
vocab_size=self.model_config.vocab_size,
|
||||
@@ -1301,6 +1306,7 @@ class ModelRunner:
|
||||
else False
|
||||
),
|
||||
speculative_draft_attention_backend=self.draft_attention_backend,
|
||||
speculative_draft_kv_cache_dtype=self.server_args.speculative_draft_kv_cache_dtype,
|
||||
)
|
||||
)
|
||||
# This runner's OWN resolved dtype string (target or draft). Attention
|
||||
|
||||
@@ -146,6 +146,13 @@ def _resolve_dflash_aux_hidden_state(
|
||||
draft_num_layers=int(draft_num_layers),
|
||||
)
|
||||
|
||||
# Native export uses HF layer-output ids; shift them.
|
||||
draft_architectures = (
|
||||
getattr(draft_model_config.hf_config, "architectures", None) or []
|
||||
)
|
||||
if "MuseGlimmerAssistantModel" in draft_architectures:
|
||||
target_layer_ids = [i + 1 for i in target_layer_ids]
|
||||
|
||||
if spec_algorithm.is_dspark():
|
||||
from sglang.srt.speculative.dspark_components.dspark_config import (
|
||||
parse_dspark_draft_config,
|
||||
@@ -190,6 +197,9 @@ def _resolve_dflash_draft_cell_size(
|
||||
try:
|
||||
_, draft_kv_cache_dtype = configure_kv_cache_dtype(
|
||||
server_args_kv_cache_dtype=server_args.kv_cache_dtype,
|
||||
speculative_draft_kv_cache_dtype=(
|
||||
server_args.speculative_draft_kv_cache_dtype
|
||||
),
|
||||
model=None,
|
||||
model_dtype=draft_model_config.dtype,
|
||||
is_draft_worker=True,
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
# Copyright 2023-2026 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Per-architecture GGUF -> HF tensor name maps.
|
||||
|
||||
``GGUFModelLoader`` normally derives this map from ``gguf.get_tensor_name_map``,
|
||||
which only covers architectures upstream gguf-py knows, and from a meta-device
|
||||
``AutoModelForCausalLM.from_config`` to enumerate the HF parameter names. Neither
|
||||
works for an architecture that lives outside transformers, so those are supplied
|
||||
here instead.
|
||||
|
||||
A builder returns the complete ``{gguf_tensor_name: hf_param_name}`` map. Any
|
||||
GGUF tensor left out of the map is skipped by ``gguf_quant_weights_iterator``,
|
||||
which is how dummy tensors are dropped.
|
||||
"""
|
||||
|
||||
from typing import Callable, Dict
|
||||
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
# Sandwich naming: ffn_norm is the pre-FFN norm.
|
||||
_MUSE_GLIMMER_LAYER_TENSORS = {
|
||||
"attn_norm": "input_layernorm",
|
||||
"post_attention_norm": "post_attn_norm",
|
||||
"ffn_norm": "post_attention_layernorm",
|
||||
"post_ffw_norm": "post_ffn_norm",
|
||||
"attn_q": "self_attn.q_proj",
|
||||
"attn_k": "self_attn.k_proj",
|
||||
"attn_v": "self_attn.v_proj",
|
||||
"attn_output": "self_attn.o_proj",
|
||||
"attn_gate": "self_attn.output_gate_proj",
|
||||
"ffn_gate": "mlp.gate_proj",
|
||||
"ffn_up": "mlp.up_proj",
|
||||
"ffn_down": "mlp.down_proj",
|
||||
}
|
||||
|
||||
_MUSE_GLIMMER_GLOBAL_TENSORS = {
|
||||
"token_embd": "model.embed_tokens",
|
||||
"output_norm": "model.norm",
|
||||
"output": "lm_head",
|
||||
}
|
||||
|
||||
# attn_q_norm/attn_k_norm omitted: Muse Glimmer's QK-norm is non-parametric.
|
||||
|
||||
|
||||
def build_muse_glimmer_name_map(config: PretrainedConfig) -> Dict[str, str]:
|
||||
name_map = {
|
||||
f"{gguf}.weight": f"{hf}.weight"
|
||||
for gguf, hf in _MUSE_GLIMMER_GLOBAL_TENSORS.items()
|
||||
}
|
||||
for layer in range(config.num_hidden_layers):
|
||||
for gguf, hf in _MUSE_GLIMMER_LAYER_TENSORS.items():
|
||||
name_map[f"blk.{layer}.{gguf}.weight"] = f"model.layers.{layer}.{hf}.weight"
|
||||
return name_map
|
||||
|
||||
|
||||
# Keyed by HF ``config.model_type`` (loader.py looks it up with that), which is
|
||||
# not the GGUF ``general.architecture`` that GGUF_NATIVE_CONFIG_BUILDERS uses:
|
||||
# llama.cpp spells the arch "muse-glimmer" while the HF config says "muse_glimmer".
|
||||
GGUF_HF_NAME_MAP_BUILDERS: Dict[str, Callable[[PretrainedConfig], Dict[str, str]]] = {
|
||||
"muse_glimmer": build_muse_glimmer_name_map,
|
||||
}
|
||||
@@ -3038,8 +3038,14 @@ class GGUFModelLoader(BaseModelLoader):
|
||||
"Please install gguf via `pip install gguf` to use gguf quantizer."
|
||||
) from err
|
||||
|
||||
from sglang.srt.model_loader.gguf_name_maps import GGUF_HF_NAME_MAP_BUILDERS
|
||||
|
||||
config = model_config.hf_config
|
||||
model_type = config.model_type
|
||||
name_map_builder = GGUF_HF_NAME_MAP_BUILDERS.get(model_type)
|
||||
if name_map_builder is not None:
|
||||
return name_map_builder(config)
|
||||
|
||||
# hack: ggufs have a different name than transformers
|
||||
if model_type == "cohere":
|
||||
model_type = "command-r"
|
||||
|
||||
@@ -43,6 +43,16 @@ if _is_npu:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _get_dflash_attention_type(config) -> AttentionType:
|
||||
"""Bidirectional over the draft block unless the checkpoint says causal."""
|
||||
text_config = getattr(config, "text_config", None) or config
|
||||
return (
|
||||
AttentionType.DECODER
|
||||
if getattr(text_config, "is_causal", False)
|
||||
else AttentionType.ENCODER_ONLY
|
||||
)
|
||||
|
||||
|
||||
def _get_dflash_layer_attention_params(
|
||||
config, layer_id: int
|
||||
) -> Tuple[int, AttentionType]:
|
||||
@@ -57,17 +67,12 @@ def _get_dflash_layer_attention_params(
|
||||
|
||||
layer_type = layer_types[layer_id]
|
||||
if layer_type == "full_attention":
|
||||
text_config = getattr(config, "text_config", None) or config
|
||||
attention_type = (
|
||||
AttentionType.DECODER
|
||||
if getattr(text_config, "is_causal", False)
|
||||
else AttentionType.ENCODER_ONLY
|
||||
)
|
||||
return -1, attention_type
|
||||
return -1, _get_dflash_attention_type(config)
|
||||
if layer_type == "sliding_attention":
|
||||
# Windowing is orthogonal to causality (mask is p1 - p0 >= window).
|
||||
sliding_window_size = get_dflash_attention_sliding_window_size(config)
|
||||
assert sliding_window_size is not None
|
||||
return sliding_window_size, AttentionType.DECODER
|
||||
return sliding_window_size, _get_dflash_attention_type(config)
|
||||
raise ValueError(
|
||||
"Unsupported DFLASH draft layer type. "
|
||||
f"layer_types[{layer_id}]={layer_type!r}."
|
||||
@@ -381,11 +386,12 @@ class DFlashDraftModel(nn.Module):
|
||||
# concat(K * hidden_size) -> hidden_size, where K is the number of target-layer
|
||||
# feature tensors concatenated per token (not necessarily equal to num_layers).
|
||||
draft_config = parse_dflash_draft_config(draft_hf_config=config)
|
||||
target_num_layers = (
|
||||
int(draft_config.num_target_layers)
|
||||
if draft_config.num_target_layers is not None
|
||||
else num_layers
|
||||
)
|
||||
if draft_config.num_target_layers is not None:
|
||||
target_num_layers = int(draft_config.num_target_layers)
|
||||
elif draft_config.target_layer_ids is not None:
|
||||
target_num_layers = max(draft_config.target_layer_ids) + 1
|
||||
else:
|
||||
target_num_layers = num_layers
|
||||
target_layer_ids = draft_config.resolve_target_layer_ids(
|
||||
target_num_layers=target_num_layers, draft_num_layers=num_layers
|
||||
)
|
||||
@@ -470,6 +476,12 @@ class DFlashDraftModel(nn.Module):
|
||||
|
||||
params_dict = dict(self.named_parameters())
|
||||
|
||||
# Alias the native export's "encoder." names.
|
||||
_VENDOR_ENCODER_ALIASES = {
|
||||
"encoder.fc.weight": "fc.weight",
|
||||
"encoder.output_norm_enc.weight": "hidden_norm.weight",
|
||||
}
|
||||
|
||||
def resolve_param_name(name: str) -> Optional[str]:
|
||||
if name in params_dict:
|
||||
return name
|
||||
@@ -481,6 +493,9 @@ class DFlashDraftModel(nn.Module):
|
||||
prefixed_name = f"model.{name}"
|
||||
if prefixed_name in params_dict:
|
||||
return prefixed_name
|
||||
aliased_name = _VENDOR_ENCODER_ALIASES.get(name)
|
||||
if aliased_name is not None and aliased_name in params_dict:
|
||||
return aliased_name
|
||||
return None
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
@@ -608,4 +623,8 @@ class DFlashLagunaForCausalLM(DFlashDraftModel):
|
||||
return self.hidden_norm(self.fc(fused))
|
||||
|
||||
|
||||
EntryClass = [DFlashDraftModel, DFlashLagunaForCausalLM]
|
||||
class MuseGlimmerAssistantModel(DFlashDraftModel):
|
||||
"""Alias for checkpoints declaring architectures=["MuseGlimmerAssistantModel"]."""
|
||||
|
||||
|
||||
EntryClass = [DFlashDraftModel, DFlashLagunaForCausalLM, MuseGlimmerAssistantModel]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,61 @@
|
||||
# Copyright 2023-2026 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""SGLang multimodal processor for Muse Glimmer (images)."""
|
||||
|
||||
from typing import Dict, List, Union
|
||||
|
||||
from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput
|
||||
from sglang.srt.models.muse_glimmer import MuseGlimmerForConditionalGeneration
|
||||
from sglang.srt.multimodal.processors.base_processor import (
|
||||
BaseMultimodalProcessor as SGLangBaseProcessor,
|
||||
)
|
||||
from sglang.srt.multimodal.processors.base_processor import (
|
||||
MultimodalSpecialTokens,
|
||||
)
|
||||
|
||||
|
||||
class MuseGlimmerMultimodalProcessor(SGLangBaseProcessor):
|
||||
models = [MuseGlimmerForConditionalGeneration]
|
||||
|
||||
def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
|
||||
super().__init__(hf_config, server_args, _processor, *args, **kwargs)
|
||||
|
||||
self.image_token_id = _processor.image_token_id
|
||||
|
||||
self.mm_tokens = MultimodalSpecialTokens(
|
||||
image_token=_processor.image_token,
|
||||
image_token_id=self.image_token_id,
|
||||
).build(_processor)
|
||||
|
||||
async def process_mm_data_async(
|
||||
self,
|
||||
image_data: List[Union[str, bytes, Dict]],
|
||||
input_text,
|
||||
request_obj,
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
base_output = await self.load_mm_data(
|
||||
prompt=input_text,
|
||||
image_data=image_data,
|
||||
multimodal_tokens=self.mm_tokens,
|
||||
)
|
||||
mm_items, input_ids, _ = await self.process_and_combine_mm_data_async(
|
||||
base_output, self.mm_tokens
|
||||
)
|
||||
return MultimodalProcessorOutput(
|
||||
input_ids=input_ids.tolist(),
|
||||
mm_items=mm_items,
|
||||
im_token_id=self.image_token_id,
|
||||
)
|
||||
@@ -23,6 +23,17 @@ from sglang.srt.function_call.kimik3_format import (
|
||||
strip_partial_marker_suffix,
|
||||
strip_response_wrappers,
|
||||
)
|
||||
from sglang.srt.function_call.muse_glimmer_format import (
|
||||
EOM,
|
||||
EOT,
|
||||
MAX_CHANNEL_MARKER,
|
||||
MESSAGE,
|
||||
RECIPIENT_RE,
|
||||
START,
|
||||
could_start_header,
|
||||
has_atem_markers,
|
||||
partial_marker_len,
|
||||
)
|
||||
from sglang.srt.parser.harmony_parser import HarmonyParser
|
||||
from sglang.srt.parser.inkling_tokenizer import (
|
||||
CONTENT_INVOKE_TOOL_JSON,
|
||||
@@ -1614,6 +1625,213 @@ class CohereCommand4Detector(BaseReasoningFormatDetector):
|
||||
return self._maybe_apply_force_nonempty_content(ret)
|
||||
|
||||
|
||||
class MuseGlimmerDetector(BaseReasoningFormatDetector):
|
||||
"""Detector for Muse Glimmer's recipient-channel format.
|
||||
|
||||
The chat template ends the generation prompt at ``<|start|>assistant`` with no
|
||||
recipient and no ``<|message|>``, so the model itself emits the channel header as
|
||||
ordinary text. A full turn looks like::
|
||||
|
||||
" to=self<|message|>" <reasoning> "<|eom|>"
|
||||
"<|start|>assistant to=user<|message|>" <answer> "<|eot|>"
|
||||
|
||||
Reasoning is the ``to=self`` channel; the answer is ``to=user``. Any other recipient
|
||||
is a tool call (``to=functions.get_weather``), whose body is an ATEM block that must
|
||||
reach the function-call detector with its markers intact — so those channels are
|
||||
emitted as normal text including their header, following GptOssDetector's precedent
|
||||
of preserving raw structural text for tool calls.
|
||||
|
||||
When a tool-call parser consumes this detector's normal text
|
||||
(``tool_call_parser_active=True``), the ``to=user`` channel keeps its framing too,
|
||||
so the downstream detector sees every channel boundary and can tell a real tool
|
||||
channel from one merely *quoted* inside the answer — unwrapping here would make a
|
||||
quoted ``<|start|>assistant to=<tool><|message|>`` indistinguishable from a real
|
||||
header and turn quoted markup into a live call. The tool detector unwraps
|
||||
``to=user`` itself, so nothing framed leaks to the client. Non-streaming
|
||||
additionally requires that a turn *without* any ATEM block come out unwrapped,
|
||||
because serving bypasses the tool detector entirely when ``has_tool_call()`` is
|
||||
false — hence the ATEM-presence branch in ``detect_and_parse``, mirroring the
|
||||
vendor's reference reasoning parser.
|
||||
|
||||
Keying on ``<|message|>`` rather than the literal " to=self" mirrors the vendor's own
|
||||
reference implementation (which slices past the last ``<|message|>`` token),
|
||||
and is robust to the header varying with
|
||||
the recipient. It does require the delimiters to survive detokenization, which is why
|
||||
``muse`` is registered in ``_patch_reasoning_skip_special_tokens``.
|
||||
|
||||
A single channel may also be cut short by the token cap, in which case there is no
|
||||
terminator and the partial body is still attributed to whichever channel was open.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stream_reasoning: bool = True,
|
||||
force_reasoning: bool = False,
|
||||
continue_final_message: bool = False,
|
||||
previous_content: str = "",
|
||||
force_nonempty_content: bool = False,
|
||||
tool_call_parser_active: bool = False,
|
||||
):
|
||||
super().__init__(
|
||||
" to=self" + MESSAGE,
|
||||
EOM,
|
||||
force_reasoning=force_reasoning,
|
||||
stream_reasoning=stream_reasoning,
|
||||
continue_final_message=continue_final_message,
|
||||
previous_content=previous_content,
|
||||
force_nonempty_content=force_nonempty_content,
|
||||
)
|
||||
self._recipient: Optional[str] = None
|
||||
self._in_body = False
|
||||
self._at_stream_start = True
|
||||
self._pending_reasoning = ""
|
||||
self._tool_call_parser_active = tool_call_parser_active
|
||||
self._saw_reasoning_block = False
|
||||
|
||||
def _sink(self, recipient: Optional[str]) -> str:
|
||||
return "reasoning" if recipient == "self" else "normal"
|
||||
|
||||
def _consume(self, flush: bool, preserve_channels: bool = False) -> Tuple[str, str]:
|
||||
"""Drain self._buffer into (reasoning, normal).
|
||||
|
||||
With flush=False, holds back a short tail that could be the prefix of a marker
|
||||
split across chunk boundaries; with flush=True, emits everything.
|
||||
|
||||
With preserve_channels=True, the ``to=user`` channel keeps its header and
|
||||
terminator like tool channels do (see the class docstring for why the
|
||||
function-call detector needs the framing intact); reasoning is always
|
||||
extracted and never framed.
|
||||
"""
|
||||
reasoning_parts: List[str] = []
|
||||
normal_parts: List[str] = []
|
||||
|
||||
while self._buffer:
|
||||
if not self._in_body:
|
||||
# Without this, unframed prose never streams: it buffers
|
||||
# forever waiting for a <|message|> that never arrives.
|
||||
if not (self._at_stream_start and could_start_header(self._buffer)):
|
||||
ws = len(self._buffer) - len(self._buffer.lstrip())
|
||||
head = self._buffer[ws : ws + len(START)]
|
||||
if not START.startswith(head):
|
||||
self._in_body = True
|
||||
self._recipient = None
|
||||
self._at_stream_start = False
|
||||
continue
|
||||
if ws:
|
||||
normal_parts.append(self._buffer[:ws])
|
||||
self._buffer = self._buffer[ws:]
|
||||
if len(head) < len(START):
|
||||
break
|
||||
|
||||
idx = self._buffer.find(MESSAGE)
|
||||
if idx == -1:
|
||||
if flush:
|
||||
normal_parts.append(self._buffer)
|
||||
self._buffer = ""
|
||||
break
|
||||
self._at_stream_start = False
|
||||
header = self._buffer[:idx]
|
||||
m = RECIPIENT_RE.search(header)
|
||||
self._recipient = m.group(1) if m else "user"
|
||||
self._buffer = self._buffer[idx + len(MESSAGE) :]
|
||||
self._in_body = True
|
||||
if self._sink(self._recipient) == "reasoning":
|
||||
if self._saw_reasoning_block:
|
||||
reasoning_parts.append("\n")
|
||||
self._saw_reasoning_block = True
|
||||
elif self._recipient != "user" or preserve_channels:
|
||||
# Keep the header so the function-call detector sees it.
|
||||
normal_parts.append(header + MESSAGE)
|
||||
continue
|
||||
|
||||
end_idx, end_tok = -1, ""
|
||||
for tok in (EOM, EOT):
|
||||
i = self._buffer.find(tok)
|
||||
if i != -1 and (end_idx == -1 or i < end_idx):
|
||||
end_idx, end_tok = i, tok
|
||||
|
||||
if end_idx != -1:
|
||||
body = self._buffer[:end_idx]
|
||||
self._buffer = self._buffer[end_idx + len(end_tok) :]
|
||||
self._in_body = False
|
||||
if self._sink(self._recipient) == "reasoning":
|
||||
reasoning_parts.append(body)
|
||||
else:
|
||||
normal_parts.append(body)
|
||||
if self._recipient != "user" or preserve_channels:
|
||||
normal_parts.append(end_tok)
|
||||
self._recipient = None
|
||||
continue
|
||||
|
||||
# Hold back only a genuine marker prefix.
|
||||
if flush:
|
||||
body, self._buffer = self._buffer, ""
|
||||
else:
|
||||
keep = partial_marker_len(
|
||||
self._buffer, (EOM, EOT, START), MAX_CHANNEL_MARKER
|
||||
)
|
||||
if keep == len(self._buffer):
|
||||
break
|
||||
body = self._buffer[: len(self._buffer) - keep]
|
||||
self._buffer = self._buffer[len(self._buffer) - keep :]
|
||||
if not body:
|
||||
break
|
||||
if self._sink(self._recipient) == "reasoning":
|
||||
reasoning_parts.append(body)
|
||||
else:
|
||||
normal_parts.append(body)
|
||||
|
||||
return "".join(reasoning_parts), "".join(normal_parts)
|
||||
|
||||
def detect_and_parse(self, text: str) -> StreamingParseResult:
|
||||
self._buffer += text
|
||||
raw = self._buffer
|
||||
reasoning, normal = self._consume(flush=True)
|
||||
if self._tool_call_parser_active and has_atem_markers(normal):
|
||||
self._buffer = raw
|
||||
self._recipient = None
|
||||
self._in_body = False
|
||||
self._at_stream_start = True
|
||||
self._saw_reasoning_block = False
|
||||
reasoning, normal = self._consume(flush=True, preserve_channels=True)
|
||||
return self._maybe_apply_force_nonempty_content(
|
||||
StreamingParseResult(normal_text=normal, reasoning_text=reasoning)
|
||||
)
|
||||
|
||||
def parse_streaming_increment(self, new_text: str) -> StreamingParseResult:
|
||||
self._buffer += new_text
|
||||
reasoning, normal = self._consume(
|
||||
flush=False, preserve_channels=self._tool_call_parser_active
|
||||
)
|
||||
if not self.stream_reasoning:
|
||||
self._pending_reasoning += reasoning
|
||||
reasoning = ""
|
||||
if not self._in_body and self._pending_reasoning:
|
||||
reasoning, self._pending_reasoning = self._pending_reasoning, ""
|
||||
if self._force_nonempty_content:
|
||||
# Kept so finish() can promote it to content if the turn produces
|
||||
# none. Dropped on real content, NOT when the channel closes --
|
||||
# <|eom|> lands in the same chunk as the last reasoning text.
|
||||
self._accumulated_reasoning += reasoning
|
||||
if normal:
|
||||
self._accumulated_reasoning = ""
|
||||
return StreamingParseResult(normal_text=normal, reasoning_text=reasoning)
|
||||
|
||||
def finish(self) -> StreamingParseResult:
|
||||
reasoning, normal = self._consume(
|
||||
flush=True, preserve_channels=self._tool_call_parser_active
|
||||
)
|
||||
if self._pending_reasoning:
|
||||
reasoning = self._pending_reasoning + reasoning
|
||||
self._pending_reasoning = ""
|
||||
if self._force_nonempty_content:
|
||||
promoted = self._accumulated_reasoning + reasoning
|
||||
self._accumulated_reasoning = ""
|
||||
if not normal and promoted:
|
||||
return StreamingParseResult(normal_text=promoted)
|
||||
return StreamingParseResult(normal_text=normal, reasoning_text=reasoning)
|
||||
|
||||
|
||||
class ReasoningParser:
|
||||
"""
|
||||
Parser that handles both streaming and non-streaming scenarios for extracting
|
||||
@@ -1623,6 +1841,10 @@ class ReasoningParser:
|
||||
model_type (str): Type of model to parse reasoning from
|
||||
stream_reasoning (bool): If False, accumulates reasoning content until complete.
|
||||
If True, streams reasoning content as it arrives.
|
||||
tool_call_parser_active (bool): True when this parser's normal text feeds a
|
||||
function-call parser rather than going straight to the client. Passed on
|
||||
to detectors that accept it (channel-framed formats keep tool framing
|
||||
intact for the downstream detector).
|
||||
"""
|
||||
|
||||
DetectorMap: Dict[str, Type[BaseReasoningFormatDetector]] = {
|
||||
@@ -1637,6 +1859,7 @@ class ReasoningParser:
|
||||
"kimi_k2": KimiK2Detector,
|
||||
"kimi_k3": KimiK3Detector,
|
||||
"mimo": _MimoDetector,
|
||||
"muse": MuseGlimmerDetector,
|
||||
"poolside_v1": _PoolsideV1Detector,
|
||||
"qwen3": Qwen3Detector,
|
||||
"qwen3-thinking": Qwen3Detector,
|
||||
@@ -1660,6 +1883,7 @@ class ReasoningParser:
|
||||
force_reasoning: Optional[bool] = None,
|
||||
request: ChatCompletionRequest = None,
|
||||
tokenizer=None,
|
||||
tool_call_parser_active: bool = False,
|
||||
):
|
||||
if not model_type:
|
||||
raise ValueError("Model type must be specified")
|
||||
@@ -1705,6 +1929,11 @@ class ReasoningParser:
|
||||
if "tokenizer" in sig.parameters:
|
||||
kwargs["tokenizer"] = tokenizer
|
||||
|
||||
if tool_call_parser_active:
|
||||
sig = inspect.signature(detector_class)
|
||||
if "tool_call_parser_active" in sig.parameters:
|
||||
kwargs["tool_call_parser_active"] = True
|
||||
|
||||
self.detector = detector_class(**kwargs)
|
||||
|
||||
def parse_non_stream(self, full_text: str) -> Tuple[Optional[str], Optional[str]]:
|
||||
|
||||
@@ -1747,6 +1747,7 @@ class ServerArgs:
|
||||
help="Choose the runner backend for NVFP4 GEMM operations. Options: 'auto' (default; selects flashinfer_cutedsl on SM100, marlin on SM80-SM90, flashinfer_cutlass otherwise (including SM120)), 'flashinfer_cutlass' (FlashInfer CUTLASS backend), 'flashinfer_cudnn' (FlashInfer cuDNN backend, optimal on CUDA 13+ with cuDNN 9.15+), 'flashinfer_cutedsl' (FlashInfer CuTe DSL backend), 'flashinfer_trtllm' (FlashInfer TensorRT-LLM backend, requires different weight preparation with shuffling), 'marlin' (weight-only W4A16 fallback for SM80+). ",
|
||||
cli_name="--fp4-gemm-backend",
|
||||
choices=FP4_GEMM_RUNNER_BACKEND_CHOICES,
|
||||
resolvable=True,
|
||||
),
|
||||
NS("exec.kernel"),
|
||||
] = "auto"
|
||||
@@ -2158,6 +2159,20 @@ class ServerArgs:
|
||||
"Attention backend for speculative decoding drafting.",
|
||||
NS("spec"),
|
||||
] = None
|
||||
speculative_draft_kv_cache_dtype: A[
|
||||
Optional[str],
|
||||
Arg(
|
||||
help="KV cache dtype for the speculative draft model only. The draft pool is "
|
||||
"allocated with one slot per target token (draft and target share a slot index "
|
||||
"space), so for a small draft it can still rival the target pool: a 5-layer "
|
||||
"DFLASH draft costs 10240 bytes/token in bf16. Setting fp8_e4m3 halves the draft "
|
||||
"pool; the saving shows up as free device memory, so raise "
|
||||
"--mem-fraction-static to convert it into KV capacity. Default follows "
|
||||
"--kv-cache-dtype.",
|
||||
choices=["auto", "fp8_e5m2", "fp8_e4m3", "bf16", "bfloat16"],
|
||||
),
|
||||
NS("spec"),
|
||||
] = None
|
||||
speculative_draft_window_size: A[
|
||||
Optional[int],
|
||||
"Sliding window size for the draft model. Honored by Llama EAGLE-3 (`LlamaForCausalLMEagle3`) and DFLASH only; other EAGLE-3 backends (e.g. MLA-based drafters) silently ignore it. For Llama EAGLE-3, the drafter only attends to the most recent N keys (verifier hidden states + its own outputs); the verifier is unaffected. For DFLASH, the draft worker keeps a recent target-token window in its local KV cache (paged backends may retain up to one extra page on the left for alignment). Default is full attention/context.",
|
||||
@@ -3093,6 +3108,14 @@ class ServerArgs:
|
||||
language_only: A[
|
||||
bool, "For VLM, load weights for the language model only.", NS("disagg")
|
||||
] = False
|
||||
language_model_only: A[
|
||||
bool,
|
||||
"Skip the multimodal encoder entirely: its weights are never loaded and the "
|
||||
"tower is never built, freeing that GPU memory for KV cache. Multimodal "
|
||||
"requests are rejected. Unlike --language-only this is a standalone mode, "
|
||||
"not part of encoder/decoder disaggregation.",
|
||||
NS("disagg"),
|
||||
] = False
|
||||
encoder_transfer_backend: A[
|
||||
str,
|
||||
Arg(
|
||||
@@ -3549,6 +3572,7 @@ class ServerArgs:
|
||||
# resolution (the declarative registry materializes too late to affect
|
||||
# it). Inkling opts into full-graph prefill capture here.
|
||||
self._apply_inkling_prefill_cuda_graph_default()
|
||||
self._apply_muse_glimmer_prefill_cuda_graph_max_bs_default()
|
||||
|
||||
# must run before _handle_cuda_graph_config and _handle_data_parallelism
|
||||
self._handle_dwdp()
|
||||
@@ -3852,6 +3876,8 @@ class ServerArgs:
|
||||
|
||||
def _handle_model_source_paths(self):
|
||||
"""Prepare metadata for model paths backed by remote object stores."""
|
||||
self._resolve_hf_gguf_model_path()
|
||||
|
||||
seen_paths = set()
|
||||
for model_path in (
|
||||
self.model_path,
|
||||
@@ -4303,6 +4329,16 @@ class ServerArgs:
|
||||
):
|
||||
self.cuda_graph_backend_prefill = Backend.FULL
|
||||
|
||||
def _apply_muse_glimmer_prefill_cuda_graph_max_bs_default(self):
|
||||
if (
|
||||
self.cuda_graph_max_bs_prefill is not None
|
||||
or parse_connector_type(self.model_path) == ConnectorType.INSTANCE
|
||||
):
|
||||
return
|
||||
arch = self.get_model_config().hf_config.architectures[0]
|
||||
if arch in ("MuseGlimmerForCausalLM", "MuseGlimmerForConditionalGeneration"):
|
||||
self.cuda_graph_max_bs_prefill = 512
|
||||
|
||||
def _handle_cuda_graph_config(self):
|
||||
from sglang.srt.arg_groups.kimi_k3_hook import disable_kimi_k3_symm_mem
|
||||
|
||||
@@ -4861,6 +4897,7 @@ class ServerArgs:
|
||||
if (
|
||||
model_config.is_multimodal
|
||||
and not self.language_only
|
||||
and not self.language_model_only
|
||||
and self.disaggregation_mode != "decode"
|
||||
):
|
||||
self.adjust_mem_fraction_for_vlm(model_config)
|
||||
@@ -6392,7 +6429,11 @@ class ServerArgs:
|
||||
raise ValueError(
|
||||
"MiMo V2 CP-v2 only supports --cp-strategy zigzag."
|
||||
)
|
||||
if model_config.is_multimodal and not self.language_only:
|
||||
if (
|
||||
model_config.is_multimodal
|
||||
and not self.language_only
|
||||
and not self.language_model_only
|
||||
):
|
||||
raise ValueError(
|
||||
"MiMo V2 CP-v2 only supports text inference; add "
|
||||
"--language-only."
|
||||
@@ -7320,6 +7361,32 @@ class ServerArgs:
|
||||
f"switching to {new_layout} layout for {self.hicache_io_backend} io backend"
|
||||
)
|
||||
|
||||
def _resolve_hf_gguf_model_path(self):
|
||||
"""Turn a Hub reference to a .gguf into a local file path."""
|
||||
from sglang.srt.utils.hf_transformers_utils import resolve_hf_gguf_reference
|
||||
|
||||
resolved = resolve_hf_gguf_reference(self.model_path, revision=self.revision)
|
||||
if resolved is not None:
|
||||
logger.info("Resolved GGUF %s -> %s", self.model_path, resolved)
|
||||
if self.tokenizer_path == self.model_path:
|
||||
self.tokenizer_path = resolved
|
||||
self.model_path = resolved
|
||||
|
||||
# A speculative draft can be a .gguf too, and it is loaded by path, so it
|
||||
# needs the same Hub-reference resolution as the target.
|
||||
if self.speculative_draft_model_path:
|
||||
resolved_draft = resolve_hf_gguf_reference(
|
||||
self.speculative_draft_model_path,
|
||||
revision=self.speculative_draft_model_revision,
|
||||
)
|
||||
if resolved_draft is not None:
|
||||
logger.info(
|
||||
"Resolved draft GGUF %s -> %s",
|
||||
self.speculative_draft_model_path,
|
||||
resolved_draft,
|
||||
)
|
||||
self.speculative_draft_model_path = resolved_draft
|
||||
|
||||
def _handle_load_format(self):
|
||||
# The quantization side of the gguf coupling moved to the pipeline
|
||||
# (arg_groups/overrides.py: _gguf_quantization); load_format itself is
|
||||
@@ -7474,7 +7541,39 @@ class ServerArgs:
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
LANGUAGE_MODEL_ONLY_ARCHITECTURES = ("MuseGlimmerForConditionalGeneration",)
|
||||
|
||||
def _handle_language_model_only(self):
|
||||
if not self.language_model_only:
|
||||
return
|
||||
for flag, name in (
|
||||
(self.encoder_only, "--encoder-only"),
|
||||
(self.language_only, "--language-only"),
|
||||
(self.enable_prefix_mm_cache, "--enable-prefix-mm-cache"),
|
||||
(
|
||||
self.enable_broadcast_mm_inputs_process,
|
||||
"--enable-broadcast-mm-inputs-process",
|
||||
),
|
||||
(self.mm_enable_dp_encoder, "--mm-enable-dp-encoder"),
|
||||
):
|
||||
if flag:
|
||||
raise ValueError(
|
||||
f"--language-model-only cannot be combined with {name}"
|
||||
)
|
||||
if self.disaggregation_mode != "null":
|
||||
raise ValueError(
|
||||
"--language-model-only is incompatible with --disaggregation-mode "
|
||||
"prefill/decode"
|
||||
)
|
||||
architectures = self.get_model_config().hf_config.architectures
|
||||
if not any(a in self.LANGUAGE_MODEL_ONLY_ARCHITECTURES for a in architectures):
|
||||
raise ValueError(
|
||||
f"--language-model-only does not support {architectures}. "
|
||||
f"Supported: {list(self.LANGUAGE_MODEL_ONLY_ARCHITECTURES)}."
|
||||
)
|
||||
|
||||
def _handle_encoder_disaggregation(self):
|
||||
self._handle_language_model_only()
|
||||
if self.enable_prefix_mm_cache and not self.encoder_only:
|
||||
raise ValueError(
|
||||
"--enable-prefix-mm-cache requires --encoder-only to be enabled"
|
||||
|
||||
@@ -515,7 +515,9 @@ def parse_dflash_draft_config(*, draft_hf_config: Any) -> DFlashDraftConfig:
|
||||
f"Got len(target_layer_ids)={len(parsed_target_layer_ids)}."
|
||||
)
|
||||
|
||||
mask_token = dflash_cfg.get("mask_token", None)
|
||||
mask_token = dflash_cfg.get(
|
||||
"mask_token", _cfg_get(draft_hf_config, "mask_token", None)
|
||||
)
|
||||
if mask_token is None:
|
||||
mask_token = DEFAULT_DFLASH_MASK_TOKEN
|
||||
if not isinstance(mask_token, str) or not mask_token:
|
||||
@@ -524,7 +526,9 @@ def parse_dflash_draft_config(*, draft_hf_config: Any) -> DFlashDraftConfig:
|
||||
f"got {mask_token!r}."
|
||||
)
|
||||
|
||||
mask_token_id = dflash_cfg.get("mask_token_id", None)
|
||||
mask_token_id = dflash_cfg.get(
|
||||
"mask_token_id", _cfg_get(draft_hf_config, "mask_token_id", None)
|
||||
)
|
||||
if mask_token_id is not None:
|
||||
if not isinstance(mask_token_id, Integral) or isinstance(mask_token_id, bool):
|
||||
raise ValueError(
|
||||
|
||||
@@ -79,6 +79,14 @@ def _get_fused_kv_materialize_helper():
|
||||
return _FusedKVMaterializeHelper
|
||||
|
||||
|
||||
# is_floating_point() is True for fp8; list dtypes explicitly.
|
||||
_DENSE_HEAD_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
|
||||
|
||||
|
||||
def _is_dense_head_weight(weight) -> bool:
|
||||
return weight is not None and weight.dtype in _DENSE_HEAD_DTYPES
|
||||
|
||||
|
||||
class _DflashDraftSampler:
|
||||
"""Capture-safe greedy argmax over the target LM head, run inside the draft
|
||||
cuda graph so the draft sampling is captured and counted in fwd_occupancy.
|
||||
@@ -231,6 +239,12 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
mask_token=self._mask_token,
|
||||
mask_token_id=self._mask_token_id_override,
|
||||
)
|
||||
target_model = self._target_worker.model_runner.model
|
||||
self._noise_embed_scale = (
|
||||
float(target_model.get_dflash_noise_embedding_scale())
|
||||
if hasattr(target_model, "get_dflash_noise_embedding_scale")
|
||||
else 1.0
|
||||
)
|
||||
if self.ps.tp_rank == 0:
|
||||
logger.info(
|
||||
"Initialized DFLASH draft runner. attention_backend=%s, model=%s, block_size=%s, draft_window_size=%s, compact_cache=%s",
|
||||
@@ -241,10 +255,11 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
self.use_compact_draft_cache,
|
||||
)
|
||||
logger.info(
|
||||
"DFLASH draft runner ready. mask_token=%s, mask_token_id=%s, mask_token_id_override=%s",
|
||||
"DFLASH draft runner ready. mask_token=%s, mask_token_id=%s, mask_token_id_override=%s, noise_embed_scale=%s",
|
||||
self._mask_token,
|
||||
self._mask_token_id,
|
||||
self._mask_token_id_override,
|
||||
self._noise_embed_scale,
|
||||
)
|
||||
|
||||
self._block_pos_offsets = build_block_pos_offsets(
|
||||
@@ -375,9 +390,11 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
return _eager("block_size<=1")
|
||||
target_model = self._target_worker.model_runner.model
|
||||
lm_head = getattr(target_model, "lm_head", None)
|
||||
if lm_head is None or not hasattr(lm_head, "weight"):
|
||||
if lm_head is None:
|
||||
return _eager("no target lm_head")
|
||||
if not torch.is_floating_point(lm_head.weight):
|
||||
if not hasattr(lm_head, "weight"):
|
||||
return _eager("quantized lm_head has no dense weight")
|
||||
if not _is_dense_head_weight(lm_head.weight):
|
||||
# Quantized lm_head (FP8/INT) would break the static matmul.
|
||||
return _eager("quantized lm_head")
|
||||
tp_group = get_tp_group()
|
||||
@@ -805,6 +822,42 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
|
||||
return int(resolved_id)
|
||||
|
||||
def _greedy_sample_from_quantized_head(
|
||||
self,
|
||||
*,
|
||||
hidden_states: torch.Tensor,
|
||||
lm_head,
|
||||
chunk_size: int,
|
||||
) -> torch.Tensor:
|
||||
"""Greedy argmax over a target LM head that has no dense ``weight``.
|
||||
|
||||
A GGUF head stores packed ``qweight`` plus a type tag, so the dense path's
|
||||
``weight[:num_org]`` slicing has nothing to slice. Logits come from the
|
||||
layer's own kernel instead -- the same call ``LogitsProcessor._get_logits``
|
||||
makes for GGUF models. Padding rows are excluded so argmax cannot return
|
||||
an id outside the real vocabulary.
|
||||
"""
|
||||
tp_size = int(get_tp_group().world_size)
|
||||
if tp_size != 1:
|
||||
raise RuntimeError(
|
||||
"DFLASH with a quantized target lm_head is only supported at "
|
||||
f"tp=1, got tp_size={tp_size}."
|
||||
)
|
||||
|
||||
num_tokens = int(hidden_states.shape[0])
|
||||
out_tokens = torch.empty(
|
||||
(num_tokens,), dtype=torch.long, device=hidden_states.device
|
||||
)
|
||||
num_org = int(getattr(lm_head, "org_vocab_size", 0)) or None
|
||||
|
||||
for start in range(0, num_tokens, int(chunk_size)):
|
||||
end = min(num_tokens, start + int(chunk_size))
|
||||
logits = lm_head.quant_method.apply(lm_head, hidden_states[start:end], None)
|
||||
if num_org is not None and logits.shape[-1] > num_org:
|
||||
logits = logits[:, :num_org]
|
||||
out_tokens[start:end] = torch.argmax(logits, dim=-1).to(torch.long)
|
||||
return out_tokens
|
||||
|
||||
def _greedy_sample_from_vocab_parallel_head(
|
||||
self,
|
||||
*,
|
||||
@@ -822,6 +875,11 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
if hidden_states.numel() == 0:
|
||||
return torch.empty((0,), dtype=torch.long, device=hidden_states.device)
|
||||
|
||||
if not _is_dense_head_weight(getattr(lm_head, "weight", None)):
|
||||
return self._greedy_sample_from_quantized_head(
|
||||
hidden_states=hidden_states, lm_head=lm_head, chunk_size=chunk_size
|
||||
)
|
||||
|
||||
weight = lm_head.weight # [local_vocab_padded, hidden]
|
||||
weight_dtype = weight.dtype
|
||||
num_tokens = int(hidden_states.shape[0])
|
||||
@@ -1487,9 +1545,13 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
target_model = self.target_worker.model_runner.model
|
||||
embed_module = target_model.get_input_embeddings()
|
||||
lm_head = getattr(target_model, "lm_head", None)
|
||||
if lm_head is None or not hasattr(lm_head, "weight"):
|
||||
if lm_head is None or not (
|
||||
hasattr(lm_head, "weight")
|
||||
or callable(getattr(getattr(lm_head, "quant_method", None), "apply", None))
|
||||
):
|
||||
raise RuntimeError(
|
||||
"DFLASH requires the target model to expose `lm_head` with `weight`."
|
||||
"DFLASH requires the target model to expose `lm_head` with either "
|
||||
"`weight` or a `quant_method` that can produce logits."
|
||||
)
|
||||
|
||||
block_size = int(self.block_size)
|
||||
@@ -1562,6 +1624,8 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
verify_out_cache_loc_2d.copy_(verify_out_cache_loc.view(bs, block_size))
|
||||
|
||||
noise_embedding = embed_module(block_ids)
|
||||
if self._noise_embed_scale != 1.0:
|
||||
noise_embedding = noise_embedding * self._noise_embed_scale
|
||||
input_embeds = noise_embedding.view(-1, noise_embedding.shape[-1])
|
||||
|
||||
positions = positions_2d.reshape(-1)
|
||||
|
||||
@@ -101,6 +101,10 @@ def build_draft_tp_worker(
|
||||
|
||||
draft_model_runner = draft_worker.model_runner
|
||||
draft_worker.draft_runner = draft_model_runner
|
||||
|
||||
# DFlash drafts have no vocab; borrow the target's.
|
||||
if draft_model_runner.model_config.vocab_size is None:
|
||||
draft_model_runner.model_config.vocab_size = target_model_config.vocab_size
|
||||
return DraftWorkerBundle(
|
||||
draft_worker=draft_worker,
|
||||
draft_model_runner=draft_model_runner,
|
||||
|
||||
@@ -35,6 +35,7 @@ from .common import (
|
||||
get_rope_config,
|
||||
get_sparse_attention_config,
|
||||
get_tokenizer_from_processor,
|
||||
resolve_hf_gguf_reference,
|
||||
)
|
||||
from .config import get_config
|
||||
from .processor import get_processor, resolve_image_processor_backend
|
||||
@@ -51,6 +52,7 @@ __all__ = [
|
||||
"_fix_v5_add_bos_eos_token",
|
||||
"attach_additional_stop_token_ids",
|
||||
"check_gguf_file",
|
||||
"resolve_hf_gguf_reference",
|
||||
"download_from_hf",
|
||||
"get_config",
|
||||
"get_context_length",
|
||||
|
||||
@@ -52,6 +52,8 @@ from sglang.srt.configs import (
|
||||
MiniCPMV4_6VisionConfig,
|
||||
MiniMaxM3VLConfig,
|
||||
MultiModalityConfig,
|
||||
MuseGlimmerAssistantConfig,
|
||||
MuseGlimmerConfig,
|
||||
NemotronH_Nano_Omni_Reasoning_V3_Config,
|
||||
NemotronH_Nano_VL_V2_Config,
|
||||
NemotronHConfig,
|
||||
@@ -101,6 +103,8 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
|
||||
Step3VLConfig,
|
||||
LongcatFlashConfig,
|
||||
Olmo3Config,
|
||||
MuseGlimmerConfig,
|
||||
MuseGlimmerAssistantConfig,
|
||||
KimiK3Config,
|
||||
KimiLinearConfig,
|
||||
Qwen3NextConfig,
|
||||
@@ -284,6 +288,64 @@ def check_gguf_file(model: Union[str, os.PathLike]) -> bool:
|
||||
return header == b"GGUF"
|
||||
|
||||
|
||||
def resolve_hf_gguf_reference(
|
||||
model: str, revision: Optional[str] = None
|
||||
) -> Optional[str]:
|
||||
"""Download a .gguf named by Hub reference and return its local path.
|
||||
|
||||
owner/repo/path/inside/repo.gguf -> exactly that file
|
||||
owner/repo -> the only .gguf in the repo
|
||||
"""
|
||||
from sglang.srt.utils import is_remote_url
|
||||
|
||||
if not model or os.path.exists(model) or is_remote_url(model):
|
||||
return None
|
||||
|
||||
parts = model.strip("/").split("/")
|
||||
if len(parts) < 2:
|
||||
return None
|
||||
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
if len(parts) > 2 and model.endswith(".gguf"):
|
||||
repo_id = "/".join(parts[:2])
|
||||
filename = "/".join(parts[2:])
|
||||
return hf_hub_download(repo_id, filename, revision=revision)
|
||||
|
||||
if len(parts) != 2:
|
||||
return None
|
||||
|
||||
from huggingface_hub import HfApi
|
||||
|
||||
try:
|
||||
files = [
|
||||
s.rfilename for s in HfApi().repo_info(model, revision=revision).siblings
|
||||
]
|
||||
except Exception:
|
||||
return None
|
||||
if any(f == "config.json" for f in files):
|
||||
return None
|
||||
|
||||
candidates = [f for f in files if f.endswith(".gguf")]
|
||||
if not candidates:
|
||||
return None
|
||||
if len(candidates) > 1:
|
||||
listing = "\n ".join(f"{model}/{f}" for f in sorted(candidates))
|
||||
raise ValueError(
|
||||
f"{model} contains {len(candidates)} .gguf files; name the one to "
|
||||
f"serve:\n {listing}"
|
||||
)
|
||||
return hf_hub_download(model, candidates[0], revision=revision)
|
||||
|
||||
|
||||
def gguf_sidecar_dir(
|
||||
gguf_path: Union[str, os.PathLike], sentinel: str
|
||||
) -> Optional[Path]:
|
||||
"""Directory containing *sentinel* next to a .gguf file, if there is one."""
|
||||
directory = Path(gguf_path).parent
|
||||
return directory if (directory / sentinel).is_file() else None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Rope / text config helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -491,6 +553,19 @@ def get_generation_config(
|
||||
revision: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
if check_gguf_file(model):
|
||||
sidecar = gguf_sidecar_dir(model, "generation_config.json")
|
||||
if sidecar is not None:
|
||||
model = str(sidecar)
|
||||
else:
|
||||
from .gguf_native import (
|
||||
build_gguf_generation_config,
|
||||
has_native_gguf_support,
|
||||
)
|
||||
|
||||
if has_native_gguf_support(model):
|
||||
return build_gguf_generation_config(model)
|
||||
|
||||
try:
|
||||
return GenerationConfig.from_pretrained(
|
||||
model, trust_remote_code=trust_remote_code, revision=revision, **kwargs
|
||||
|
||||
@@ -37,8 +37,10 @@ from .common import (
|
||||
_override_v_head_dim_if_zero,
|
||||
check_gguf_file,
|
||||
get_hf_text_config,
|
||||
gguf_sidecar_dir,
|
||||
resolve_runai_obj_uri,
|
||||
)
|
||||
from .gguf_native import build_gguf_config, has_native_gguf_support
|
||||
from .mistral_utils import is_mistral_model, load_mistral_config
|
||||
|
||||
|
||||
@@ -224,6 +226,7 @@ def get_config(
|
||||
**kwargs,
|
||||
):
|
||||
is_gguf = check_gguf_file(model)
|
||||
gguf_has_sidecar_config = False
|
||||
if is_gguf:
|
||||
if model_config_parser not in ("auto", "hf"):
|
||||
raise ValueError(
|
||||
@@ -231,7 +234,14 @@ def get_config(
|
||||
"with GGUF inputs; only 'hf' (or 'auto') is supported."
|
||||
)
|
||||
_ensure_gguf_version()
|
||||
kwargs["gguf_file"] = model
|
||||
gguf_has_sidecar_config = gguf_sidecar_dir(model, "config.json") is not None
|
||||
if not gguf_has_sidecar_config and has_native_gguf_support(model):
|
||||
config = build_gguf_config(model)
|
||||
if model_override_args:
|
||||
config.update(model_override_args)
|
||||
return config
|
||||
if not gguf_has_sidecar_config:
|
||||
kwargs["gguf_file"] = model
|
||||
model = Path(model).parent
|
||||
# Skip auto-resolution for GGUF: the name-based Mistral heuristic
|
||||
# would misfire on the rewritten parent dir.
|
||||
@@ -264,9 +274,13 @@ def get_config(
|
||||
else:
|
||||
setattr(config, key, value)
|
||||
|
||||
if is_gguf:
|
||||
if is_gguf and not gguf_has_sidecar_config:
|
||||
if config.model_type not in MODEL_FOR_CAUSAL_LM_MAPPING_NAMES:
|
||||
raise RuntimeError(f"Can't get gguf config for {config.model_type}.")
|
||||
raise RuntimeError(
|
||||
f"Can't get gguf config for {config.model_type}. Place a "
|
||||
"config.json next to the .gguf file to load the config from "
|
||||
"there instead."
|
||||
)
|
||||
_set_architectures(config, MODEL_FOR_CAUSAL_LM_MAPPING_NAMES[config.model_type])
|
||||
|
||||
return config
|
||||
|
||||
@@ -0,0 +1,258 @@
|
||||
# Copyright 2023-2026 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Reading config and tokenizer from a GGUF whose architecture transformers lacks.
|
||||
|
||||
``load_gguf_checkpoint`` refuses any architecture outside its own
|
||||
``GGUF_SUPPORTED_ARCHITECTURES``, and it does so before touching a single field,
|
||||
so both the config and the tokenizer are unreachable for such a checkpoint --
|
||||
even though the tokenizer half of that reader is entirely architecture-agnostic
|
||||
(it dispatches on ``tokenizer.ggml.model``, not on the model architecture).
|
||||
|
||||
This module carries SGLang's own path for those checkpoints:
|
||||
|
||||
* ``GGUF_NATIVE_CONFIG_BUILDERS`` maps a GGUF ``general.architecture`` to a
|
||||
builder returning a fully populated config.
|
||||
* ``build_gguf_tokenizer`` reuses transformers' own converters, which work fine
|
||||
once they are reached directly instead of through the gated loader.
|
||||
|
||||
Reaching for these is a last resort: a config.json next to the .gguf still wins,
|
||||
because the checkpoint author's own config outranks anything reconstructed.
|
||||
"""
|
||||
|
||||
from typing import Any, Callable, Dict, Optional
|
||||
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from sglang.srt.configs.muse_glimmer import MuseGlimmerConfig
|
||||
|
||||
GGUF_NATIVE_CONFIG_BUILDERS: Dict[str, Callable[[str], PretrainedConfig]] = {
|
||||
"muse-glimmer": MuseGlimmerConfig.from_gguf,
|
||||
}
|
||||
|
||||
|
||||
def read_gguf_architecture(gguf_path: str) -> Optional[str]:
|
||||
"""The ``general.architecture`` string, or None if it cannot be read."""
|
||||
try:
|
||||
from gguf import GGUFReader
|
||||
|
||||
reader = GGUFReader(gguf_path)
|
||||
field = reader.fields.get("general.architecture")
|
||||
if field is None:
|
||||
return None
|
||||
value = field.contents()
|
||||
return value if isinstance(value, str) else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def has_native_gguf_support(gguf_path: str) -> bool:
|
||||
return read_gguf_architecture(gguf_path) in GGUF_NATIVE_CONFIG_BUILDERS
|
||||
|
||||
|
||||
def build_gguf_config(gguf_path: str) -> PretrainedConfig:
|
||||
arch = read_gguf_architecture(gguf_path)
|
||||
return GGUF_NATIVE_CONFIG_BUILDERS[arch](gguf_path)
|
||||
|
||||
|
||||
_GPT4O_SPLIT_REGEX = (
|
||||
r"[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}]*"
|
||||
r"[\p{Ll}\p{Lm}\p{Lo}\p{M}]+(?i:'s|'t|'re|'ve|'m|'ll|'d)?|"
|
||||
r"[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}]+"
|
||||
r"[\p{Ll}\p{Lm}\p{Lo}\p{M}]*(?i:'s|'t|'re|'ve|'m|'ll|'d)?|"
|
||||
r"\p{N}{1,3}| ?[^\s\p{L}\p{N}]+[\r\n/]*|\s*[\r\n]+|\s+(?!\S)|\s+"
|
||||
)
|
||||
|
||||
_PRE_TOKENIZER_REGEX = {
|
||||
# LLAMA_VOCAB_PRE_TYPE_LLAMA3
|
||||
"llama-bpe": (
|
||||
r"(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|"
|
||||
r"[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}{1,3}| ?[^\s\p{L}\p{N}]+[\r\n]*|"
|
||||
r"\s*[\r\n]+|\s+(?!\S)|\s+"
|
||||
),
|
||||
"gpt-4o": _GPT4O_SPLIT_REGEX,
|
||||
"llama4": _GPT4O_SPLIT_REGEX,
|
||||
}
|
||||
|
||||
_GGML_TOKEN_TYPE_CONTROL = 3
|
||||
|
||||
|
||||
def build_gguf_generation_config(gguf_path: str):
|
||||
"""GenerationConfig from GGUF metadata, or None if there is nothing to say.
|
||||
|
||||
llama.cpp records the end-of-generation ids explicitly, and for a
|
||||
Harmony-style model the distinction matters: ``eos_token_id`` ends the
|
||||
sequence and ``eot_token_id`` ends a turn, so both must stop generation while
|
||||
an end-of-*message* id must not -- stopping on that truncates the model
|
||||
mid-reasoning, before it answers.
|
||||
"""
|
||||
from gguf import GGUFReader
|
||||
from transformers import GenerationConfig
|
||||
|
||||
reader = GGUFReader(gguf_path)
|
||||
meta = {key: field.contents() for key, field in reader.fields.items()}
|
||||
|
||||
stop_ids = []
|
||||
for key in ("tokenizer.ggml.eos_token_id", "tokenizer.ggml.eot_token_id"):
|
||||
if key in meta:
|
||||
value = int(meta[key])
|
||||
if value not in stop_ids:
|
||||
stop_ids.append(value)
|
||||
if not stop_ids:
|
||||
return None
|
||||
|
||||
fields: Dict[str, Any] = {
|
||||
"eos_token_id": stop_ids if len(stop_ids) > 1 else stop_ids[0]
|
||||
}
|
||||
if "tokenizer.ggml.bos_token_id" in meta:
|
||||
fields["bos_token_id"] = int(meta["tokenizer.ggml.bos_token_id"])
|
||||
if "tokenizer.ggml.padding_token_id" in meta:
|
||||
fields["pad_token_id"] = int(meta["tokenizer.ggml.padding_token_id"])
|
||||
return GenerationConfig(**fields)
|
||||
|
||||
|
||||
def build_gguf_tokenizer(gguf_path: str, **kwargs: Any):
|
||||
"""Build a fast tokenizer from GGUF metadata alone.
|
||||
|
||||
transformers' own GGUF tokenizer path is unreachable for an architecture its
|
||||
checkpoint loader rejects, and its converters key on ``tokenizer.ggml.model``
|
||||
(here "gpt2") which loses both the special-token block and the pre-tokenizer
|
||||
regex. So the tokenizers spec is assembled directly instead: a byte-level BPE
|
||||
over the NORMAL tokens, the CONTROL tokens registered as added specials, and
|
||||
the split regex named by ``tokenizer.ggml.pre``.
|
||||
"""
|
||||
import json
|
||||
|
||||
from gguf import GGUFReader
|
||||
from tokenizers import Tokenizer
|
||||
from transformers import PreTrainedTokenizerFast
|
||||
|
||||
reader = GGUFReader(gguf_path)
|
||||
meta = {key: field.contents() for key, field in reader.fields.items()}
|
||||
|
||||
tokens = list(meta["tokenizer.ggml.tokens"])
|
||||
token_types = [int(t) for t in meta["tokenizer.ggml.token_type"]]
|
||||
merges = [tuple(m.split(" ", 1)) for m in meta["tokenizer.ggml.merges"]]
|
||||
|
||||
pre_name = meta.get("tokenizer.ggml.pre")
|
||||
if pre_name not in _PRE_TOKENIZER_REGEX:
|
||||
raise ValueError(
|
||||
f"No pre-tokenizer regex known for tokenizer.ggml.pre={pre_name!r}; "
|
||||
f"known: {sorted(_PRE_TOKENIZER_REGEX)}"
|
||||
)
|
||||
|
||||
control_ids = [
|
||||
i for i, t in enumerate(token_types) if t == _GGML_TOKEN_TYPE_CONTROL
|
||||
]
|
||||
control = set(control_ids)
|
||||
vocab = {tok: i for i, tok in enumerate(tokens) if i not in control}
|
||||
|
||||
def token_of(key):
|
||||
idx = meta.get(f"tokenizer.ggml.{key}")
|
||||
return None if idx is None else tokens[int(idx)]
|
||||
|
||||
bos = token_of("bos_token_id")
|
||||
|
||||
spec = {
|
||||
"version": "1.0",
|
||||
"truncation": None,
|
||||
"padding": None,
|
||||
"added_tokens": [
|
||||
{
|
||||
"id": i,
|
||||
"content": tokens[i],
|
||||
"single_word": False,
|
||||
"lstrip": False,
|
||||
"rstrip": False,
|
||||
"normalized": False,
|
||||
"special": True,
|
||||
}
|
||||
for i in control_ids
|
||||
],
|
||||
"normalizer": None,
|
||||
"pre_tokenizer": {
|
||||
"type": "Sequence",
|
||||
"pretokenizers": [
|
||||
{
|
||||
"type": "Split",
|
||||
"pattern": {"Regex": _PRE_TOKENIZER_REGEX[pre_name]},
|
||||
"behavior": "Isolated",
|
||||
"invert": False,
|
||||
},
|
||||
{
|
||||
"type": "ByteLevel",
|
||||
"add_prefix_space": False,
|
||||
"trim_offsets": True,
|
||||
"use_regex": False,
|
||||
},
|
||||
],
|
||||
},
|
||||
"post_processor": None,
|
||||
"decoder": {
|
||||
"type": "ByteLevel",
|
||||
"add_prefix_space": True,
|
||||
"trim_offsets": True,
|
||||
"use_regex": True,
|
||||
},
|
||||
"model": {
|
||||
"type": "BPE",
|
||||
"dropout": None,
|
||||
"unk_token": None,
|
||||
"continuing_subword_prefix": None,
|
||||
"end_of_word_suffix": None,
|
||||
"fuse_unk": False,
|
||||
"byte_fallback": False,
|
||||
"ignore_merges": True,
|
||||
"vocab": vocab,
|
||||
"merges": [list(m) for m in merges],
|
||||
},
|
||||
}
|
||||
|
||||
if meta.get("tokenizer.ggml.add_bos_token") and bos is not None:
|
||||
bos_id = int(meta["tokenizer.ggml.bos_token_id"])
|
||||
spec["post_processor"] = {
|
||||
"type": "TemplateProcessing",
|
||||
"single": [
|
||||
{"SpecialToken": {"id": bos, "type_id": 0}},
|
||||
{"Sequence": {"id": "A", "type_id": 0}},
|
||||
],
|
||||
"pair": [
|
||||
{"SpecialToken": {"id": bos, "type_id": 0}},
|
||||
{"Sequence": {"id": "A", "type_id": 0}},
|
||||
{"Sequence": {"id": "B", "type_id": 0}},
|
||||
],
|
||||
"special_tokens": {
|
||||
bos: {"id": bos, "ids": [bos_id], "tokens": [bos]},
|
||||
},
|
||||
}
|
||||
|
||||
backend = Tokenizer.from_str(json.dumps(spec))
|
||||
|
||||
named = {
|
||||
bos,
|
||||
token_of("eos_token_id"),
|
||||
token_of("padding_token_id"),
|
||||
token_of("unknown_token_id"),
|
||||
}
|
||||
additional = [tokens[i] for i in control_ids if tokens[i] not in named]
|
||||
|
||||
return PreTrainedTokenizerFast(
|
||||
tokenizer_object=backend,
|
||||
bos_token=bos,
|
||||
eos_token=token_of("eos_token_id"),
|
||||
unk_token=token_of("unknown_token_id"),
|
||||
pad_token=token_of("padding_token_id"),
|
||||
additional_special_tokens=additional,
|
||||
chat_template=meta.get("tokenizer.chat_template"),
|
||||
**kwargs,
|
||||
)
|
||||
@@ -34,8 +34,10 @@ from .common import (
|
||||
_resolve_local_or_cached_file,
|
||||
attach_additional_stop_token_ids,
|
||||
check_gguf_file,
|
||||
gguf_sidecar_dir,
|
||||
resolve_runai_obj_uri,
|
||||
)
|
||||
from .gguf_native import build_gguf_tokenizer, has_native_gguf_support
|
||||
from .mistral_utils import (
|
||||
_MISTRAL_TOKENIZER_REDIRECTS,
|
||||
is_bare_tekken_checkpoint,
|
||||
@@ -144,7 +146,8 @@ def _resolve_tokenizer_name(tokenizer_name, kwargs):
|
||||
|
||||
if check_gguf_file(tokenizer_name):
|
||||
_ensure_gguf_version()
|
||||
kwargs["gguf_file"] = tokenizer_name
|
||||
if gguf_sidecar_dir(tokenizer_name, "tokenizer_config.json") is None:
|
||||
kwargs["gguf_file"] = tokenizer_name
|
||||
tokenizer_name = Path(tokenizer_name).parent
|
||||
|
||||
tokenizer_name = resolve_runai_obj_uri(tokenizer_name)
|
||||
@@ -487,6 +490,17 @@ def get_tokenizer(
|
||||
if "use_fast" not in kwargs:
|
||||
kwargs["use_fast"] = True
|
||||
|
||||
if (
|
||||
check_gguf_file(tokenizer_name)
|
||||
and gguf_sidecar_dir(tokenizer_name, "tokenizer_config.json") is None
|
||||
and has_native_gguf_support(tokenizer_name)
|
||||
):
|
||||
_ensure_gguf_version()
|
||||
tokenizer = build_gguf_tokenizer(tokenizer_name)
|
||||
_fix_special_tokens_pattern(tokenizer)
|
||||
attach_additional_stop_token_ids(tokenizer)
|
||||
return patch_tokenizer(tokenizer)
|
||||
|
||||
tokenizer_name = _resolve_tokenizer_name(tokenizer_name, kwargs)
|
||||
|
||||
common_kwargs = dict(
|
||||
|
||||
Reference in New Issue
Block a user