Merge branch 'main' into studio-fla-tilelang-qwen3.5
This commit is contained in:
commit
57afa6287e
11 changed files with 1206 additions and 47 deletions
11
.github/workflows/consolidated-tests-ci.yml
vendored
11
.github/workflows/consolidated-tests-ci.yml
vendored
|
|
@ -304,6 +304,17 @@ jobs:
|
|||
run: |
|
||||
python -m pytest -v --tb=short tests/test_import_fixes_drift.py
|
||||
|
||||
- name: public-api surface drift detectors (9 tests, HARD GATE)
|
||||
# Companion to test_import_fixes_drift.py: that file catches
|
||||
# third-party drift; this one catches drift in unsloth's OWN
|
||||
# public surface (FastLanguageModel / FastVisionModel /
|
||||
# FastModel + their classmethods + is_bf16_supported). A
|
||||
# rename here would silently break the unslothai/notebooks tree
|
||||
# one PR cycle later -- this gate catches it BEFORE the
|
||||
# breakage reaches users.
|
||||
run: |
|
||||
python -m pytest -v --tb=short tests/test_public_api_surface.py
|
||||
|
||||
- name: unsloth Bucket-A — CPU tests not in Repo tests (CPU)
|
||||
# 16 tests across 5 files. They live inside tests/saving/ and
|
||||
# tests/utils/, both of which Repo tests (CPU) excludes via --ignore
|
||||
|
|
|
|||
|
|
@ -2367,8 +2367,20 @@ class LlamaCppBackend:
|
|||
if not Path(mmproj_path).is_file():
|
||||
logger.warning(f"mmproj file not found: {mmproj_path}")
|
||||
else:
|
||||
cmd.extend(["--mmproj", mmproj_path])
|
||||
logger.info(f"Using mmproj for vision: {mmproj_path}")
|
||||
# #5347 guard for paths that bypass detect_mmproj_file.
|
||||
from utils.models.model_config import (
|
||||
mmproj_matches_model_family,
|
||||
)
|
||||
|
||||
if not mmproj_matches_model_family(model_path, mmproj_path):
|
||||
logger.warning(
|
||||
f"Skipping mmproj with mismatched family: "
|
||||
f"model={Path(model_path).name}, "
|
||||
f"mmproj={Path(mmproj_path).name}"
|
||||
)
|
||||
else:
|
||||
cmd.extend(["--mmproj", mmproj_path])
|
||||
logger.info(f"Using mmproj for vision: {mmproj_path}")
|
||||
|
||||
# Option C: add --api-key for direct client access when enabled
|
||||
import os as _os
|
||||
|
|
@ -3747,7 +3759,7 @@ class LlamaCppBackend:
|
|||
|
||||
except json.JSONDecodeError:
|
||||
logger.debug(
|
||||
f"Skipping malformed SSE line: " f"{line[:100]}"
|
||||
f"Skipping malformed SSE line: {line[:100]}"
|
||||
)
|
||||
if _stream_done:
|
||||
break # exit outer for
|
||||
|
|
|
|||
326
studio/backend/tests/test_detect_mmproj_file.py
Normal file
326
studio/backend/tests/test_detect_mmproj_file.py
Normal file
|
|
@ -0,0 +1,326 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""Tests for :func:`utils.models.model_config.detect_mmproj_file` (#5347)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import struct
|
||||
|
||||
from utils.models.model_config import (
|
||||
_detect_family_token,
|
||||
detect_mmproj_file,
|
||||
mmproj_matches_model_family,
|
||||
)
|
||||
|
||||
|
||||
_GGUF_MAGIC = 0x46554747
|
||||
|
||||
|
||||
def _gguf_with_general(path: Path, fields: dict) -> Path:
|
||||
"""Write a minimal GGUF with only ``general.*`` string KVs."""
|
||||
body = b""
|
||||
for k, v in fields.items():
|
||||
kb = k.encode("utf-8")
|
||||
vb = v.encode("utf-8")
|
||||
body += struct.pack("<Q", len(kb)) + kb
|
||||
body += struct.pack("<I", 8) # STRING vtype
|
||||
body += struct.pack("<Q", len(vb)) + vb
|
||||
header = struct.pack("<IIQQ", _GGUF_MAGIC, 3, 0, len(fields))
|
||||
path.parent.mkdir(parents = True, exist_ok = True)
|
||||
path.write_bytes(header + body)
|
||||
return path
|
||||
|
||||
|
||||
def _touch(path: Path) -> Path:
|
||||
path.parent.mkdir(parents = True, exist_ok = True)
|
||||
path.write_bytes(b"")
|
||||
return path
|
||||
|
||||
|
||||
def test_returns_none_when_no_mmproj(tmp_path: Path):
|
||||
model = _touch(tmp_path / "Qwen3.5-9B-Q4_K_M.gguf")
|
||||
assert detect_mmproj_file(str(model)) is None
|
||||
|
||||
|
||||
def test_single_matching_family_mmproj_picked(tmp_path: Path):
|
||||
"""Single same-family projector: returned (historical behaviour)."""
|
||||
model = _touch(tmp_path / "Qwen3.5-9B-Q4_K_M.gguf")
|
||||
mmproj = _touch(tmp_path / "Qwen3.5-9B-BF16-mmproj.gguf")
|
||||
assert detect_mmproj_file(str(model)) == str(mmproj.resolve())
|
||||
|
||||
|
||||
def test_hf_style_unprefixed_mmproj_still_works(tmp_path: Path):
|
||||
"""HF convention: weight + ``mmproj-F16.gguf`` sibling."""
|
||||
model = _touch(tmp_path / "model.gguf")
|
||||
mmproj = _touch(tmp_path / "mmproj-F16.gguf")
|
||||
assert detect_mmproj_file(str(model)) == str(mmproj.resolve())
|
||||
|
||||
|
||||
def test_blocks_single_cross_family_projector(tmp_path: Path):
|
||||
"""#5347 core: Qwen weight + lone Gemma mmproj returns None."""
|
||||
model = _touch(tmp_path / "Qwen3.5-9B-Q4_K_M.gguf")
|
||||
_touch(tmp_path / "gemma-4-26B-A4B-it.mmproj-q8_0.gguf")
|
||||
assert detect_mmproj_file(str(model)) is None
|
||||
|
||||
|
||||
def test_picks_matching_family_among_mixed_candidates(tmp_path: Path):
|
||||
"""Mixed Qwen + Gemma projectors: pick Qwen, drop Gemma."""
|
||||
model = _touch(tmp_path / "Qwen3.5-9B-Q4_K_M.gguf")
|
||||
qwen_mm = _touch(tmp_path / "Qwen3.5-9B-BF16-mmproj.gguf")
|
||||
_touch(tmp_path / "gemma-4-26B-A4B-it.mmproj-q8_0.gguf")
|
||||
assert detect_mmproj_file(str(model)) == str(qwen_mm.resolve())
|
||||
|
||||
|
||||
def test_prefers_longest_prefix_within_same_family(tmp_path: Path):
|
||||
"""Same family, different sizes: longest shared stem prefix wins."""
|
||||
model = _touch(tmp_path / "Qwen3.5-35B-A3B-UD-Q4_K_L.gguf")
|
||||
_touch(tmp_path / "Qwen3.5-9B-BF16-mmproj.gguf")
|
||||
big_mm = _touch(tmp_path / "Qwen3.5-35B-A3B-BF16-mmproj.gguf")
|
||||
assert detect_mmproj_file(str(model)) == str(big_mm.resolve())
|
||||
|
||||
|
||||
def test_unrecognised_family_does_not_break_detection(tmp_path: Path):
|
||||
"""Unknown model family must not return None on a sole candidate."""
|
||||
model = _touch(tmp_path / "MyCustomBrand-7B-Q4_K_M.gguf")
|
||||
mmproj = _touch(tmp_path / "MyCustomBrand-7B-BF16-mmproj.gguf")
|
||||
assert detect_mmproj_file(str(model)) == str(mmproj.resolve())
|
||||
|
||||
|
||||
def test_directory_path_returns_first_candidate(tmp_path: Path):
|
||||
"""Directory path: no model stem to compare; legacy first-candidate."""
|
||||
_touch(tmp_path / "Qwen3.5-9B-BF16-mmproj.gguf")
|
||||
_touch(tmp_path / "gemma-4-26B-A4B-it.mmproj-q8_0.gguf")
|
||||
result = detect_mmproj_file(str(tmp_path))
|
||||
assert result is not None
|
||||
assert "mmproj" in Path(result).name.lower()
|
||||
|
||||
|
||||
def test_search_root_walk_still_works(tmp_path: Path):
|
||||
"""Snapshot layout: weight in quant subdir, mmproj at snapshot root."""
|
||||
snapshot = tmp_path / "snapshot"
|
||||
weight = _touch(snapshot / "BF16" / "Qwen3.5-9B-BF16.gguf")
|
||||
mmproj = _touch(snapshot / "Qwen3.5-9B-BF16-mmproj.gguf")
|
||||
result = detect_mmproj_file(str(weight), search_root = str(snapshot))
|
||||
assert result == str(mmproj.resolve())
|
||||
|
||||
|
||||
# -- Family token detection: word-bounded matching ----------------------
|
||||
|
||||
|
||||
def test_family_token_phi_does_not_match_sapphire():
|
||||
"""``phi`` substring inside ``sapphire`` must not tag Phi."""
|
||||
assert _detect_family_token("sapphire-7b-q4_k_m.gguf") is None
|
||||
|
||||
|
||||
def test_family_token_yi_does_not_match_tinyish_names():
|
||||
"""``yi`` must not cross letter boundaries (``yip``)."""
|
||||
assert _detect_family_token("yip-7b.gguf") is None
|
||||
assert _detect_family_token("yi-vl-6b.gguf") == "yi"
|
||||
|
||||
|
||||
def test_family_token_mimo_does_not_match_mimosa():
|
||||
"""``mimo`` must not tag ``mimosa``."""
|
||||
assert _detect_family_token("mimosa-rosa-7b.gguf") is None
|
||||
assert _detect_family_token("MiMo-VL-7B-RL-BF16.gguf") == "mimo"
|
||||
|
||||
|
||||
def test_family_token_mistral_does_not_match_ministral():
|
||||
"""Pin Mistral-derivative tagging."""
|
||||
assert _detect_family_token("Ministral-3-8B-Instruct-2512-BF16.gguf") == "ministral"
|
||||
assert _detect_family_token("Mistral-7B-Instruct-v0.3.gguf") == "mistral"
|
||||
assert _detect_family_token("Magistral-Small-2506-BF16.gguf") == "magistral"
|
||||
assert (
|
||||
_detect_family_token("Devstral-Small-2-24B-Instruct-2512-BF16.gguf")
|
||||
== "devstral"
|
||||
)
|
||||
|
||||
|
||||
def test_family_token_picks_leftmost_when_multiple_present():
|
||||
"""Leftmost family token wins, not tuple order."""
|
||||
assert _detect_family_token("llama-phi-merge.gguf") == "llama"
|
||||
assert _detect_family_token("phi-llama-merge.gguf") == "phi"
|
||||
assert _detect_family_token("llama3-3b-instruct.gguf") == "llama"
|
||||
|
||||
|
||||
def test_family_token_new_families_recognised():
|
||||
"""Catalogue-audit additions tag correctly."""
|
||||
assert _detect_family_token("NVIDIA-Nemotron-3-Nano-Omni-30B.gguf") == "nemotron"
|
||||
assert _detect_family_token("Kimi-K2.6-BF16.gguf") == "kimi"
|
||||
assert _detect_family_token("Nanonets-OCR-s-BF16.gguf") == "nanonets"
|
||||
assert _detect_family_token("Cosmos-Reason1-7B-BF16.gguf") == "cosmos"
|
||||
assert _detect_family_token("Apriel-1.5-15b-Thinker-BF16.gguf") == "apriel"
|
||||
assert _detect_family_token("LFM2.5-VL-1.6B-BF16.gguf") == "lfm"
|
||||
|
||||
|
||||
# -- Cross-family rejection with the expanded token list ----------------
|
||||
|
||||
|
||||
def test_blocks_cross_family_for_new_token_pair(tmp_path: Path):
|
||||
"""Nemotron weight + lone Gemma projector returns None."""
|
||||
model = _touch(
|
||||
tmp_path / "NVIDIA-Nemotron-3-Nano-Omni-30B-A3B-Reasoning-MXFP4_MOE.gguf"
|
||||
)
|
||||
_touch(tmp_path / "gemma-4-26B-A4B-it.mmproj-q8_0.gguf")
|
||||
assert detect_mmproj_file(str(model)) is None
|
||||
|
||||
|
||||
def test_picks_devstral_mmproj_in_mixed_dir(tmp_path: Path):
|
||||
"""Devstral weight + Devstral mmproj + a Qwen mmproj: pick Devstral."""
|
||||
model = _touch(tmp_path / "Devstral-Small-2-24B-Instruct-2512-BF16.gguf")
|
||||
dev_mm = _touch(tmp_path / "Devstral-Small-2-mmproj-bf16.gguf")
|
||||
_touch(tmp_path / "Qwen3.5-9B-BF16-mmproj.gguf")
|
||||
assert detect_mmproj_file(str(model)) == str(dev_mm.resolve())
|
||||
|
||||
|
||||
# -- Launcher-level family guard ----------------------------------------
|
||||
|
||||
|
||||
def test_mmproj_family_guard_blocks_cross_family():
|
||||
assert (
|
||||
mmproj_matches_model_family(
|
||||
"/models/Qwen3.5-9B-Q4_K_M.gguf",
|
||||
"/models/gemma-4-26B-A4B-it.mmproj-q8_0.gguf",
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
def test_mmproj_family_guard_allows_same_family():
|
||||
assert (
|
||||
mmproj_matches_model_family(
|
||||
"/models/Qwen3.5-9B-Q4_K_M.gguf",
|
||||
"/models/Qwen3.5-9B-BF16-mmproj.gguf",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_mmproj_family_guard_allows_generic_hf_mmproj():
|
||||
"""No family token on the projector: wildcard."""
|
||||
assert (
|
||||
mmproj_matches_model_family(
|
||||
"/models/Qwen3.5-9B-Q4_K_M.gguf",
|
||||
"/models/mmproj-F16.gguf",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_mmproj_family_guard_allows_unrecognised_model_family():
|
||||
"""No family token on the model: wildcard."""
|
||||
assert (
|
||||
mmproj_matches_model_family(
|
||||
"/models/Apriel-1.5-15b-Thinker-BF16.gguf",
|
||||
"/models/mmproj-F16.gguf",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
# -- Metadata-primary pairing in detect_mmproj_file ---------------------
|
||||
|
||||
|
||||
def test_metadata_url_match_picked_over_filename_lookalike(tmp_path: Path):
|
||||
"""URL match beats a longer-prefix sibling."""
|
||||
weight = _gguf_with_general(
|
||||
tmp_path / "Qwen3.5-9B-Q4_K_M.gguf",
|
||||
{
|
||||
"general.architecture": "qwen2vl",
|
||||
"general.type": "model",
|
||||
"general.basename": "Qwen3.5",
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
},
|
||||
)
|
||||
# Closer filename prefix, wrong upstream.
|
||||
_gguf_with_general(
|
||||
tmp_path / "Qwen3.5-9B-mmproj-bf16.gguf",
|
||||
{
|
||||
"general.architecture": "clip",
|
||||
"general.type": "mmproj",
|
||||
"general.basename": "Qwen3.5",
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-1.5B",
|
||||
},
|
||||
)
|
||||
# Matching upstream.
|
||||
correct = _gguf_with_general(
|
||||
tmp_path / "mmproj-BF16.gguf",
|
||||
{
|
||||
"general.architecture": "clip",
|
||||
"general.type": "mmproj",
|
||||
"general.basename": "Qwen3.5",
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
},
|
||||
)
|
||||
assert detect_mmproj_file(str(weight)) == str(correct.resolve())
|
||||
|
||||
|
||||
def test_metadata_url_mismatch_dropped(tmp_path: Path):
|
||||
"""Filenames match family but metadata disagrees: returns None."""
|
||||
weight = _gguf_with_general(
|
||||
tmp_path / "qwen-9b.gguf",
|
||||
{
|
||||
"general.architecture": "qwen2vl",
|
||||
"general.type": "model",
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
},
|
||||
)
|
||||
_gguf_with_general(
|
||||
tmp_path / "qwen-9b-mmproj.gguf",
|
||||
{
|
||||
"general.architecture": "clip",
|
||||
"general.type": "mmproj",
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/google/gemma-3-9B",
|
||||
},
|
||||
)
|
||||
assert detect_mmproj_file(str(weight)) is None
|
||||
|
||||
|
||||
def test_metadata_identifies_mmproj_without_filename_hint(tmp_path: Path):
|
||||
"""Projector named ``vision-projector.gguf`` discovered via header."""
|
||||
weight = _gguf_with_general(
|
||||
tmp_path / "Qwen3.5-9B.gguf",
|
||||
{
|
||||
"general.architecture": "qwen2vl",
|
||||
"general.type": "model",
|
||||
"general.basename": "Qwen3.5",
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
},
|
||||
)
|
||||
projector = _gguf_with_general(
|
||||
tmp_path / "vision-projector.gguf",
|
||||
{
|
||||
"general.architecture": "clip",
|
||||
"general.type": "mmproj",
|
||||
"general.basename": "Qwen3.5",
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
},
|
||||
)
|
||||
assert detect_mmproj_file(str(weight)) == str(projector.resolve())
|
||||
|
||||
|
||||
def test_metadata_score_outranks_filename_prefix(tmp_path: Path):
|
||||
"""Score 100 (URL match) beats score 0 (long filename prefix)."""
|
||||
weight = _gguf_with_general(
|
||||
tmp_path / "Qwen3.5-9B-Q4_K_M.gguf",
|
||||
{
|
||||
"general.architecture": "qwen2vl",
|
||||
"general.type": "model",
|
||||
"general.basename": "Qwen3.5",
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
},
|
||||
)
|
||||
# Headerless: long shared stem, score 0.
|
||||
_touch(tmp_path / "Qwen3.5-9B-Q4_K_M-mmproj.gguf")
|
||||
# Headered: generic name, score 100.
|
||||
correct = _gguf_with_general(
|
||||
tmp_path / "mmproj-BF16.gguf",
|
||||
{
|
||||
"general.architecture": "clip",
|
||||
"general.type": "mmproj",
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
},
|
||||
)
|
||||
assert detect_mmproj_file(str(weight)) == str(correct.resolve())
|
||||
216
studio/backend/tests/test_gguf_metadata.py
Normal file
216
studio/backend/tests/test_gguf_metadata.py
Normal file
|
|
@ -0,0 +1,216 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""Tests for :mod:`utils.models.gguf_metadata`. Synthesise small GGUF
|
||||
headers in tmp dirs so we never depend on real model files."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import struct
|
||||
from pathlib import Path
|
||||
from typing import Iterable, Mapping
|
||||
|
||||
from utils.models.gguf_metadata import (
|
||||
is_mmproj_by_metadata,
|
||||
pairing_score,
|
||||
read_gguf_general_metadata,
|
||||
)
|
||||
|
||||
|
||||
_GGUF_MAGIC = 0x46554747
|
||||
_VTYPE_STRING = 8
|
||||
_VTYPE_UINT32 = 4
|
||||
_VTYPE_ARRAY = 9
|
||||
|
||||
|
||||
def _enc_string(s: str) -> bytes:
|
||||
b = s.encode("utf-8")
|
||||
return struct.pack("<Q", len(b)) + b
|
||||
|
||||
|
||||
def _enc_kv_string(key: str, value: str) -> bytes:
|
||||
return _enc_string(key) + struct.pack("<I", _VTYPE_STRING) + _enc_string(value)
|
||||
|
||||
|
||||
def _enc_kv_uint32(key: str, value: int) -> bytes:
|
||||
return (
|
||||
_enc_string(key) + struct.pack("<I", _VTYPE_UINT32) + struct.pack("<I", value)
|
||||
)
|
||||
|
||||
|
||||
def _enc_kv_string_array(key: str, values: Iterable[str]) -> bytes:
|
||||
vals = list(values)
|
||||
out = _enc_string(key) + struct.pack("<I", _VTYPE_ARRAY)
|
||||
out += struct.pack("<I", _VTYPE_STRING) + struct.pack("<Q", len(vals))
|
||||
for v in vals:
|
||||
out += _enc_string(v)
|
||||
return out
|
||||
|
||||
|
||||
def _write_synthetic_gguf(
|
||||
path: Path,
|
||||
general_strings: Mapping[str, str],
|
||||
*,
|
||||
extra_uint32: Mapping[str, int] | None = None,
|
||||
extra_string_arrays: Mapping[str, Iterable[str]] | None = None,
|
||||
) -> Path:
|
||||
"""Minimal GGUF: header + KV body, no tensors."""
|
||||
extra_uint32 = extra_uint32 or {}
|
||||
extra_string_arrays = extra_string_arrays or {}
|
||||
kv_count = len(general_strings) + len(extra_uint32) + len(extra_string_arrays)
|
||||
body = b""
|
||||
for k, v in general_strings.items():
|
||||
body += _enc_kv_string(k, v)
|
||||
for k, v in extra_uint32.items():
|
||||
body += _enc_kv_uint32(k, v)
|
||||
for k, v in extra_string_arrays.items():
|
||||
body += _enc_kv_string_array(k, v)
|
||||
header = struct.pack(
|
||||
"<IIQQ",
|
||||
_GGUF_MAGIC,
|
||||
3, # version
|
||||
0, # tensor_count
|
||||
kv_count,
|
||||
)
|
||||
path.parent.mkdir(parents = True, exist_ok = True)
|
||||
path.write_bytes(header + body)
|
||||
return path
|
||||
|
||||
|
||||
# --- read_gguf_general_metadata ----------------------------------------
|
||||
|
||||
|
||||
def test_returns_none_for_missing_file(tmp_path: Path):
|
||||
assert read_gguf_general_metadata(str(tmp_path / "nope.gguf")) is None
|
||||
|
||||
|
||||
def test_returns_none_for_non_gguf(tmp_path: Path):
|
||||
p = tmp_path / "garbage.gguf"
|
||||
p.write_bytes(b"not a gguf file at all, just bytes")
|
||||
assert read_gguf_general_metadata(str(p)) is None
|
||||
|
||||
|
||||
def test_extracts_general_string_fields(tmp_path: Path):
|
||||
p = _write_synthetic_gguf(
|
||||
tmp_path / "model.gguf",
|
||||
{
|
||||
"general.architecture": "qwen2vl",
|
||||
"general.type": "model",
|
||||
"general.basename": "Qwen3.5",
|
||||
"general.organization": "Qwen",
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
"general.base_model.0.name": "Qwen3.5 9B",
|
||||
"general.base_model.0.organization": "Qwen",
|
||||
},
|
||||
)
|
||||
meta = read_gguf_general_metadata(str(p))
|
||||
assert meta is not None
|
||||
assert meta["general.architecture"] == "qwen2vl"
|
||||
assert meta["general.basename"] == "Qwen3.5"
|
||||
assert (
|
||||
meta["general.base_model.0.repo_url"]
|
||||
== "https://huggingface.co/Qwen/Qwen3.5-9B"
|
||||
)
|
||||
|
||||
|
||||
def test_skips_unrelated_fields_without_breaking(tmp_path: Path):
|
||||
"""Skip unwanted arrays and uint32s without losing position."""
|
||||
p = _write_synthetic_gguf(
|
||||
tmp_path / "model.gguf",
|
||||
{"general.basename": "Foo"},
|
||||
extra_uint32 = {"qwen2vl.context_length": 32768},
|
||||
extra_string_arrays = {"tokenizer.ggml.tokens": ["a", "bc", "def"]},
|
||||
)
|
||||
meta = read_gguf_general_metadata(str(p))
|
||||
assert meta == {"general.basename": "Foo"}
|
||||
|
||||
|
||||
def test_metadata_is_cached(tmp_path: Path):
|
||||
"""Cache invalidates on size change."""
|
||||
p = _write_synthetic_gguf(
|
||||
tmp_path / "model.gguf",
|
||||
{"general.basename": "First"},
|
||||
)
|
||||
first = read_gguf_general_metadata(str(p))
|
||||
assert first == {"general.basename": "First"}
|
||||
# Force size change so the (path, mtime, size) key invalidates.
|
||||
_write_synthetic_gguf(
|
||||
tmp_path / "model.gguf",
|
||||
{"general.basename": "Second", "general.organization": "X"},
|
||||
)
|
||||
second = read_gguf_general_metadata(str(p))
|
||||
assert second == {"general.basename": "Second", "general.organization": "X"}
|
||||
|
||||
|
||||
# --- is_mmproj_by_metadata --------------------------------------------
|
||||
|
||||
|
||||
def test_is_mmproj_by_metadata_signals():
|
||||
assert is_mmproj_by_metadata({"general.type": "mmproj"}) is True
|
||||
assert is_mmproj_by_metadata({"general.type": "MMProj"}) is True
|
||||
assert is_mmproj_by_metadata({"general.type": "model"}) is False
|
||||
assert is_mmproj_by_metadata({"general.basename": "foo"}) is None
|
||||
assert is_mmproj_by_metadata({}) is None
|
||||
assert is_mmproj_by_metadata(None) is None
|
||||
|
||||
|
||||
# --- pairing_score -----------------------------------------------------
|
||||
|
||||
|
||||
def test_pairing_score_base_model_url_match():
|
||||
weight = {
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
}
|
||||
mmproj = {
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
}
|
||||
assert pairing_score(weight, mmproj) == 100
|
||||
|
||||
|
||||
def test_pairing_score_base_model_url_mismatch():
|
||||
weight = {
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
}
|
||||
mmproj = {
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/google/gemma-3-9B",
|
||||
}
|
||||
assert pairing_score(weight, mmproj) == -1
|
||||
|
||||
|
||||
def test_pairing_score_base_model_url_trailing_slash_normalised():
|
||||
weight = {
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B/",
|
||||
}
|
||||
mmproj = {
|
||||
"general.base_model.0.repo_url": "https://huggingface.co/Qwen/Qwen3.5-9B",
|
||||
}
|
||||
assert pairing_score(weight, mmproj) == 100
|
||||
|
||||
|
||||
def test_pairing_score_basename_plus_org_fallback():
|
||||
weight = {
|
||||
"general.basename": "Nanonets-Ocr-S",
|
||||
"general.base_model.0.organization": "Nanonets",
|
||||
}
|
||||
mmproj = {
|
||||
"general.basename": "Nanonets-Ocr-S",
|
||||
"general.base_model.0.organization": "Nanonets",
|
||||
}
|
||||
assert pairing_score(weight, mmproj) == 80
|
||||
|
||||
|
||||
def test_pairing_score_basename_only_fallback():
|
||||
assert (
|
||||
pairing_score(
|
||||
{"general.basename": "Nanonets-Ocr-S"},
|
||||
{"general.basename": "Nanonets-Ocr-S"},
|
||||
)
|
||||
== 60
|
||||
)
|
||||
|
||||
|
||||
def test_pairing_score_no_overlap_returns_zero():
|
||||
"""One side empty: scorer punts to filename fallback."""
|
||||
assert pairing_score({"general.basename": "Foo"}, {}) == 0
|
||||
assert pairing_score({}, {"general.basename": "Foo"}) == 0
|
||||
assert pairing_score(None, {"general.basename": "Foo"}) == 0
|
||||
236
studio/backend/utils/models/gguf_metadata.py
Normal file
236
studio/backend/utils/models/gguf_metadata.py
Normal file
|
|
@ -0,0 +1,236 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""Free-function ``general.*`` reader for GGUF headers, used by
|
||||
``detect_mmproj_file`` to pair weights and projectors via
|
||||
``general.base_model.0.repo_url``. ~30 ms per file, cached by
|
||||
(path, mtime, size)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import struct
|
||||
import threading
|
||||
from pathlib import Path
|
||||
from typing import Dict, Optional, Tuple
|
||||
|
||||
from loggers import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
_GGUF_MAGIC = 0x46554747 # b"GGUF" LE u32
|
||||
|
||||
_WANTED_GENERAL_KEYS: frozenset[str] = frozenset(
|
||||
{
|
||||
"general.architecture",
|
||||
"general.type",
|
||||
"general.name",
|
||||
"general.basename",
|
||||
"general.organization",
|
||||
"general.size_label",
|
||||
"general.finetune",
|
||||
"general.base_model.0.name",
|
||||
"general.base_model.0.organization",
|
||||
"general.base_model.0.repo_url",
|
||||
"general.repo_url",
|
||||
"general.source.url",
|
||||
"general.source.repo_url",
|
||||
"general.source.huggingface.repository",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
# Cache failed parses too so a broken file is not retried each scan.
|
||||
_CacheKey = Tuple[str, int, int]
|
||||
_METADATA_CACHE: Dict[_CacheKey, Optional[Dict[str, str]]] = {}
|
||||
_CACHE_LOCK = threading.Lock()
|
||||
_CACHE_MAX_ENTRIES = 4096
|
||||
|
||||
|
||||
def _cache_key(path: str) -> Optional[_CacheKey]:
|
||||
try:
|
||||
st = os.stat(path)
|
||||
except OSError:
|
||||
return None
|
||||
try:
|
||||
resolved = str(Path(path).resolve())
|
||||
except OSError:
|
||||
resolved = str(path)
|
||||
return (resolved, st.st_mtime_ns, st.st_size)
|
||||
|
||||
|
||||
def read_gguf_general_metadata(path: str) -> Optional[Dict[str, str]]:
|
||||
"""Return ``general.*`` strings from a GGUF header, or ``None`` if
|
||||
the file is missing, unreadable, or not a GGUF. ``{}`` means the
|
||||
file is valid but carries none of the wanted keys."""
|
||||
key = _cache_key(path)
|
||||
if key is None:
|
||||
return None
|
||||
with _CACHE_LOCK:
|
||||
if key in _METADATA_CACHE:
|
||||
return _METADATA_CACHE[key]
|
||||
result = _parse_gguf_header(path)
|
||||
with _CACHE_LOCK:
|
||||
# Arbitrary eviction; header reads are cheap so true LRU is overkill.
|
||||
while len(_METADATA_CACHE) >= _CACHE_MAX_ENTRIES:
|
||||
try:
|
||||
_METADATA_CACHE.pop(next(iter(_METADATA_CACHE)))
|
||||
except StopIteration:
|
||||
break
|
||||
_METADATA_CACHE[key] = result
|
||||
return result
|
||||
|
||||
|
||||
def _parse_gguf_header(path: str) -> Optional[Dict[str, str]]:
|
||||
out: Dict[str, str] = {}
|
||||
try:
|
||||
with open(path, "rb") as f:
|
||||
head = f.read(24)
|
||||
if len(head) < 24:
|
||||
return None
|
||||
magic, _version, _tcount, kv_count = struct.unpack("<IIQQ", head)
|
||||
if magic != _GGUF_MAGIC:
|
||||
return None
|
||||
|
||||
for _ in range(kv_count):
|
||||
try:
|
||||
klen_bytes = f.read(8)
|
||||
if len(klen_bytes) < 8:
|
||||
break
|
||||
klen = struct.unpack("<Q", klen_bytes)[0]
|
||||
if klen > 1 << 20: # 1 MB sanity bound
|
||||
break
|
||||
kbytes = f.read(klen)
|
||||
if len(kbytes) < klen:
|
||||
break
|
||||
key = kbytes.decode("utf-8", "replace")
|
||||
vt_bytes = f.read(4)
|
||||
if len(vt_bytes) < 4:
|
||||
break
|
||||
vtype = struct.unpack("<I", vt_bytes)[0]
|
||||
|
||||
if vtype == 8 and key in _WANTED_GENERAL_KEYS:
|
||||
slen_bytes = f.read(8)
|
||||
if len(slen_bytes) < 8:
|
||||
break
|
||||
slen = struct.unpack("<Q", slen_bytes)[0]
|
||||
if slen > 1 << 22: # 4 MB sanity bound
|
||||
break
|
||||
sbytes = f.read(slen)
|
||||
if len(sbytes) < slen:
|
||||
break
|
||||
out[key] = sbytes.decode("utf-8", "replace")
|
||||
else:
|
||||
if not _skip_gguf_value(f, vtype):
|
||||
break
|
||||
except (struct.error, UnicodeDecodeError):
|
||||
break
|
||||
except OSError as e:
|
||||
logger.debug(f"read_gguf_general_metadata: cannot open {path}: {e}")
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.debug(f"read_gguf_general_metadata: parse failure on {path}: {e}")
|
||||
return None
|
||||
return out
|
||||
|
||||
|
||||
# Strings (8) and arrays (9) are handled inline.
|
||||
_FIXED_VTYPE_SIZES: Dict[int, int] = {
|
||||
0: 1, # uint8
|
||||
1: 1, # int8
|
||||
2: 2, # uint16
|
||||
3: 2, # int16
|
||||
4: 4, # uint32
|
||||
5: 4, # int32
|
||||
6: 4, # float32
|
||||
7: 1, # bool
|
||||
10: 8, # uint64
|
||||
11: 8, # int64
|
||||
12: 8, # float64
|
||||
}
|
||||
|
||||
|
||||
def _skip_gguf_value(f, vtype: int) -> bool:
|
||||
"""Advance past one GGUF value. ``f.seek(.., 1)`` past EOF is legal
|
||||
on a regular file so truncation is detected on the next read; we
|
||||
only return False for unknown types or sanity-bound overflow."""
|
||||
if vtype == 8: # STRING
|
||||
slen_bytes = f.read(8)
|
||||
if len(slen_bytes) < 8:
|
||||
return False
|
||||
slen = struct.unpack("<Q", slen_bytes)[0]
|
||||
if slen > 1 << 30: # 1 GB sanity bound
|
||||
return False
|
||||
f.seek(slen, 1)
|
||||
return True
|
||||
if vtype == 9: # ARRAY
|
||||
head = f.read(12)
|
||||
if len(head) < 12:
|
||||
return False
|
||||
atype, alen = struct.unpack("<IQ", head)
|
||||
if alen > 1 << 30:
|
||||
return False
|
||||
if atype == 8:
|
||||
for _ in range(alen):
|
||||
slen_bytes = f.read(8)
|
||||
if len(slen_bytes) < 8:
|
||||
return False
|
||||
slen = struct.unpack("<Q", slen_bytes)[0]
|
||||
if slen > 1 << 30:
|
||||
return False
|
||||
f.seek(slen, 1)
|
||||
return True
|
||||
sz = _FIXED_VTYPE_SIZES.get(atype)
|
||||
if sz is None:
|
||||
return False
|
||||
f.seek(sz * alen, 1)
|
||||
return True
|
||||
sz = _FIXED_VTYPE_SIZES.get(vtype)
|
||||
if sz is None:
|
||||
return False
|
||||
f.seek(sz, 1)
|
||||
return True
|
||||
|
||||
|
||||
def is_mmproj_by_metadata(meta: Optional[Dict[str, str]]) -> Optional[bool]:
|
||||
"""True/False from ``general.type``; None means fall back to filename."""
|
||||
if not meta:
|
||||
return None
|
||||
t = meta.get("general.type")
|
||||
if t is None:
|
||||
return None
|
||||
return t.lower() == "mmproj"
|
||||
|
||||
|
||||
def pairing_score(
|
||||
weight_meta: Optional[Dict[str, str]],
|
||||
mmproj_meta: Optional[Dict[str, str]],
|
||||
) -> int:
|
||||
"""Pairing confidence: 100 = base_model URL match, 80 = basename + org,
|
||||
60 = basename, -1 = definitive mismatch, 0 = decide from filename."""
|
||||
if not weight_meta or not mmproj_meta:
|
||||
return 0
|
||||
|
||||
w_url = weight_meta.get("general.base_model.0.repo_url")
|
||||
p_url = mmproj_meta.get("general.base_model.0.repo_url")
|
||||
if w_url and p_url:
|
||||
return 100 if w_url.strip().rstrip("/") == p_url.strip().rstrip("/") else -1
|
||||
|
||||
w_base = weight_meta.get("general.basename")
|
||||
p_base = mmproj_meta.get("general.basename")
|
||||
w_org = weight_meta.get("general.base_model.0.organization") or weight_meta.get(
|
||||
"general.organization"
|
||||
)
|
||||
p_org = mmproj_meta.get("general.base_model.0.organization") or mmproj_meta.get(
|
||||
"general.organization"
|
||||
)
|
||||
if w_base and p_base and w_org and p_org:
|
||||
if w_base.lower() == p_base.lower() and w_org.lower() == p_org.lower():
|
||||
return 80
|
||||
return -1
|
||||
|
||||
if w_base and p_base:
|
||||
return 60 if w_base.lower() == p_base.lower() else -1
|
||||
|
||||
return 0
|
||||
|
|
@ -19,6 +19,11 @@ from utils.paths import (
|
|||
resolve_export_dir,
|
||||
)
|
||||
from utils.utils import without_hf_auth
|
||||
from utils.models.gguf_metadata import (
|
||||
is_mmproj_by_metadata,
|
||||
pairing_score,
|
||||
read_gguf_general_metadata,
|
||||
)
|
||||
import structlog
|
||||
from loggers import get_logger
|
||||
import os
|
||||
|
|
@ -801,12 +806,15 @@ _AUDIO_TOKEN_PATTERNS = {
|
|||
"whisper": lambda tokens: "<|startoftranscript|>" in tokens,
|
||||
"audio_vlm": lambda tokens: "<audio_soft_token>" in tokens,
|
||||
"bicodec": lambda tokens: any(t.startswith("<|bicodec_") for t in tokens),
|
||||
"dac": lambda tokens: "<|audio_start|>" in tokens
|
||||
and "<|audio_end|>" in tokens
|
||||
and "<|text_start|>" in tokens
|
||||
and "<|text_end|>" in tokens,
|
||||
"snac": lambda tokens: sum(1 for t in tokens if t.startswith("<custom_token_"))
|
||||
> 10000,
|
||||
"dac": lambda tokens: (
|
||||
"<|audio_start|>" in tokens
|
||||
and "<|audio_end|>" in tokens
|
||||
and "<|text_start|>" in tokens
|
||||
and "<|text_end|>" in tokens
|
||||
),
|
||||
"snac": lambda tokens: (
|
||||
sum(1 for t in tokens if t.startswith("<custom_token_")) > 10000
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -913,6 +921,85 @@ def _is_mmproj(filename: str) -> bool:
|
|||
return "mmproj" in filename.lower()
|
||||
|
||||
|
||||
# Family tokens for #5347's filename fallback. Lowercase. Order does not
|
||||
# matter (see ``_detect_family_token``).
|
||||
_MODEL_FAMILY_TOKENS: tuple[str, ...] = (
|
||||
"qwen",
|
||||
"gemma",
|
||||
"llama",
|
||||
"mistral",
|
||||
"ministral",
|
||||
"magistral",
|
||||
"devstral",
|
||||
"phi",
|
||||
"deepseek",
|
||||
"internvl",
|
||||
"minicpm",
|
||||
"llava",
|
||||
"glm",
|
||||
"yi",
|
||||
"command-r",
|
||||
"molmo",
|
||||
"pixtral",
|
||||
"smolvlm",
|
||||
"moondream",
|
||||
"granite",
|
||||
"ovis",
|
||||
"nemotron",
|
||||
"kimi",
|
||||
"nanonets",
|
||||
"cosmos",
|
||||
"mimo",
|
||||
"apriel",
|
||||
"lfm",
|
||||
)
|
||||
|
||||
|
||||
# Word-bounded match: any letter on either side disqualifies. Stops
|
||||
# ``phi`` matching ``sapphire``, ``yi`` matching ``tiny``, etc.
|
||||
_FAMILY_TOKEN_RE_CACHE: Dict[str, "_re.Pattern[str]"] = {}
|
||||
|
||||
|
||||
def _family_token_re(token: str) -> "_re.Pattern[str]":
|
||||
pat = _FAMILY_TOKEN_RE_CACHE.get(token)
|
||||
if pat is None:
|
||||
pat = _re.compile(rf"(?:^|[^a-z])({_re.escape(token)})(?:[^a-z]|$)")
|
||||
_FAMILY_TOKEN_RE_CACHE[token] = pat
|
||||
return pat
|
||||
|
||||
|
||||
def _detect_family_token(filename: str) -> Optional[str]:
|
||||
"""Leftmost-position match; ties prefer the longer token."""
|
||||
name = filename.lower()
|
||||
best: Optional[tuple[int, int, str]] = None # (start, -len, token)
|
||||
for token in _MODEL_FAMILY_TOKENS:
|
||||
m = _family_token_re(token).search(name)
|
||||
if m is None:
|
||||
continue
|
||||
key = (m.start(1), -len(token), token)
|
||||
if best is None or key < best:
|
||||
best = key
|
||||
return None if best is None else best[2]
|
||||
|
||||
|
||||
def mmproj_matches_model_family(model_path: str, mmproj_path: str) -> bool:
|
||||
"""Defense-in-depth guard for the launcher: True unless both filenames
|
||||
carry recognised family tokens that disagree."""
|
||||
model_fam = _detect_family_token(Path(model_path).name)
|
||||
mmproj_fam = _detect_family_token(Path(mmproj_path).name)
|
||||
if model_fam is None or mmproj_fam is None:
|
||||
return True
|
||||
return model_fam == mmproj_fam
|
||||
|
||||
|
||||
def _shared_prefix_len(a: str, b: str) -> int:
|
||||
n = min(len(a), len(b))
|
||||
for i in range(n):
|
||||
if a[i] != b[i]:
|
||||
return i
|
||||
return n
|
||||
|
||||
|
||||
def _is_gguf_filename(filename: str) -> bool:
|
||||
return filename.lower().endswith(".gguf")
|
||||
|
||||
|
|
@ -927,33 +1014,18 @@ def _iter_gguf_files(directory: Path, recursive: bool = False):
|
|||
|
||||
|
||||
def detect_mmproj_file(path: str, search_root: Optional[str] = None) -> Optional[str]:
|
||||
"""
|
||||
Find the mmproj (vision projection) GGUF file for a given model.
|
||||
"""Find the mmproj GGUF for a model.
|
||||
|
||||
Args:
|
||||
path: Directory to search — or a .gguf file (uses its parent dir
|
||||
as the starting point).
|
||||
search_root: Optional outer directory that should also be scanned
|
||||
(and any directory between it and ``path``). This handles
|
||||
local layouts where the model weights live in a quant-named
|
||||
subdir (``snapshot/BF16/foo.gguf``) but the mmproj sits at
|
||||
the snapshot root (``snapshot/mmproj-BF16.gguf``). When
|
||||
``None``, only the immediate parent dir is scanned, matching
|
||||
the historical behavior.
|
||||
|
||||
Returns:
|
||||
Full path to the mmproj .gguf file, or None if not found.
|
||||
"""
|
||||
``path``: directory or a .gguf file. ``search_root``: optional ancestor
|
||||
to also walk (snapshot layouts where the weight is in ``snapshot/BF16/``
|
||||
but the projector sits at ``snapshot/``). Returns the projector path or
|
||||
``None``."""
|
||||
p = Path(path)
|
||||
start_dir = p.parent if p.is_file() else p
|
||||
if not start_dir.is_dir():
|
||||
return None
|
||||
|
||||
# Build the list of dirs to scan: immediate dir first, then walk up
|
||||
# to (and including) ``search_root`` if it is an ancestor. We walk
|
||||
# incrementally rather than recursing into ``search_root`` so we
|
||||
# don't accidentally pick up an mmproj from a sibling subdir
|
||||
# belonging to a different model variant.
|
||||
# Walk incrementally so a sibling subdir's mmproj cannot leak in.
|
||||
seen: set[Path] = set()
|
||||
scan_order: list[Path] = []
|
||||
|
||||
|
|
@ -969,12 +1041,7 @@ def detect_mmproj_file(path: str, search_root: Optional[str] = None) -> Optional
|
|||
|
||||
_add(start_dir)
|
||||
|
||||
# When ``path`` is a symlink (e.g. Ollama's ``.studio_links/...gguf``
|
||||
# -> ``blobs/sha256-...``), the symlink's parent directory rarely
|
||||
# contains the mmproj sibling; the real mmproj file lives next to
|
||||
# the symlink target. Add the target's parent to the scan so vision
|
||||
# GGUFs that are surfaced via symlinks are still recognised as
|
||||
# vision models.
|
||||
# Ollama's .studio_links/foo.gguf -> blobs/sha256-...: also scan target dir.
|
||||
try:
|
||||
if p.is_symlink() and p.is_file():
|
||||
target_parent = p.resolve().parent
|
||||
|
|
@ -986,14 +1053,12 @@ def detect_mmproj_file(path: str, search_root: Optional[str] = None) -> Optional
|
|||
try:
|
||||
root_resolved = Path(search_root).resolve()
|
||||
start_resolved = start_dir.resolve()
|
||||
# Only walk if start_dir is inside (or equal to) search_root.
|
||||
if root_resolved == start_resolved or (
|
||||
start_resolved.is_relative_to(root_resolved)
|
||||
if hasattr(start_resolved, "is_relative_to")
|
||||
else str(start_resolved).startswith(str(root_resolved) + "/")
|
||||
):
|
||||
cur = start_resolved
|
||||
# Walk up from start_dir to (and including) root_resolved.
|
||||
while cur != root_resolved and cur.parent != cur:
|
||||
cur = cur.parent
|
||||
_add(cur)
|
||||
|
|
@ -1002,11 +1067,66 @@ def detect_mmproj_file(path: str, search_root: Optional[str] = None) -> Optional
|
|||
except OSError:
|
||||
pass
|
||||
|
||||
candidates: list[Path] = []
|
||||
seen_resolved: set[Path] = set()
|
||||
for d in scan_order:
|
||||
for f in _iter_gguf_files(d):
|
||||
if _is_mmproj(f.name):
|
||||
return str(f.resolve())
|
||||
return None
|
||||
try:
|
||||
resolved = f.resolve()
|
||||
except OSError:
|
||||
continue
|
||||
if resolved in seen_resolved:
|
||||
continue
|
||||
# Prefer ``general.type=='mmproj'``; fall back to filename.
|
||||
meta = read_gguf_general_metadata(str(resolved))
|
||||
by_meta = is_mmproj_by_metadata(meta)
|
||||
if by_meta is True or (by_meta is None and _is_mmproj(f.name)):
|
||||
seen_resolved.add(resolved)
|
||||
candidates.append(resolved)
|
||||
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
# Directory path: no model name to compare against; legacy behaviour.
|
||||
if not p.is_file():
|
||||
return str(candidates[0])
|
||||
|
||||
# Stage 1: GGUF metadata. Stage 2: filename family token (#5347).
|
||||
model_stem = p.stem.lower()
|
||||
model_family = _detect_family_token(p.name)
|
||||
weight_meta = read_gguf_general_metadata(str(p))
|
||||
|
||||
scored: list[tuple[int, Path]] = []
|
||||
for c in candidates:
|
||||
cand_meta = read_gguf_general_metadata(str(c))
|
||||
meta_score = pairing_score(weight_meta, cand_meta)
|
||||
if meta_score == -1:
|
||||
logger.info(f"detect_mmproj_file: dropped {c.name} (metadata mismatch)")
|
||||
continue
|
||||
if meta_score == 0 and model_family is not None:
|
||||
# Unrecognised candidate family is a wildcard (``mmproj-F16.gguf``).
|
||||
cand_family = _detect_family_token(c.name)
|
||||
if cand_family is not None and cand_family != model_family:
|
||||
logger.info(
|
||||
f"detect_mmproj_file: dropped {c.name} "
|
||||
f"(filename family {cand_family!r} vs model {model_family!r})"
|
||||
)
|
||||
continue
|
||||
scored.append((meta_score, c))
|
||||
|
||||
if not scored:
|
||||
return None
|
||||
|
||||
# Score first, then longest shared prefix, then shorter stem.
|
||||
best = max(
|
||||
scored,
|
||||
key = lambda sc: (
|
||||
sc[0],
|
||||
_shared_prefix_len(model_stem, sc[1].stem.lower()),
|
||||
-len(sc[1].stem),
|
||||
),
|
||||
)
|
||||
return str(best[1])
|
||||
|
||||
|
||||
def detect_gguf_model(path: str) -> Optional[str]:
|
||||
|
|
@ -1360,7 +1480,7 @@ def detect_gguf_model_remote(
|
|||
if attempt < 2:
|
||||
time.sleep(2**attempt)
|
||||
logger.warning(
|
||||
f"Could not check GGUF files for '{repo_id}' after 3 attempts: " f"{last_err}"
|
||||
f"Could not check GGUF files for '{repo_id}' after 3 attempts: {last_err}"
|
||||
)
|
||||
return None
|
||||
|
||||
|
|
|
|||
216
tests/test_public_api_surface.py
Normal file
216
tests/test_public_api_surface.py
Normal file
|
|
@ -0,0 +1,216 @@
|
|||
# Unsloth - 2x faster, 60% less VRAM LLM training and finetuning
|
||||
# Copyright 2023-present Daniel Han-Chen, Michael Han-Chen & the Unsloth team. All rights reserved.
|
||||
#
|
||||
# This program is free software: you can redistribute it and/or modify
|
||||
# it under the terms of the GNU Lesser General Public License as published by
|
||||
# the Free Software Foundation, either version 3 of the License, or
|
||||
# (at your option) any later version.
|
||||
#
|
||||
# This program is distributed in the hope that it will be useful,
|
||||
# but WITHOUT ANY WARRANTY; without even the implied warranty of
|
||||
# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
|
||||
# GNU Lesser General Public License for more details.
|
||||
|
||||
"""Public-API surface drift detectors for unsloth itself.
|
||||
|
||||
Companion to tests/test_import_fixes_drift.py: that file catches drift
|
||||
in THIRD-PARTY libraries (transformers / trl / triton / peft / etc.)
|
||||
that unsloth's import_fixes patches around. This file catches drift in
|
||||
unsloth's OWN public-surface API -- the top-10 symbols and classmethods
|
||||
that the unslothai/notebooks tree (and therefore every user on Colab)
|
||||
calls. If a refactor on this repo renames FastLanguageModel.from_pretrained
|
||||
or drops one of the documented kwargs, the test fires DRIFT DETECTED
|
||||
here BEFORE the breakage reaches users.
|
||||
|
||||
Call-site counts measured against unslothai/notebooks @ main:
|
||||
FastLanguageModel.from_pretrained 506
|
||||
FastLanguageModel.for_inference 370
|
||||
FastLanguageModel.get_peft_model 304
|
||||
FastVisionModel.for_inference 183
|
||||
FastVisionModel.from_pretrained 176
|
||||
FastVisionModel.get_peft_model 99
|
||||
FastVisionModel.for_training 60
|
||||
FastModel.from_pretrained 103
|
||||
FastModel.get_peft_model 67
|
||||
|
||||
Mirrors the unsloth-zoo / unsloth drift-detector skeleton:
|
||||
``pytest.importorskip("unsloth")`` to gate, assert the healthy upstream
|
||||
shape, ``pytest.fail("DRIFT DETECTED: ...")`` (never ``pytest.skip``) on
|
||||
regression so the matrix cell goes red.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _signature_param_names(callable_obj) -> set[str]:
|
||||
try:
|
||||
sig = inspect.signature(callable_obj)
|
||||
except (TypeError, ValueError):
|
||||
return set()
|
||||
return set(sig.parameters)
|
||||
|
||||
|
||||
def _accepts(callable_obj, kwargs: set[str]) -> tuple[bool, set[str]]:
|
||||
"""True if every name in ``kwargs`` is either a named parameter on
|
||||
``callable_obj`` OR the callable's signature has a ``**kwargs``
|
||||
catch-all. Returns (ok, missing_set)."""
|
||||
try:
|
||||
sig = inspect.signature(callable_obj)
|
||||
except (TypeError, ValueError):
|
||||
return True, set()
|
||||
params = sig.parameters
|
||||
has_var_kw = any(p.kind == inspect.Parameter.VAR_KEYWORD for p in params.values())
|
||||
if has_var_kw:
|
||||
return True, set()
|
||||
missing = kwargs - set(params)
|
||||
return (not missing), missing
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# FastLanguageModel: the headline class. 506 from_pretrained + 370
|
||||
# for_inference + 304 get_peft_model call sites across the notebooks.
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
def test_fast_language_model_class_present():
|
||||
unsloth = pytest.importorskip("unsloth")
|
||||
if not hasattr(unsloth, "FastLanguageModel"):
|
||||
pytest.fail(
|
||||
"DRIFT DETECTED: unsloth.FastLanguageModel is missing; every "
|
||||
"LoRA notebook fails at the first import cell."
|
||||
)
|
||||
|
||||
|
||||
def test_fast_language_model_from_pretrained_kwargs():
|
||||
"""from_pretrained must accept the canonical kwargs the notebooks pass."""
|
||||
unsloth = pytest.importorskip("unsloth")
|
||||
required = {"model_name", "max_seq_length", "dtype", "load_in_4bit"}
|
||||
ok, missing = _accepts(unsloth.FastLanguageModel.from_pretrained, required)
|
||||
if not ok:
|
||||
pytest.fail(
|
||||
f"DRIFT DETECTED: FastLanguageModel.from_pretrained dropped "
|
||||
f"kwargs {sorted(missing)}; 506 notebook call sites would "
|
||||
f"crash with TypeError."
|
||||
)
|
||||
|
||||
|
||||
def test_fast_language_model_get_peft_model_kwargs():
|
||||
unsloth = pytest.importorskip("unsloth")
|
||||
required = {
|
||||
"r",
|
||||
"lora_alpha",
|
||||
"lora_dropout",
|
||||
"target_modules",
|
||||
"bias",
|
||||
"use_gradient_checkpointing",
|
||||
"random_state",
|
||||
}
|
||||
ok, missing = _accepts(unsloth.FastLanguageModel.get_peft_model, required)
|
||||
if not ok:
|
||||
pytest.fail(
|
||||
f"DRIFT DETECTED: FastLanguageModel.get_peft_model dropped "
|
||||
f"kwargs {sorted(missing)}; 304 notebook call sites would crash."
|
||||
)
|
||||
|
||||
|
||||
def test_fast_language_model_for_inference_callable():
|
||||
unsloth = pytest.importorskip("unsloth")
|
||||
if not callable(getattr(unsloth.FastLanguageModel, "for_inference", None)):
|
||||
pytest.fail(
|
||||
"DRIFT DETECTED: FastLanguageModel.for_inference is missing; "
|
||||
"370 inference-cell call sites would crash."
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# FastVisionModel: 183 + 176 + 99 + 60 call sites across vision notebooks.
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
def test_fast_vision_model_class_and_methods():
|
||||
unsloth = pytest.importorskip("unsloth")
|
||||
if not hasattr(unsloth, "FastVisionModel"):
|
||||
pytest.fail(
|
||||
"DRIFT DETECTED: unsloth.FastVisionModel is missing; every "
|
||||
"vision fine-tuning notebook fails at import."
|
||||
)
|
||||
cls = unsloth.FastVisionModel
|
||||
missing = [
|
||||
m
|
||||
for m in ("from_pretrained", "get_peft_model", "for_inference", "for_training")
|
||||
if not callable(getattr(cls, m, None))
|
||||
]
|
||||
if missing:
|
||||
pytest.fail(f"DRIFT DETECTED: FastVisionModel is missing methods {missing}.")
|
||||
|
||||
|
||||
def test_fast_vision_model_get_peft_model_vision_kwargs():
|
||||
"""Vision-specific kwargs the notebooks pass on the vision LoRA path."""
|
||||
unsloth = pytest.importorskip("unsloth")
|
||||
required = {
|
||||
"finetune_vision_layers",
|
||||
"finetune_language_layers",
|
||||
"finetune_attention_modules",
|
||||
"finetune_mlp_modules",
|
||||
}
|
||||
ok, missing = _accepts(unsloth.FastVisionModel.get_peft_model, required)
|
||||
if not ok:
|
||||
pytest.fail(
|
||||
f"DRIFT DETECTED: FastVisionModel.get_peft_model dropped "
|
||||
f"vision kwargs {sorted(missing)}."
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# FastModel: the modern unified entry point. 103 + 67 call sites.
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
def test_fast_model_class_and_methods():
|
||||
unsloth = pytest.importorskip("unsloth")
|
||||
if not hasattr(unsloth, "FastModel"):
|
||||
pytest.fail(
|
||||
"DRIFT DETECTED: unsloth.FastModel is missing; the modern "
|
||||
"unified entry point used by 100+ notebooks would crash."
|
||||
)
|
||||
missing = [
|
||||
m
|
||||
for m in ("from_pretrained", "get_peft_model")
|
||||
if not callable(getattr(unsloth.FastModel, m, None))
|
||||
]
|
||||
if missing:
|
||||
pytest.fail(f"DRIFT DETECTED: FastModel is missing methods {missing}.")
|
||||
|
||||
|
||||
def test_fast_model_from_pretrained_kwargs():
|
||||
unsloth = pytest.importorskip("unsloth")
|
||||
required = {"model_name", "max_seq_length", "dtype", "load_in_4bit"}
|
||||
ok, missing = _accepts(unsloth.FastModel.from_pretrained, required)
|
||||
if not ok:
|
||||
pytest.fail(
|
||||
f"DRIFT DETECTED: FastModel.from_pretrained dropped kwargs "
|
||||
f"{sorted(missing)}; 103 notebook call sites would crash."
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Bf16 helper alias (renamed once already; keep both accepted).
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
def test_is_bf16_supported_or_alias_callable():
|
||||
"""48 notebook import sites for is_bf16_supported plus 8 for the
|
||||
legacy is_bfloat16_supported alias. Either must remain importable."""
|
||||
unsloth = pytest.importorskip("unsloth")
|
||||
has_new = callable(getattr(unsloth, "is_bf16_supported", None))
|
||||
has_old = callable(getattr(unsloth, "is_bfloat16_supported", None))
|
||||
if not (has_new or has_old):
|
||||
pytest.fail(
|
||||
"DRIFT DETECTED: neither unsloth.is_bf16_supported nor "
|
||||
"unsloth.is_bfloat16_supported is callable; dtype probing "
|
||||
"in 50+ notebooks fails."
|
||||
)
|
||||
|
|
@ -209,7 +209,7 @@ del fix_peft_transformers_weight_conversion_import
|
|||
del patch_peft_weight_converter_compatibility
|
||||
|
||||
# Torch 2.4 has including_emulation
|
||||
if DEVICE_TYPE == "cuda":
|
||||
if DEVICE_TYPE == "cuda" and torch.cuda.is_available():
|
||||
major_version, minor_version = torch.cuda.get_device_capability()
|
||||
SUPPORTS_BFLOAT16 = major_version >= 8
|
||||
|
||||
|
|
@ -233,12 +233,18 @@ elif DEVICE_TYPE == "xpu":
|
|||
# torch.xpu.is_bf16_supported() does not have including_emulation
|
||||
# set SUPPORTS_BFLOAT16 as torch.xpu.is_bf16_supported()
|
||||
SUPPORTS_BFLOAT16 = torch.xpu.is_bf16_supported()
|
||||
else:
|
||||
# CPU-only CI under UNSLOTH_ALLOW_CPU=1. We can't probe device
|
||||
# capability, so assume no bf16 -- training won't run on this host
|
||||
# anyway, this branch only exists to let `import unsloth.trainer`
|
||||
# succeed for source-inspection tests.
|
||||
SUPPORTS_BFLOAT16 = False
|
||||
|
||||
# For Gradio HF Spaces?
|
||||
# if "SPACE_AUTHOR_NAME" not in os.environ and "SPACE_REPO_NAME" not in os.environ:
|
||||
import triton
|
||||
|
||||
if DEVICE_TYPE == "cuda":
|
||||
if DEVICE_TYPE == "cuda" and torch.cuda.is_available():
|
||||
libcuda_dirs = lambda: None
|
||||
if Version(triton.__version__) >= Version("3.0.0"):
|
||||
try:
|
||||
|
|
@ -349,5 +355,10 @@ from unsloth_zoo.rl_environments import (
|
|||
launch_openenv,
|
||||
)
|
||||
|
||||
# Patch TRL trainers for backwards compatibility
|
||||
_patch_trl_trainer()
|
||||
# Patch TRL trainers for backwards compatibility.
|
||||
# Skipped under UNSLOTH_ALLOW_CPU=1 (CPU-only CI) because rebinding
|
||||
# trl.SFTTrainer.__init__ to a generic wrapper changes
|
||||
# inspect.getsource(SFTTrainer.__init__) and corrupts downstream
|
||||
# drift detectors that anchor on the pristine upstream source.
|
||||
if os.environ.get("UNSLOTH_ALLOW_CPU", "0") != "1":
|
||||
_patch_trl_trainer()
|
||||
|
|
|
|||
|
|
@ -63,6 +63,10 @@ def get_device_type():
|
|||
# Check torch.accelerator
|
||||
if hasattr(torch, "accelerator"):
|
||||
if not torch.accelerator.is_available():
|
||||
# Test-only CPU fallback. The env var is read exactly once per
|
||||
# process because get_device_type is @functools.cache'd.
|
||||
if os.environ.get("UNSLOTH_ALLOW_CPU", "0") == "1":
|
||||
return "cuda"
|
||||
raise NotImplementedError(
|
||||
"Unsloth cannot find any torch accelerator? You need a GPU."
|
||||
)
|
||||
|
|
@ -73,6 +77,8 @@ def get_device_type():
|
|||
f"But `torch.accelerator.current_accelerator()` works with it being = `{accelerator}`\n"
|
||||
f"Please reinstall torch - it's most likely broken :("
|
||||
)
|
||||
if os.environ.get("UNSLOTH_ALLOW_CPU", "0") == "1":
|
||||
return "cuda"
|
||||
raise NotImplementedError(
|
||||
"Unsloth currently only works on NVIDIA, AMD and Intel GPUs."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1193,7 +1193,7 @@ SUPPORTS_BFLOAT16 = False
|
|||
HAS_FLASH_ATTENTION = False
|
||||
HAS_FLASH_ATTENTION_SOFTCAPPING = False
|
||||
|
||||
if DEVICE_TYPE == "cuda":
|
||||
if DEVICE_TYPE == "cuda" and torch.cuda.is_available():
|
||||
major_version, minor_version = torch.cuda.get_device_capability()
|
||||
torch.cuda.get_device_capability = functools.cache(torch.cuda.get_device_capability)
|
||||
|
||||
|
|
|
|||
|
|
@ -2270,6 +2270,11 @@ def patch_trl_vllm_generation():
|
|||
def PatchFastRL(algorithm = None, FastLanguageModel = None):
|
||||
if FastLanguageModel is not None:
|
||||
PatchRL(FastLanguageModel)
|
||||
# Under UNSLOTH_ALLOW_CPU=1 (CPU-only CI), skip TRL trainer rewriting so
|
||||
# downstream `inspect.getsource(trl.SFTTrainer)` drift detectors see the
|
||||
# pristine upstream class, not the compiled Unsloth* wrappers.
|
||||
if os.environ.get("UNSLOTH_ALLOW_CPU", "0") == "1":
|
||||
return
|
||||
# Install the disable_gradient_checkpointing noop BEFORE
|
||||
# patch_trl_rl_trainers. patch_trl_rl_trainers imports extra trl.* trainer
|
||||
# submodules while generating the compiled cache; any new trl.* modules
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue