Merge branch 'main' into fix-issue-5344-quantization-guardrail
This commit is contained in:
commit
8bad7ac1b7
38 changed files with 1790 additions and 117 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
|
||||
|
|
|
|||
|
|
@ -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}]")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
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
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
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."
|
||||
)
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue