Merge branch 'main' into fix-issue-5344-quantization-guardrail

This commit is contained in:
Daniel Han 2026-05-15 00:12:10 -07:00 committed by GitHub
commit 8bad7ac1b7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
38 changed files with 1790 additions and 117 deletions

View file

@ -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

View file

@ -35,6 +35,7 @@ Example:
git diff --name-only origin/main..HEAD \\
| xargs python scripts/verify_comment_only_diff.py --base origin/main
"""
from __future__ import annotations
import argparse
@ -49,7 +50,9 @@ import yaml
def _git_show(rev: str, path: str) -> str:
return subprocess.check_output(
["git", "show", f"{rev}:{path}"], text = True, stderr = subprocess.DEVNULL,
["git", "show", f"{rev}:{path}"],
text = True,
stderr = subprocess.DEVNULL,
)
@ -145,8 +148,7 @@ def _walk_yaml_diff(b: Any, a: Any, prefix: str = "") -> None:
elif isinstance(b, list):
if len(b) != len(a):
print(
f" list len at {prefix or '/'}: "
f"{len(b)} -> {len(a)}",
f" list len at {prefix or '/'}: " f"{len(b)} -> {len(a)}",
)
for i, (bi, ai) in enumerate(zip(b, a)):
_walk_yaml_diff(bi, ai, f"{prefix}[{i}]")

View file

@ -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

View file

@ -116,11 +116,11 @@ class MLXInferenceBackend:
)
try:
from unsloth_zoo.mlx_loader import FastMLXModel
from unsloth_zoo.mlx.loader import FastMLXModel
except ImportError as e:
raise ImportError(
"Unsloth: MLX inference requires unsloth-zoo with the MLX modules "
"(unsloth_zoo.mlx_loader). Reinstall via install.sh on Apple Silicon."
"(unsloth_zoo.mlx.loader). Reinstall via install.sh on Apple Silicon."
) from e
model, tokenizer_or_processor = FastMLXModel.from_pretrained(

View file

@ -3208,6 +3208,9 @@ class UnslothTrainer:
if eval_steps_val > 0:
config_args["eval_strategy"] = "steps"
config_args["eval_steps"] = eval_steps_val
config_args["per_device_eval_batch_size"] = config_args[
"per_device_train_batch_size"
]
logger.info(
f"✅ Evaluation enabled: eval_steps={eval_steps_val} (fraction of total steps)\n"
)

View file

@ -214,6 +214,7 @@ class TrainingBackend:
"max_steps": kwargs.get("max_steps", 0),
"save_steps": kwargs.get("save_steps", 0),
"weight_decay": kwargs.get("weight_decay", 0.001),
"max_grad_norm": kwargs.get("max_grad_norm", 0.0),
"random_seed": kwargs.get("random_seed", 3407),
"packing": kwargs.get("packing", False),
"optim": kwargs.get("optim", "adamw_8bit"),

View file

@ -30,6 +30,7 @@ from utils.hardware import apply_gpu_ids
from utils.wheel_utils import (
direct_wheel_url,
flash_attn_wheel_url,
has_blackwell_gpu,
install_wheel,
probe_torch_wheel_env,
url_exists,
@ -313,6 +314,12 @@ def _should_try_runtime_flash_attn_install(max_seq_length: int) -> bool:
def _ensure_flash_attn_for_long_context(event_queue: Any, max_seq_length: int) -> None:
if not _should_try_runtime_flash_attn_install(max_seq_length):
return
if has_blackwell_gpu():
_send_status(
event_queue,
"Skipping flash-attn install: Blackwell GPU detected (sm_100+); no compatible prebuilt wheel",
)
return
installed = _install_package_wheel_first(
event_queue = event_queue,
@ -417,6 +424,55 @@ def _normalize_mlx_studio_scheduler(value):
return raw
def _resolve_mlx_local_dataset_files(file_paths: list) -> list[str]:
"""Resolve Studio local dataset uploads without importing the GPU trainer."""
from utils.paths import resolve_dataset_path
all_files: list[str] = []
for dataset_file in file_paths or []:
file_path = (
dataset_file
if os.path.isabs(dataset_file)
else str(resolve_dataset_path(dataset_file))
)
file_path_obj = Path(file_path)
if file_path_obj.is_dir():
parquet_dir = (
file_path_obj / "parquet-files"
if (file_path_obj / "parquet-files").exists()
else file_path_obj
)
parquet_files = sorted(parquet_dir.glob("*.parquet"))
if parquet_files:
all_files.extend(str(p) for p in parquet_files)
continue
candidates: list[Path] = []
for ext in (".json", ".jsonl", ".csv", ".parquet"):
candidates.extend(sorted(file_path_obj.glob(f"*{ext}")))
if candidates:
all_files.extend(str(c) for c in candidates)
continue
raise ValueError(f"No supported data files in directory: {file_path_obj}")
all_files.append(str(file_path_obj))
return all_files
def _mlx_local_dataset_loader_for_files(files: list[str]) -> str:
first_ext = Path(files[0]).suffix.lower()
if first_ext in (".json", ".jsonl"):
return "json"
if first_ext == ".csv":
return "csv"
if first_ext == ".parquet":
return "parquet"
raise ValueError(f"Unsupported dataset format: {files[0]}")
def _run_mlx_training(event_queue, stop_queue, config):
"""Self-contained MLX training path for Apple Silicon.
@ -442,8 +498,8 @@ def _run_mlx_training(event_queue, stop_queue, config):
import mlx.core as mx
try:
from unsloth_zoo.mlx_loader import FastMLXModel
from unsloth_zoo.mlx_trainer import (
from unsloth_zoo.mlx.loader import FastMLXModel
from unsloth_zoo.mlx.trainer import (
MLXTrainer,
MLXTrainingConfig,
train_on_responses_only,
@ -451,7 +507,7 @@ def _run_mlx_training(event_queue, stop_queue, config):
except ImportError as e:
raise ImportError(
"Unsloth: MLX training requires unsloth-zoo with the MLX modules "
"(unsloth_zoo.mlx_loader / unsloth_zoo.mlx_trainer). Reinstall via "
"(unsloth_zoo.mlx.loader / unsloth_zoo.mlx.trainer). Reinstall via "
"install.sh on Apple Silicon."
) from e
from datasets import load_dataset
@ -572,7 +628,6 @@ def _run_mlx_training(event_queue, stop_queue, config):
return ds
def _load_local(file_paths):
from core.training.trainer import UnslothTrainer
from datasets import load_from_disk
if len(file_paths) == 1:
@ -581,10 +636,10 @@ def _run_mlx_training(event_queue, stop_queue, config):
(p / "dataset_info.json").exists() or (p / "state.json").exists()
):
return load_from_disk(str(p))
all_files = UnslothTrainer._resolve_local_files(file_paths)
all_files = _resolve_mlx_local_dataset_files(file_paths)
if not all_files:
raise ValueError("No local dataset files found")
loader = UnslothTrainer._loader_for_files(all_files)
loader = _mlx_local_dataset_loader_for_files(all_files)
return load_dataset(loader, data_files = all_files, split = "train")
if hf_dataset:
@ -718,6 +773,10 @@ def _run_mlx_training(event_queue, stop_queue, config):
else:
eval_steps_val = int(eval_steps_val)
# MLX: value-clip grads to [-5, 5]; norm clipping disabled for compile-friendliness.
max_grad_norm = 0.0
max_grad_value = 5.0 # TODO: expose MLX grad-clip in Studio UI for power users
trainer = MLXTrainer(
model = model,
tokenizer = tokenizer,
@ -732,6 +791,8 @@ def _run_mlx_training(event_queue, stop_queue, config):
lr_scheduler_type = lr_scheduler_type,
optim = optim_name,
weight_decay = float(config.get("weight_decay", 0.001) or 0.001),
max_grad_norm = max_grad_norm,
max_grad_value = max_grad_value,
logging_steps = 1,
max_seq_length = max_seq_length,
seed = config.get("random_seed", 3407),
@ -820,7 +881,17 @@ def _run_mlx_training(event_queue, stop_queue, config):
# ── 9. Real-time progress callback ──
_send("status", status_message = f"Training {model_name}...")
def _on_step(step, total, loss, lr, tok_s, peak_gb, elapsed, num_tokens):
def _on_step(
step,
total,
loss,
lr,
tok_s,
peak_gb,
elapsed,
num_tokens,
grad_norm = None,
):
eta = (elapsed / step * (total - step)) if step > 0 else 0
_send(
"progress",
@ -831,7 +902,7 @@ def _run_mlx_training(event_queue, stop_queue, config):
total_steps = total,
elapsed_seconds = elapsed,
eta_seconds = max(0, eta),
grad_norm = None,
grad_norm = grad_norm,
num_tokens = num_tokens,
eval_loss = None,
status_message = None,
@ -846,6 +917,11 @@ def _run_mlx_training(event_queue, stop_queue, config):
"train/tokens_per_sec": tok_s,
"train/peak_gb": peak_gb,
"train/num_tokens": num_tokens,
**(
{"train/grad_norm": grad_norm}
if grad_norm is not None
else {}
),
},
step = step,
)
@ -857,6 +933,8 @@ def _run_mlx_training(event_queue, stop_queue, config):
tb_writer.add_scalar("train/learning_rate", lr, step)
tb_writer.add_scalar("train/tokens_per_sec", tok_s, step)
tb_writer.add_scalar("train/peak_gb", peak_gb, step)
if grad_norm is not None:
tb_writer.add_scalar("train/grad_norm", grad_norm, step)
except Exception:
pass

View file

@ -262,6 +262,11 @@ class TrainingStartRequest(BaseModel):
max_steps: Optional[int] = Field(None, description = "Maximum training steps")
save_steps: int = Field(100, description = "Steps between checkpoints")
weight_decay: float = Field(0.001, description = "Weight decay")
max_grad_norm: float = Field(
0.0,
ge = 0,
description = "Global gradient norm clipping threshold. Set 0 to disable.",
)
random_seed: int = Field(42, description = "Random seed")
packing: bool = Field(False, description = "Enable sequence packing")
optim: str = Field("adamw_8bit", description = "Optimizer")

View file

@ -215,6 +215,7 @@ async def start_training(
"max_steps": request.max_steps,
"save_steps": request.save_steps,
"weight_decay": request.weight_decay,
"max_grad_norm": request.max_grad_norm,
"random_seed": request.random_seed,
"packing": request.packing,
"optim": request.optim,

View 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())

View 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

View file

@ -56,11 +56,14 @@ def _install_fake_fast_mlx(monkeypatch, calls):
return _DummyModel(), _DummyTokenizer()
unsloth_zoo_pkg = types.ModuleType("unsloth_zoo")
mlx_loader = types.ModuleType("unsloth_zoo.mlx_loader")
mlx_pkg = types.ModuleType("unsloth_zoo.mlx")
mlx_loader = types.ModuleType("unsloth_zoo.mlx.loader")
mlx_loader.FastMLXModel = _FastMLXModel
unsloth_zoo_pkg.mlx_loader = mlx_loader
unsloth_zoo_pkg.mlx = mlx_pkg
mlx_pkg.loader = mlx_loader
monkeypatch.setitem(sys.modules, "unsloth_zoo", unsloth_zoo_pkg)
monkeypatch.setitem(sys.modules, "unsloth_zoo.mlx_loader", mlx_loader)
monkeypatch.setitem(sys.modules, "unsloth_zoo.mlx", mlx_pkg)
monkeypatch.setitem(sys.modules, "unsloth_zoo.mlx.loader", mlx_loader)
def test_mlx_inference_text_load_forwards_studio_settings(monkeypatch):

View file

@ -37,6 +37,7 @@ def _load_worker_module():
for name in (
"direct_wheel_url",
"flash_attn_wheel_url",
"has_blackwell_gpu",
"install_wheel",
"probe_torch_wheel_env",
"url_exists",

View file

@ -70,6 +70,48 @@ class TestTrainingRawSupport(unittest.TestCase):
self.assertTrue(config["load_in_4bit"])
self.assertEqual(config["embedding_learning_rate"], 1e-5)
def test_training_backend_forwards_grad_clipping_controls(self):
backend = TrainingBackend()
class DummyProcess:
pid = 12345
def start(self):
return None
class DummyThread:
def start(self):
return None
dummy_queue = object()
with (
patch(
"core.training.training.prepare_gpu_selection",
return_value = ([0], {"selection_mode": "auto"}),
),
patch(
"core.training.training._CTX.Queue",
side_effect = [dummy_queue, dummy_queue],
),
patch(
"core.training.training._CTX.Process", return_value = DummyProcess()
) as mock_process,
patch(
"core.training.training.threading.Thread",
return_value = DummyThread(),
),
):
backend.start_training(
job_id = "test-grad-clip",
model_name = "unsloth/test",
training_type = "LoRA/QLoRA",
max_grad_norm = 0.7,
)
config = mock_process.call_args.kwargs["kwargs"]["config"]
self.assertEqual(config["max_grad_norm"], 0.7)
def test_training_route_forwards_embedding_learning_rate(self):
training_route = _load_route_module(
"training_route_module_raw_support",

View file

@ -37,6 +37,7 @@ def test_runtime_flash_attn_prefers_prebuilt_wheel(monkeypatch):
statuses: list[str] = []
monkeypatch.delenv(worker._FLASH_ATTN_SKIP_ENV, raising = False)
monkeypatch.setattr(worker, "has_blackwell_gpu", lambda: False)
monkeypatch.setattr(builtins, "__import__", _missing_flash_attn_import())
monkeypatch.setattr(
worker,
@ -65,6 +66,7 @@ def test_runtime_flash_attn_falls_back_to_pypi(monkeypatch):
statuses: list[str] = []
monkeypatch.delenv(worker._FLASH_ATTN_SKIP_ENV, raising = False)
monkeypatch.setattr(worker, "has_blackwell_gpu", lambda: False)
monkeypatch.setattr(builtins, "__import__", _missing_flash_attn_import())
monkeypatch.setattr(
worker,
@ -112,6 +114,29 @@ def test_runtime_flash_attn_skip_env_avoids_all_install_work(monkeypatch):
worker._sp.run.assert_not_called()
def test_runtime_flash_attn_skips_on_blackwell(monkeypatch):
statuses: list[str] = []
install_mock = mock.Mock()
monkeypatch.delenv(worker._FLASH_ATTN_SKIP_ENV, raising = False)
monkeypatch.setattr(
worker, "_should_try_runtime_flash_attn_install", lambda max_seq: True
)
monkeypatch.setattr(worker, "has_blackwell_gpu", lambda: True)
monkeypatch.setattr(worker, "_install_package_wheel_first", install_mock)
monkeypatch.setattr(
worker,
"_send_status",
lambda queue, message: statuses.append(message),
)
worker._ensure_flash_attn_for_long_context(event_queue = [], max_seq_length = 65536)
install_mock.assert_not_called()
assert len(statuses) == 1
assert "Blackwell" in statuses[0]
def test_causal_conv1d_fast_path_preserves_wheel_first_install_args(monkeypatch):
install_mock = mock.Mock(return_value = True)
monkeypatch.setattr(worker, "_install_package_wheel_first", install_mock)

View file

@ -28,6 +28,23 @@ DEFAULT_ALPACA_TEMPLATE = """Below is an instruction that describes a task, pair
{}"""
def _is_mlx_runtime() -> bool:
try:
from unsloth_zoo.mlx import is_mlx_available
except ImportError:
return False
return is_mlx_available()
def _chat_template_kwargs() -> dict:
if not _is_mlx_runtime():
return {}
return {
"patch_saving": False,
"use_zoo_tokenizer_patch": True,
}
def get_tokenizer_chat_template(tokenizer, model_name):
"""
Gets appropriate chat template for tokenizer based on model.
@ -60,6 +77,7 @@ def get_tokenizer_chat_template(tokenizer, model_name):
tokenizer = get_chat_template(
tokenizer,
chat_template = matched_template,
**_chat_template_kwargs(),
)
except Exception as e:
logger.info(f"⚠️ Failed to apply Unsloth template '{matched_template}': {e}")
@ -79,6 +97,7 @@ def get_tokenizer_chat_template(tokenizer, model_name):
tokenizer = get_chat_template(
tokenizer,
chat_template = "chatml",
**_chat_template_kwargs(),
)
except Exception as e:
logger.info(f"⚠️ Failed to apply default ChatML template: {e}")
@ -255,7 +274,11 @@ def apply_chat_template_to_dataset(
if not (hasattr(tokenizer, 'chat_template') and tokenizer.chat_template):
try:
from unsloth.chat_templates import get_chat_template
tokenizer = get_chat_template(tokenizer, chat_template = "alpaca")
tokenizer = get_chat_template(
tokenizer,
chat_template = "alpaca",
**_chat_template_kwargs(),
)
logger.info(f"📝 Set alpaca chat template on tokenizer for model saving")
except Exception as e:
logger.info(f"⚠️ Could not set alpaca template on tokenizer: {e}")

View 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

View file

@ -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

View file

@ -3,6 +3,7 @@
from __future__ import annotations
import functools
import json
import logging
import platform
@ -22,6 +23,49 @@ FLASH_ATTN_RELEASE_BASE_URL = (
)
@functools.lru_cache(maxsize = 1)
def has_blackwell_gpu() -> bool:
"""Return True if any visible NVIDIA GPU has compute capability >= 10.0
(Blackwell: sm_100, sm_120, sm_121, ...).
Dao-AILab does not publish prebuilt flash-attention wheels for these
architectures, and the older-arch wheels fail to load on Blackwell, so
callers use this gate to skip the flash-attn install/upgrade path.
Result is cached for the process lifetime since GPU hardware does not
change. Tests that mock subprocess/nvidia-smi must call
``has_blackwell_gpu.cache_clear()`` before each invocation.
"""
exe = shutil.which("nvidia-smi")
if not exe:
return False
try:
result = subprocess.run(
[exe, "--query-gpu=compute_cap", "--format=csv,noheader"],
stdout = subprocess.PIPE,
stderr = subprocess.DEVNULL,
text = True,
timeout = 10,
env = child_env_without_native_path_secret(),
)
except (OSError, subprocess.TimeoutExpired):
return False
if result.returncode != 0:
return False
for line in result.stdout.splitlines():
cap = line.strip()
if not cap:
continue
major_part = cap.split(".", 1)[0]
try:
major = int(major_part)
except ValueError:
continue
if major >= 10:
return True
return False
def linux_wheel_platform_tag() -> str | None:
machine = platform.machine().lower()
if sys.platform.startswith("linux"):

View file

@ -83,6 +83,7 @@ export function buildTrainingStartPayload(
save_steps: config.saveSteps,
eval_steps: config.evalSteps,
weight_decay: config.weightDecay,
max_grad_norm: 0.0,
random_seed: config.randomSeed,
packing: isEmbedding ? false : config.packing,
optim: config.optimizerType,

View file

@ -31,6 +31,7 @@ export interface TrainingStartRequest {
save_steps: number;
eval_steps: number;
weight_decay: number;
max_grad_norm: number;
random_seed: number;
packing: boolean;
optim: string;

View file

@ -28,6 +28,7 @@ if str(_BACKEND_DIR) not in sys.path:
from backend.utils.wheel_utils import (
flash_attn_package_version,
flash_attn_wheel_url,
has_blackwell_gpu,
install_wheel,
probe_torch_wheel_env,
url_exists,
@ -628,10 +629,19 @@ def _flash_attn_install_disabled() -> bool:
def _ensure_flash_attn() -> None:
if NO_TORCH or IS_WINDOWS or IS_MACOS:
return
if _flash_attn_install_disabled():
return
if NO_TORCH:
return
if has_blackwell_gpu():
_step(
"warning",
"Skipping flash-attn: Blackwell GPU detected (sm_100+); no compatible prebuilt wheel",
_cyan,
)
return
if IS_WINDOWS or IS_MACOS:
return
if (
subprocess.run(
[sys.executable, "-c", "import flash_attn"],

View file

@ -10,8 +10,133 @@ from unittest import mock
STUDIO_DIR = Path(__file__).resolve().parents[2] / "studio"
sys.path.insert(0, str(STUDIO_DIR))
sys.path.insert(0, str(STUDIO_DIR / "backend"))
import install_python_stack as ips
from backend.utils import wheel_utils
def _smi_result(stdout: str, returncode: int = 0) -> subprocess.CompletedProcess:
return subprocess.CompletedProcess(["nvidia-smi"], returncode, stdout, "")
class TestHasBlackwellGpu:
def setup_method(self):
wheel_utils.has_blackwell_gpu.cache_clear()
def teardown_method(self):
wheel_utils.has_blackwell_gpu.cache_clear()
def test_returns_false_when_nvidia_smi_missing(self):
with mock.patch.object(wheel_utils.shutil, "which", return_value = None):
assert wheel_utils.has_blackwell_gpu() is False
def test_returns_true_for_sm_100(self):
with (
mock.patch.object(
wheel_utils.shutil, "which", return_value = "/usr/bin/nvidia-smi"
),
mock.patch.object(
wheel_utils.subprocess, "run", return_value = _smi_result("10.0\n")
),
):
assert wheel_utils.has_blackwell_gpu() is True
def test_returns_true_for_sm_120(self):
with (
mock.patch.object(
wheel_utils.shutil, "which", return_value = "/usr/bin/nvidia-smi"
),
mock.patch.object(
wheel_utils.subprocess, "run", return_value = _smi_result("12.0\n")
),
):
assert wheel_utils.has_blackwell_gpu() is True
def test_returns_true_for_sm_121(self):
with (
mock.patch.object(
wheel_utils.shutil, "which", return_value = "/usr/bin/nvidia-smi"
),
mock.patch.object(
wheel_utils.subprocess, "run", return_value = _smi_result("12.1\n")
),
):
assert wheel_utils.has_blackwell_gpu() is True
def test_returns_false_for_sm_90(self):
with (
mock.patch.object(
wheel_utils.shutil, "which", return_value = "/usr/bin/nvidia-smi"
),
mock.patch.object(
wheel_utils.subprocess, "run", return_value = _smi_result("9.0\n")
),
):
assert wheel_utils.has_blackwell_gpu() is False
def test_returns_false_for_sm_89(self):
with (
mock.patch.object(
wheel_utils.shutil, "which", return_value = "/usr/bin/nvidia-smi"
),
mock.patch.object(
wheel_utils.subprocess, "run", return_value = _smi_result("8.9\n")
),
):
assert wheel_utils.has_blackwell_gpu() is False
def test_mixed_gpus_with_one_blackwell_returns_true(self):
with (
mock.patch.object(
wheel_utils.shutil, "which", return_value = "/usr/bin/nvidia-smi"
),
mock.patch.object(
wheel_utils.subprocess,
"run",
return_value = _smi_result("8.0\n10.0\n"),
),
):
assert wheel_utils.has_blackwell_gpu() is True
def test_returns_false_when_nvidia_smi_fails(self):
with (
mock.patch.object(
wheel_utils.shutil, "which", return_value = "/usr/bin/nvidia-smi"
),
mock.patch.object(
wheel_utils.subprocess,
"run",
return_value = _smi_result("", returncode = 1),
),
):
assert wheel_utils.has_blackwell_gpu() is False
def test_returns_false_on_subprocess_timeout(self):
with (
mock.patch.object(
wheel_utils.shutil, "which", return_value = "/usr/bin/nvidia-smi"
),
mock.patch.object(
wheel_utils.subprocess,
"run",
side_effect = subprocess.TimeoutExpired(cmd = "nvidia-smi", timeout = 10),
),
):
assert wheel_utils.has_blackwell_gpu() is False
def test_returns_false_on_malformed_output(self):
with (
mock.patch.object(
wheel_utils.shutil, "which", return_value = "/usr/bin/nvidia-smi"
),
mock.patch.object(
wheel_utils.subprocess,
"run",
return_value = _smi_result("not-a-number\n\n"),
),
):
assert wheel_utils.has_blackwell_gpu() is False
class TestFlashAttnWheelSelection:
@ -234,6 +359,76 @@ class TestEnsureFlashAttn:
mock_probe.assert_not_called()
mock_install_wheel.assert_not_called()
def test_blackwell_gpu_skips_install_with_warning(self):
step_messages: list[tuple[str, str]] = []
def fake_step(label: str, value: str, color_fn = None):
step_messages.append((label, value))
with (
mock.patch.object(ips, "NO_TORCH", False),
mock.patch.object(ips, "IS_WINDOWS", False),
mock.patch.object(ips, "IS_MACOS", False),
mock.patch.object(ips, "has_blackwell_gpu", return_value = True),
mock.patch.object(ips, "probe_torch_wheel_env") as mock_probe,
mock.patch.object(ips, "install_wheel") as mock_install_wheel,
mock.patch.object(ips, "_step", side_effect = fake_step),
mock.patch("subprocess.run", return_value = self._import_check()),
):
ips._ensure_flash_attn()
mock_probe.assert_not_called()
mock_install_wheel.assert_not_called()
assert any(
label == "warning" and "Blackwell" in msg for label, msg in step_messages
)
def test_blackwell_gpu_on_windows_emits_blackwell_warning(self):
step_messages: list[tuple[str, str]] = []
def fake_step(label: str, value: str, color_fn = None):
step_messages.append((label, value))
with (
mock.patch.object(ips, "NO_TORCH", False),
mock.patch.object(ips, "IS_WINDOWS", True),
mock.patch.object(ips, "IS_MACOS", False),
mock.patch.object(ips, "has_blackwell_gpu", return_value = True),
mock.patch.object(ips, "probe_torch_wheel_env") as mock_probe,
mock.patch.object(ips, "install_wheel") as mock_install_wheel,
mock.patch.object(ips, "_step", side_effect = fake_step),
mock.patch("subprocess.run", return_value = self._import_check()),
):
ips._ensure_flash_attn()
mock_probe.assert_not_called()
mock_install_wheel.assert_not_called()
assert any(
label == "warning" and "Blackwell" in msg for label, msg in step_messages
)
def test_non_blackwell_windows_does_not_emit_blackwell_warning(self):
step_messages: list[tuple[str, str]] = []
def fake_step(label: str, value: str, color_fn = None):
step_messages.append((label, value))
with (
mock.patch.object(ips, "NO_TORCH", False),
mock.patch.object(ips, "IS_WINDOWS", True),
mock.patch.object(ips, "IS_MACOS", False),
mock.patch.object(ips, "has_blackwell_gpu", return_value = False),
mock.patch.object(ips, "probe_torch_wheel_env") as mock_probe,
mock.patch.object(ips, "install_wheel") as mock_install_wheel,
mock.patch.object(ips, "_step", side_effect = fake_step),
mock.patch("subprocess.run", return_value = self._import_check()),
):
ips._ensure_flash_attn()
mock_probe.assert_not_called()
mock_install_wheel.assert_not_called()
assert not any("Blackwell" in msg for _, msg in step_messages)
class TestInstallPythonStackFlashAttnIntegration:
def _run_install(self, *, no_torch: bool, is_macos: bool, is_windows: bool) -> int:

View file

@ -211,8 +211,8 @@ def cmd_train(args) -> int:
workdir.mkdir(parents = True, exist_ok = True)
import mlx.core as mx
from unsloth_zoo.mlx_loader import FastMLXModel
from unsloth_zoo.mlx_trainer import MLXTrainer, MLXTrainingConfig
from unsloth_zoo.mlx.loader import FastMLXModel
from unsloth_zoo.mlx.trainer import MLXTrainer, MLXTrainingConfig
hf_token = os.environ.get("HF_TOKEN") or None
@ -440,7 +440,7 @@ def cmd_reload(args) -> int:
return _reload_gguf(save_dir, metrics)
import mlx.core as mx
from unsloth_zoo.mlx_loader import FastMLXModel
from unsloth_zoo.mlx.loader import FastMLXModel
from mlx_lm import generate
hf_token = os.environ.get("HF_TOKEN") or None

View file

@ -7,8 +7,9 @@ Two gates drive every dispatch decision in Studio's MLX path:
1. ``unsloth._IS_MLX`` at the top of ``unsloth/__init__.py`` -- evaluated
once at import time and read by Studio worker code to choose between
the GPU and MLX trainer / inference / export paths. Defined as
``Darwin AND arm64 AND find_spec("mlx") is not None``.
the GPU and MLX trainer / inference / export paths. It delegates to
the shared zoo MLX runtime gate, with a local import barrier while the
paired unsloth-zoo runtime rollout is in flight.
2. ``utils.hardware.detect_hardware()`` -- runtime probe in the Studio
backend. Priority order: CUDA -> XPU -> MLX -> CPU. The MLX branch is
@ -18,8 +19,8 @@ Two gates drive every dispatch decision in Studio's MLX path:
These gates are the canaries for "MLX support accidentally hijacks
CUDA/AMD/Intel users". The tests here:
* verify the source-level structure of the ``_IS_MLX`` expression so an
accidental rewrite (e.g. dropping the ``arm64`` check) is caught,
* verify the source-level structure of the ``_IS_MLX`` helper so an
accidental rewrite importing zoo before the local MLX precheck is caught,
* exercise the runtime gate logic under a spoofed Darwin+arm64 platform
with a fake ``mlx`` module in ``sys.modules`` to confirm both gates
flip True together,
@ -64,20 +65,36 @@ def test_is_mlx_gate_uses_three_required_predicates():
target = node.value
break
assert target is not None, "_IS_MLX assignment not found in unsloth/__init__.py"
assert isinstance(target, ast.BoolOp) and isinstance(
target.op, ast.And
), "_IS_MLX must be a BoolOp(And) of platform + mlx checks"
assert isinstance(target, ast.Call), "_IS_MLX must call the shared MLX helper"
expr_src = ast.unparse(target)
assert (
"platform.system()" in expr_src and "Darwin" in expr_src
), "_IS_MLX must check platform.system() == 'Darwin'"
expr_src == "_is_mlx_available()"
), "_IS_MLX must delegate to the shared MLX runtime gate"
helper = None
for node in ast.walk(tree):
if isinstance(node, ast.FunctionDef) and node.name == "_is_mlx_available":
helper = node
break
assert helper is not None, "_is_mlx_available helper not found"
helper_src = ast.unparse(helper)
assert (
"platform.machine()" in expr_src and "arm64" in expr_src
), "_IS_MLX must check platform.machine() == 'arm64'"
"platform.system()" in helper_src
and "'Darwin'" in helper_src
and "platform.machine()" in helper_src
and "'arm64'" in helper_src
and "find_spec" in helper_src
and "'mlx'" in helper_src
and "from unsloth_zoo.mlx import is_mlx_available" in helper_src
), "_IS_MLX helper must precheck local MLX predicates before importing zoo"
assert (
"find_spec" in expr_src and "'mlx'" in expr_src
), "_IS_MLX must check importlib.util.find_spec('mlx')"
"from unsloth_zoo.mlx import is_mlx_available" in helper_src
and "return is_mlx_available()" in helper_src
), "_IS_MLX helper must delegate final detection to the shared zoo MLX runtime gate"
assert helper_src.index("UNSLOTH_FORCE_GPU_PATH") < helper_src.index(
"from unsloth_zoo.mlx import is_mlx_available"
), "_IS_MLX helper must run the local MLX precheck before importing zoo"
# ---------------------------------------------------------------------------
@ -87,13 +104,14 @@ def test_is_mlx_gate_uses_three_required_predicates():
# ---------------------------------------------------------------------------
def _evaluate_is_mlx_gate(platform_module, importlib_util):
"""Re-evaluate the _IS_MLX expression using injected dependencies.
def _evaluate_is_mlx_precheck(platform_module, importlib_util, os_module):
"""Re-evaluate the local _is_mlx_available precheck using injected dependencies.
Mirrors the assignment in unsloth/__init__.py exactly.
Mirrors only the cheap import barrier before unsloth imports unsloth_zoo.
"""
return (
platform_module.system() == "Darwin"
os_module.environ.get("UNSLOTH_FORCE_GPU_PATH", "0") != "1"
and platform_module.system() == "Darwin"
and platform_module.machine() == "arm64"
and importlib_util.find_spec("mlx") is not None
)
@ -112,7 +130,9 @@ def test_is_mlx_gate_true_on_apple_silicon_with_mlx_present(monkeypatch):
monkeypatch.setattr(platform, "system", lambda: "Darwin")
monkeypatch.setattr(platform, "machine", lambda: "arm64")
assert _evaluate_is_mlx_gate(platform, importlib.util) is True
import os
assert _evaluate_is_mlx_precheck(platform, importlib.util, os) is True
def test_is_mlx_gate_false_when_mlx_missing(monkeypatch):
@ -133,7 +153,9 @@ def test_is_mlx_gate_false_when_mlx_missing(monkeypatch):
monkeypatch.setattr(importlib.util, "find_spec", _no_mlx)
assert _evaluate_is_mlx_gate(platform, importlib.util) is False
import os
assert _evaluate_is_mlx_precheck(platform, importlib.util, os) is False
def test_is_mlx_gate_false_on_non_apple_silicon():
@ -147,7 +169,9 @@ def test_is_mlx_gate_false_on_non_apple_silicon():
pytest.skip("Test host is Apple Silicon; CUDA-side canary doesn't apply.")
assert _evaluate_is_mlx_gate(platform, importlib.util) is False
import os
assert _evaluate_is_mlx_precheck(platform, importlib.util, os) is False
# ---------------------------------------------------------------------------

View file

@ -46,8 +46,8 @@ def test_wandb_init_strips_secret_keys():
def test_local_dataset_loader_uses_load_dataset_path():
src = WORKER.read_text()
assert "_resolve_local_files" in src
assert "_loader_for_files" in src
assert "_resolve_mlx_local_dataset_files" in src
assert "_mlx_local_dataset_loader_for_files" in src
assert "data_files = all_files" in src or "data_files=all_files" in src
@ -84,7 +84,7 @@ def test_poll_stop_returns_on_broken_pipe():
def test_unsloth_zoo_mlx_imports_have_friendly_error():
src = WORKER.read_text()
assert "from unsloth_zoo.mlx_loader import FastMLXModel" in src
assert "from unsloth_zoo.mlx_trainer import" in src
assert "from unsloth_zoo.mlx.loader import FastMLXModel" in src
assert "from unsloth_zoo.mlx.trainer import" in src
assert "raise ImportError" in src
assert "install.sh" in src

View 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."
)

View file

@ -12,16 +12,33 @@
# See the License for the specific language governing permissions and
# limitations under the License.
import os, platform, importlib.util
import os, importlib.util, platform
os.environ["UNSLOTH_IS_PRESENT"] = "1"
def _is_mlx_available():
# Transitional import barrier: while the paired unsloth-zoo MLX runtime
# rollout is in flight, keep non-Apple-Silicon imports from touching
# unsloth_zoo here. After both PRs are released together and
# unsloth_zoo.mlx is guaranteed to be import-safe on GPU hosts,
# this helper can collapse back to the centralized zoo runtime call below.
if (
os.environ.get("UNSLOTH_FORCE_GPU_PATH", "0") == "1"
or platform.system() != "Darwin"
or platform.machine() != "arm64"
or importlib.util.find_spec("mlx") is None
):
return False
try:
from unsloth_zoo.mlx import is_mlx_available
except ImportError:
return False
return is_mlx_available()
# Detect Apple Silicon + MLX before any torch/numpy imports
_IS_MLX = (
platform.system() == "Darwin"
and platform.machine() == "arm64"
and importlib.util.find_spec("mlx") is not None
)
_IS_MLX = _is_mlx_available()
if _IS_MLX:
try:
@ -31,18 +48,18 @@ if _IS_MLX:
"Unsloth: MLX support requires `unsloth-zoo` with MLX modules. "
"Reinstall with `pip install unsloth-zoo` or rerun install.sh."
) from _e
# The mlx_trainer / mlx_loader submodules ship with unsloth-zoo's MLX
# The mlx.trainer / mlx.loader submodules ship with unsloth-zoo's MLX
# support. An older installed unsloth-zoo (e.g. from PyPI before the
# MLX release lands) will satisfy `import unsloth_zoo` but be missing
# these submodules. Surface the same friendly install hint instead of
# a raw ImportError on the submodule path.
try:
from unsloth_zoo.mlx_trainer import MLXTrainer, MLXTrainingConfig
from unsloth_zoo.mlx_loader import FastMLXModel
from unsloth_zoo.mlx.trainer import MLXTrainer, MLXTrainingConfig
from unsloth_zoo.mlx.loader import FastMLXModel
except ImportError as _e:
raise ImportError(
"Unsloth: MLX support requires an unsloth-zoo build that includes "
"`unsloth_zoo.mlx_trainer` and `unsloth_zoo.mlx_loader`. Upgrade with "
"`unsloth_zoo.mlx.trainer` and `unsloth_zoo.mlx.loader`. Upgrade with "
"`pip install -U unsloth-zoo` or rerun install.sh."
) from _e

View file

@ -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()

View file

@ -30,11 +30,9 @@ __all__ = [
from transformers import StoppingCriteria, StoppingCriteriaList
from torch import LongTensor, FloatTensor
from transformers.models.llama.modeling_llama import logger
from .save import patch_saving_functions
import os
import shutil
from .tokenizer_utils import *
from .models._utils import patch_tokenizer
import re
from .ollama_template_mappers import OLLAMA_TEMPLATES
from unsloth_zoo.dataset_utils import (
@ -213,7 +211,7 @@ vicuna_ollama = _ollama_template("vicuna")
vicuna_eos_token = "eos_token"
CHAT_TEMPLATES["vicuna"] = (vicuna_template, vicuna_eos_token, False, vicuna_ollama,)
DEFAULT_SYSTEM_MESSAGE["vicuna"] = "A chat between a curious user and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the user's questions."
DEFAULT_SYSTEM_MESSAGE["vicuna"] = "A chat between a curious user and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the user\\'s questions."
# =========================================== Vicuna Old
# https://github.com/lm-sys/FastChat/blob/main/docs/vicuna_weights_version.md#prompt-template
@ -1844,6 +1842,8 @@ def get_chat_template(
mapping = {"role" : "role", "content" : "content", "user" : "user", "assistant" : "assistant"},
map_eos_token = True,
system_message = None,
patch_saving = True,
use_zoo_tokenizer_patch = False,
):
assert(type(map_eos_token) is bool)
old_tokenizer = tokenizer
@ -2026,6 +2026,12 @@ def get_chat_template(
.replace("'user'", "'" + mapping["user"] + "'")\
.replace("'assistant'", "'" + mapping["assistant"] + "'")
if use_zoo_tokenizer_patch:
# Studio MLX avoids the model-utils tokenizer wrapper because that
# import path pulls in Torch/GPU-specific modules before MLX training.
from unsloth_zoo.tokenizer_utils import patch_tokenizer
else:
from .models._utils import patch_tokenizer
_, tokenizer = patch_tokenizer(model = None, tokenizer = tokenizer)
tokenizer.padding_side = old_padding_side
@ -2059,7 +2065,9 @@ def get_chat_template(
# stopping_criteria = create_stopping_criteria(tokenizer, stop_word)
# Patch saving functions
tokenizer = patch_saving_functions(tokenizer)
if patch_saving:
from .save import patch_saving_functions
tokenizer = patch_saving_functions(tokenizer)
# Add Ollama
tokenizer._ollama_modelfile = ollama_modelfile

View file

@ -20,21 +20,40 @@ __all__ = [
"DEVICE_COUNT",
"ALLOW_PREQUANTIZED_MODELS",
"ALLOW_BITSANDBYTES",
"is_mlx_available",
]
import torch
import functools
import inspect
import os
from unsloth_zoo.utils import Version
def is_mlx_available():
try:
from unsloth_zoo.mlx import is_mlx_available as _is_mlx_available
except ImportError:
return False
return _is_mlx_available()
_IS_MLX = is_mlx_available()
if not _IS_MLX:
import torch
@functools.cache
def is_hip():
if _IS_MLX:
return False
return bool(getattr(getattr(torch, "version", None), "hip", None))
@functools.cache
def get_device_type():
if _IS_MLX:
return "mlx"
if hasattr(torch, "cuda") and torch.cuda.is_available():
if is_hip():
return "hip"
@ -44,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."
)
@ -54,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."
)
@ -64,6 +89,8 @@ DEVICE_TYPE: str = get_device_type()
DEVICE_TYPE_TORCH = DEVICE_TYPE
if DEVICE_TYPE_TORCH == "hip":
DEVICE_TYPE_TORCH = "cuda"
elif DEVICE_TYPE_TORCH == "mlx":
DEVICE_TYPE_TORCH = "mps"
@functools.cache

View file

@ -71,7 +71,7 @@ def load_cached_config(cache_key: str) -> Optional[Dict[str, Any]]:
return None
try:
with open(cache_file, "r") as f:
with open(cache_file, "r", encoding = "utf-8") as f:
cached_data = json.load(f)
# Verify cache is still valid (same device, etc.)
@ -118,7 +118,7 @@ def save_cached_config(
}
try:
with open(cache_file, "w") as f:
with open(cache_file, "w", encoding = "utf-8") as f:
json.dump(cache_data, f, indent = 2)
logger.info(f"Saved MoE kernel config cache: {cache_key}")
except Exception as e:

View file

@ -183,7 +183,7 @@ def save_autotune_results(autotune_cache, mode, ref_time, fused_time, results_di
filename = "_".join(key)
save_path = f"{save_dir}/{filename}.json"
print(f"Saving autotune results to {save_path}")
with open(save_path, "w") as f:
with open(save_path, "w", encoding = "utf-8") as f:
result = {
**config.all_kwargs(),
"ref_time": ref_time,

View file

@ -160,6 +160,10 @@ else:
# INTEL GPU Specific Logic
if DEVICE_TYPE == "xpu":
_gpu_getCurrentRawStream = torch._C._xpu_getCurrentRawStream
elif DEVICE_TYPE == "mlx":
def _gpu_getCurrentRawStream(_index = 0):
return 0
# NVIDIA GPU Default Logic
elif hasattr(torch._C, "_cuda_getCurrentRawStream"):
_gpu_getCurrentRawStream = torch._C._cuda_getCurrentRawStream
@ -206,6 +210,11 @@ if DEVICE_TYPE == "xpu":
XPU_STREAMS = ()
WEIGHT_BUFFERS = []
ABSMAX_BUFFERS = []
elif DEVICE_TYPE == "mlx":
CUDA_STREAMS = ()
XPU_STREAMS = ()
WEIGHT_BUFFERS = []
ABSMAX_BUFFERS = []
else:
# NVIDIA GPU Default Logic
if DEVICE_COUNT > 0:

View file

@ -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)

View file

@ -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

View file

@ -70,7 +70,7 @@ def _save_pretrained_torchao(
modules_path = os.path.join(save_directory, "modules.json")
if os.path.exists(modules_path):
try:
with open(modules_path, "r") as f:
with open(modules_path, "r", encoding = "utf-8") as f:
modules = json.load(f)
for m in modules:
if m.get("type", "").endswith("Transformer"):
@ -177,7 +177,7 @@ def _save_pretrained_gguf(
modules_path = os.path.join(save_directory, "modules.json")
if os.path.exists(modules_path):
try:
with open(modules_path, "r") as f:
with open(modules_path, "r", encoding = "utf-8") as f:
modules = json.load(f)
for m in modules:
if m.get("type", "").endswith("Transformer"):
@ -542,7 +542,7 @@ class FastSentenceTransformer(FastModel):
model_name, "modules.json", token = token
)
with open(modules_json_path, "r") as f:
with open(modules_json_path, "r", encoding = "utf-8") as f:
modules_config = json.load(f)
pooling_config_path = None
@ -566,7 +566,7 @@ class FastSentenceTransformer(FastModel):
break
if pooling_config_path:
with open(pooling_config_path, "r") as f:
with open(pooling_config_path, "r", encoding = "utf-8") as f:
pooling_config = json.load(f)
# from here:
# https://github.com/huggingface/sentence-transformers/blob/main/sentence_transformers/models/Pooling.py#L43

View file

@ -2641,7 +2641,7 @@ This model was finetuned and converted to GGUF format using [Unsloth](https://gi
)
readme_path = os.path.join(actual_save_directory, "README.md")
with open(readme_path, "w") as f:
with open(readme_path, "w", encoding = "utf-8") as f:
f.write(readme_content)
api.upload_file(