Fix Muse Glimmer ModelOpt mixed weight mapping (#37510)
This commit is contained in:
@@ -13,7 +13,6 @@
|
||||
# ==============================================================================
|
||||
|
||||
import logging
|
||||
import re
|
||||
from typing import Iterable, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
@@ -57,7 +56,7 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
default_weight_loader,
|
||||
maybe_remap_kv_scale_name,
|
||||
)
|
||||
from sglang.srt.models.utils import apply_qk_norm, permute_inv
|
||||
from sglang.srt.models.utils import WeightsMapper, apply_qk_norm, permute_inv
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix, is_cuda
|
||||
|
||||
@@ -84,22 +83,26 @@ _VISION_NAME_FRAGMENTS = (
|
||||
"perception_emb_norm",
|
||||
)
|
||||
|
||||
# Vendor tensor names -> this port's; applied simultaneously.
|
||||
_VENDOR_RENAMES = {
|
||||
"post_attention_layernorm": "post_attn_norm",
|
||||
"pre_feedforward_layernorm": "post_attention_layernorm",
|
||||
"post_feedforward_layernorm": "post_ffn_norm",
|
||||
"self_attn.gate_proj": "self_attn.output_gate_proj",
|
||||
}
|
||||
|
||||
_VENDOR_RENAME_RE = re.compile("|".join(re.escape(key) for key in _VENDOR_RENAMES))
|
||||
# Shared by _vendor_weight_name and hf_to_sglang_mapper;
|
||||
# a rule missing from one silently breaks the other.
|
||||
_VENDOR_TO_SGLANG = WeightsMapper(
|
||||
orig_to_new_prefix={
|
||||
"model.language_model.": "model.",
|
||||
# The vision modules hang off the entry class, not off ``model``.
|
||||
"model.vision_": "vision_",
|
||||
},
|
||||
# Only the first matching substring is applied; keep these non-overlapping.
|
||||
orig_to_new_substr={
|
||||
"post_attention_layernorm": "post_attn_norm",
|
||||
"pre_feedforward_layernorm": "post_attention_layernorm",
|
||||
"post_feedforward_layernorm": "post_ffn_norm",
|
||||
"self_attn.gate_proj": "self_attn.output_gate_proj",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _vendor_weight_name(name: str) -> str:
|
||||
name = name.replace("model.language_model.", "model.", 1)
|
||||
# The vision modules hang off the entry class, not off ``model``.
|
||||
name = name.replace("model.vision_", "vision_", 1)
|
||||
return _VENDOR_RENAME_RE.sub(lambda m: _VENDOR_RENAMES[m.group(0)], name)
|
||||
return _VENDOR_TO_SGLANG.apply_list([name])[0]
|
||||
|
||||
|
||||
def get_attention_sliding_window_size(config) -> int:
|
||||
@@ -928,6 +931,8 @@ class MuseGlimmerForCausalLM(nn.Module):
|
||||
class MuseGlimmerForConditionalGeneration(MuseGlimmerForCausalLM):
|
||||
"""Vendor multimodal HF export: the MuseGlimmerForCausalLM decoder plus the image tower."""
|
||||
|
||||
# Only this class reads vendor-named checkpoints, so only it needs the mapper.
|
||||
hf_to_sglang_mapper = _VENDOR_TO_SGLANG
|
||||
checkpoint_uses_vendor_names = True
|
||||
builds_vision_tower = True
|
||||
|
||||
|
||||
Reference in New Issue
Block a user