fix: Mistral Small 4 fails to start due to config/weight format mismatch (#21620)
Co-authored-by: mengxiancheng03 <mengxiancheng03@kuaishou.com> Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com> Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
mengxiancheng03
Baizhou Zhang
Claude Opus 4.6
parent
b246269444
commit
1d6424d5ad
@@ -3157,22 +3157,42 @@ class ServerArgs:
|
|||||||
def _is_mistral_native_format(self) -> bool:
|
def _is_mistral_native_format(self) -> bool:
|
||||||
"""Detect if the model uses Mistral native format (params.json + consolidated weights).
|
"""Detect if the model uses Mistral native format (params.json + consolidated weights).
|
||||||
|
|
||||||
Models like Mistral-7B-Instruct-v0.3 have BOTH params.json (native) and
|
When both params.json and config.json exist, default to HF format to
|
||||||
config.json (HF standard). When both exist, prefer the HF format to avoid
|
avoid weight-name mismatches (e.g. Mistral-7B-Instruct-v0.3).
|
||||||
parameter name mismatches between consolidated.safetensors (native names
|
|
||||||
like layers.0.attention.wk.weight) and HuggingFace model classes (names
|
Exception: models routed through ``_load_mistral_large_3_for_causal_LM``
|
||||||
like model.layers.0.self_attn.k_proj.weight).
|
(mistral-large-3, mistral-small-4, leanstral) build their config from
|
||||||
|
params.json and expect native weight names, so native format is required
|
||||||
|
even when config.json is also present.
|
||||||
"""
|
"""
|
||||||
|
# Keep in sync with the name checks in
|
||||||
|
# hf_transformers_utils.py::get_config / get_tokenizer.
|
||||||
|
_MISTRAL_NATIVE_CONFIG_PATTERNS = (
|
||||||
|
"mistral-large-3",
|
||||||
|
"mistral-small-4",
|
||||||
|
"leanstral",
|
||||||
|
)
|
||||||
|
|
||||||
|
def _check_format(has_params: bool, has_hf_config: bool) -> bool:
|
||||||
|
if has_params and not has_hf_config:
|
||||||
|
return True
|
||||||
|
if has_params and has_hf_config:
|
||||||
|
model_lower = str(self.model_path).lower()
|
||||||
|
if any(name in model_lower for name in _MISTRAL_NATIVE_CONFIG_PATTERNS):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
if os.path.isdir(self.model_path):
|
if os.path.isdir(self.model_path):
|
||||||
has_params = os.path.exists(os.path.join(self.model_path, "params.json"))
|
has_params = os.path.exists(os.path.join(self.model_path, "params.json"))
|
||||||
has_hf_config = os.path.exists(os.path.join(self.model_path, "config.json"))
|
has_hf_config = os.path.exists(os.path.join(self.model_path, "config.json"))
|
||||||
return has_params and not has_hf_config
|
return _check_format(has_params, has_hf_config)
|
||||||
|
|
||||||
# For hub models, check remote files
|
# For hub models, check remote files
|
||||||
try:
|
try:
|
||||||
from huggingface_hub import HfApi
|
from huggingface_hub import HfApi
|
||||||
|
|
||||||
files = {s.rfilename for s in HfApi().model_info(self.model_path).siblings}
|
files = {s.rfilename for s in HfApi().model_info(self.model_path).siblings}
|
||||||
return "params.json" in files and "config.json" not in files
|
return _check_format("params.json" in files, "config.json" in files)
|
||||||
except Exception:
|
except Exception:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,32 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
|
||||||
|
from sglang.test.kits.mmmu_vlm_kit import MMMUMixin
|
||||||
|
from sglang.test.server_fixtures.default_fixture import DefaultServerBase
|
||||||
|
from sglang.test.server_fixtures.mmmu_fixture import MMMUServerBase
|
||||||
|
|
||||||
|
register_cuda_ci(
|
||||||
|
est_time=200,
|
||||||
|
suite="stage-b-test-2-gpu-large",
|
||||||
|
)
|
||||||
|
|
||||||
|
MODEL = "mistralai/Mistral-Small-4-119B-2603"
|
||||||
|
|
||||||
|
|
||||||
|
class TestMistralSmall4TextOnly(GSM8KMixin, DefaultServerBase):
|
||||||
|
gsm8k_accuracy_thres = 0.9
|
||||||
|
model = MODEL
|
||||||
|
other_args = ["--tp-size", "2", "--trust-remote-code"]
|
||||||
|
|
||||||
|
|
||||||
|
class TestMistralSmall4MMMU(MMMUMixin, MMMUServerBase):
|
||||||
|
accuracy = 0.45
|
||||||
|
model = MODEL
|
||||||
|
other_args = ["--tp-size", "2", "--trust-remote-code"]
|
||||||
|
mmmu_args = ["--limit=0.1"]
|
||||||
|
"""`--limit=0.1`: 10 percent of each task - this is fine for testing since the nominal result isn't interesting - this run is just to prevent relative regressions."""
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user