diff --git a/.github/workflows/consolidated-tests-ci.yml b/.github/workflows/consolidated-tests-ci.yml index 7978a200c0..ae4b386589 100644 --- a/.github/workflows/consolidated-tests-ci.yml +++ b/.github/workflows/consolidated-tests-ci.yml @@ -364,6 +364,7 @@ jobs: tests/utils/test_attention_masks.py \ tests/utils/test_trunc_normal_patch.py \ tests/python/test_fast_language_model_text_only.py \ + tests/test_prefetch_snapshot_scope.py \ --deselect 'tests/utils/test_attention_masks.py::test_run_attention_flash_varlen_receives_window_and_softcap' # The deselected test monkeypatches flash_attn_varlen_func, which is # only bound on the module when `flash_attn` is importable. flash_attn diff --git a/.github/workflows/lockfile-audit.yml b/.github/workflows/lockfile-audit.yml index 9c28e21672..aaf258d615 100644 --- a/.github/workflows/lockfile-audit.yml +++ b/.github/workflows/lockfile-audit.yml @@ -60,11 +60,11 @@ jobs: runs-on: ubuntu-latest timeout-minutes: 5 steps: - - uses: actions/checkout@v4 + - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 with: persist-credentials: false - - uses: actions/setup-python@v5 + - uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0 with: python-version: '3.12' diff --git a/.github/workflows/studio-windows-inference-smoke.yml b/.github/workflows/studio-windows-inference-smoke.yml index 8186c07211..0bc216d65a 100644 --- a/.github/workflows/studio-windows-inference-smoke.yml +++ b/.github/workflows/studio-windows-inference-smoke.yml @@ -1610,6 +1610,13 @@ jobs: - name: Install Pester v5 shell: pwsh run: | + # PSGallery is intermittently absent from the repository list on GitHub's Windows + # runners, which makes `Set-PSRepository PSGallery` fail with "No repository with the + # name 'PSGallery' was found." Re-register the default gallery first so the policy + # change and module install below always have a repository to target. + if (-not (Get-PSRepository -Name PSGallery -ErrorAction SilentlyContinue)) { + Register-PSRepository -Default -ErrorAction SilentlyContinue + } Set-PSRepository PSGallery -InstallationPolicy Trusted Install-Module Pester -MinimumVersion 5.5.0 -Force -SkipPublisherCheck -Scope CurrentUser Import-Module Pester -MinimumVersion 5.5.0 diff --git a/studio/backend/tests/test_hf_xet_fallback.py b/studio/backend/tests/test_hf_xet_fallback.py index 9e40fbf508..4d73213d15 100644 --- a/studio/backend/tests/test_hf_xet_fallback.py +++ b/studio/backend/tests/test_hf_xet_fallback.py @@ -1,18 +1,16 @@ # SPDX-License-Identifier: AGPL-3.0-only # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 -"""Unit tests for utils.hf_xet_fallback: the no-progress watchdog, the Xet->HTTP -transport policy, and the HF_HUB_DISABLE_XET precondition the fallback rests on. -CPU-only, no network, no real subprocess (the per-attempt download seam is -monkeypatched). +"""Tests for the Studio shim over the shared unsloth_zoo Xet -> HTTP fallback. + +The transport-policy matrix is tested once in unsloth_zoo; here we assert only the +Studio seam: re-exporting the shared API and injecting the marker-aware +prepare_cache_for_transport on the HTTP retry. CPU-only, no network, no real subprocess. """ from __future__ import annotations -import subprocess import sys -import threading -import time import types as _types from pathlib import Path @@ -22,9 +20,8 @@ _BACKEND_DIR = str(Path(__file__).resolve().parent.parent) if _BACKEND_DIR not in sys.path: sys.path.insert(0, _BACKEND_DIR) -# Stub heavy/unavailable deps before importing the module under test. Use the -# real structlog when present; a bare stub left in sys.modules would break later -# modules that log at import time. +# Stub heavy/unavailable deps before importing the module under test. Use real structlog when present; +# a bare stub would break later modules that log at import time. _loggers_stub = _types.ModuleType("loggers") _loggers_stub.get_logger = lambda name: __import__("logging").getLogger(name) sys.modules.setdefault("loggers", _loggers_stub) @@ -34,171 +31,59 @@ except ImportError: sys.modules["structlog"] = _types.ModuleType("structlog") import huggingface_hub -from huggingface_hub import constants as hf_constants + +try: + import unsloth_zoo.hf_xet_fallback as _shared_mod + shared = _shared_mod +except Exception: # noqa: BLE001 - still collect degraded-path tests when unsloth_zoo is unavailable + shared = None import utils.hf_xet_fallback as xf -# --------------------------------------------------------------------------- # -# Watchdog: fires only on a constant-size .incomplete, sparse-aware byte total. -# --------------------------------------------------------------------------- # -REPO = "ztest/xet-watchdog" - - -@pytest.fixture -def hf_cache(tmp_path, monkeypatch): - monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path)) - return tmp_path - - -def _blobs_dir(root: Path, repo_id: str = REPO) -> Path: - d = root / f"models--{repo_id.replace('/', '--')}" / "blobs" - d.mkdir(parents = True, exist_ok = True) - return d - - -def _wait( - predicate, - timeout: float = 2.0, - step: float = 0.02, -) -> bool: - deadline = time.monotonic() + timeout - while time.monotonic() < deadline: - if predicate(): - return True - time.sleep(step) - return predicate() - - -def test_constant_incomplete_fires_stall(hf_cache): - blobs = _blobs_dir(hf_cache) - (blobs / "deadbeef.incomplete").write_bytes(b"\0" * 1024) # never grows - - calls: list[str] = [] - stop = xf.start_watchdog( - repo_ids = [REPO], on_stall = calls.append, interval = 0.05, stall_timeout = 0.3 - ) - try: - assert _wait( - lambda: len(calls) >= 1, timeout = 3.0 - ), "watchdog never fired on a constant-size .incomplete" - finally: - stop.set() - assert "stalled" in calls[0].lower() - - -def test_growing_incomplete_never_stalls(hf_cache): - blobs = _blobs_dir(hf_cache) - part = blobs / "growing.incomplete" - part.write_bytes(b"\0" * 1024) - - grow_stop = threading.Event() - - def _grow(): - size = 1024 - while not grow_stop.wait(0.05): - size += 4096 - part.write_bytes(b"\0" * size) - - grower = threading.Thread(target = _grow, daemon = True) - grower.start() - - calls: list[str] = [] - stop = xf.start_watchdog( - repo_ids = [REPO], on_stall = calls.append, interval = 0.05, stall_timeout = 0.3 - ) - try: - time.sleep(1.0) # well past stall_timeout, but bytes keep growing - assert calls == [], "watchdog fired despite continuous progress" - finally: - stop.set() - grow_stop.set() - - -def test_no_incomplete_never_stalls(hf_cache): - blobs = _blobs_dir(hf_cache) - (blobs / "finalized_blob").write_bytes(b"\0" * 4096) # no .incomplete - - calls: list[str] = [] - stop = xf.start_watchdog( - repo_ids = [REPO], on_stall = calls.append, interval = 0.05, stall_timeout = 0.3 - ) - try: - time.sleep(0.8) - assert calls == [], "watchdog fired with no active .incomplete" - finally: - stop.set() - - -def test_stall_fires_at_most_once(hf_cache): - blobs = _blobs_dir(hf_cache) - (blobs / "frozen.incomplete").write_bytes(b"\0" * 2048) - - calls: list[str] = [] - stop = xf.start_watchdog( - repo_ids = [REPO], on_stall = calls.append, interval = 0.05, stall_timeout = 0.2 - ) - try: - assert _wait(lambda: len(calls) >= 1, timeout = 3.0) - time.sleep(0.6) # keep ticking; must not fire again - assert len(calls) == 1, f"on_stall fired {len(calls)} times, expected exactly 1" - finally: - stop.set() - - -def test_get_state_empty_cache(hf_cache): - assert xf.get_hf_download_state([REPO]) == (0, False) - - -def test_get_state_absent_cache_root(tmp_path, monkeypatch): - monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path / "no-such-cache")) - assert xf.get_hf_download_state([REPO]) == (0, False) - - -def test_get_state_skips_local_paths(hf_cache): - # Filesystem paths are not HF repo IDs and must be ignored without error. - assert xf.get_hf_download_state(["/abs/path", "./rel", "~user", "c:\\x"]) == (0, False) - - -def test_get_state_sparse_aware(hf_cache): - blobs = _blobs_dir(hf_cache) - sparse = blobs / "sparse.incomplete" - with open(sparse, "wb") as f: - f.truncate(64 * 1024 * 1024) # large apparent size, few allocated blocks - st = sparse.stat() - if getattr(st, "st_blocks", 0) == 0: - pytest.skip("filesystem does not report st_blocks; sparse accounting unavailable") - total, has_incomplete = xf.get_hf_download_state([REPO]) - assert has_incomplete is True - assert total < st.st_size, "sparse partial counted at apparent size, not allocated blocks" - - -# --------------------------------------------------------------------------- # -# Transport policy: cached short-circuit, cancel, error propagation, and the -# single Xet->HTTP fallback. _run_download_attempt is faked, so no real spawn. -# --------------------------------------------------------------------------- # DL_REPO, FILE = "ztest/xet-dl", "model-Q4_K_XL.gguf" -@pytest.fixture(autouse = True) -def _no_real_cache_hit(monkeypatch): - """Default: the cached probe misses; tests override it to force a hit.""" +def _requires_shared(): + if shared is None: + pytest.skip("unsloth_zoo.hf_xet_fallback is not installed in this environment") + + +def test_shim_reexports_shared_api(): + _requires_shared() + assert xf.DownloadStallError is shared.DownloadStallError + for name in ( + "start_watchdog", + "get_hf_download_state", + "child_should_disable_xet", + "hf_hub_download_with_xet_fallback", + "snapshot_download_with_xet_fallback", + ): + assert hasattr(xf, name), f"shim missing {name}" + + +def test_child_should_disable_xet_truth_table(): + assert xf.child_should_disable_xet({"disable_xet": True}) is True + assert xf.child_should_disable_xet({"disable_xet": False}) is False + assert xf.child_should_disable_xet({}) is False + + +def test_shim_injects_studio_prepare_on_http_retry(monkeypatch): + """A Xet stall retries over HTTP and the shim runs Studio's marker-aware + ``prepare_cache_for_transport(..., 'http')`` before the retry.""" + _requires_shared() + for var in ("UNSLOTH_DISABLE_XET", "UNSLOTH_STABLE_DOWNLOADS", "HF_HUB_DISABLE_XET"): + monkeypatch.delenv(var, raising = False) monkeypatch.setattr(huggingface_hub, "try_to_load_from_cache", lambda *a, **k: None) + seen_disable_xet = [] -class _FakeAttempt: - """Records calls to the download seam and returns scripted results.""" - - def __init__(self, results): - self._results = list(results) - self.calls = [] - - def __call__( - self, + def fake_attempt( repo_id, - filename, - token, *, + kind, + params, + token, repo_type, disable_xet, cancel_event, @@ -208,146 +93,243 @@ class _FakeAttempt: on_status, force_download = False, ): - self.calls.append( - _types.SimpleNamespace( - repo_id = repo_id, - filename = filename, - disable_xet = disable_xet, - repo_type = repo_type, - ) - ) - return self._results[len(self.calls) - 1] + seen_disable_xet.append(disable_xet) + return ("ok", "/cache/model.gguf") if disable_xet else ("stall", None) + monkeypatch.setattr(shared, "_run_download_attempt", fake_attempt) -def _install(monkeypatch, results): - fake = _FakeAttempt(results) - monkeypatch.setattr(xf, "_run_download_attempt", fake) - return fake - - -def test_cached_file_short_circuits(monkeypatch, tmp_path): - cached = tmp_path / "cached.gguf" - cached.write_bytes(b"\0" * 8) - monkeypatch.setattr(huggingface_hub, "try_to_load_from_cache", lambda *a, **k: str(cached)) - fake = _install(monkeypatch, []) # must not be called - - out = xf.hf_hub_download_with_xet_fallback(DL_REPO, FILE, None) - assert out == str(cached) - assert fake.calls == [], "spawned a download for an already-cached file" - - -def test_cancel_before_start_raises_no_attempt(monkeypatch): - fake = _install(monkeypatch, []) - ev = threading.Event() - ev.set() - with pytest.raises(RuntimeError, match = "Cancelled"): - xf.hf_hub_download_with_xet_fallback(DL_REPO, FILE, None, cancel_event = ev) - assert fake.calls == [] - - -def test_nonstall_error_propagates_without_fallback(monkeypatch): - fake = _install(monkeypatch, [("error", "RepositoryNotFoundError: 404 not found")]) - with pytest.raises(RuntimeError, match = "RepositoryNotFoundError"): - xf.hf_hub_download_with_xet_fallback(DL_REPO, FILE, None) - assert len(fake.calls) == 1, "deterministic error must not trigger an HTTP fallback" - assert fake.calls[0].disable_xet is False - - -def test_immediate_success_uses_xet_only(monkeypatch): - prepared = [] - monkeypatch.setattr( - "hub.utils.download_registry.prepare_cache_for_transport", - lambda *a, **k: prepared.append(a), - ) - fake = _install(monkeypatch, [("ok", "/cache/model.gguf")]) - out = xf.hf_hub_download_with_xet_fallback(DL_REPO, FILE, None) - assert out == "/cache/model.gguf" - assert len(fake.calls) == 1 and fake.calls[0].disable_xet is False - assert prepared == [], "no cache prep should run when Xet succeeds first try" - - -def test_stall_then_http_fallback_succeeds(monkeypatch): prepared = [] monkeypatch.setattr( "hub.utils.download_registry.prepare_cache_for_transport", lambda repo_type, repo_id, mode, *a, **k: prepared.append((repo_type, repo_id, mode)), ) - fake = _install(monkeypatch, [("stall", None), ("ok", "/cache/model.gguf")]) out = xf.hf_hub_download_with_xet_fallback(DL_REPO, FILE, None) assert out == "/cache/model.gguf" - assert len(fake.calls) == 2 - assert fake.calls[0].disable_xet is False # Xet first - assert fake.calls[1].disable_xet is True # HTTP fallback - assert prepared == [("model", DL_REPO, "http")], "must prep cache for HTTP before the retry" + assert seen_disable_xet == [False, True] # Xet first, then HTTP + assert prepared == [("model", DL_REPO, "http")], "shim must run Studio's marker-aware prep" -def test_second_stall_raises_download_stall_error(monkeypatch): - monkeypatch.setattr( - "hub.utils.download_registry.prepare_cache_for_transport", lambda *a, **k: None - ) - fake = _install(monkeypatch, [("stall", None), ("stall", None)]) - with pytest.raises(xf.DownloadStallError): - xf.hf_hub_download_with_xet_fallback(DL_REPO, FILE, None) - assert len(fake.calls) == 2 +def test_shim_snapshot_injects_studio_prepare(monkeypatch): + """The snapshot wrapper forwards Studio's marker-aware prep, like the file wrapper.""" + captured = {} + + def fake_snapshot(repo_id, **kwargs): + captured["repo_id"] = repo_id + captured["prepare_for_http_fn"] = kwargs.get("prepare_for_http_fn") + return "/tmp/snap-dir" + + monkeypatch.setattr(xf, "_shared_snapshot_download_with_xet_fallback", fake_snapshot) + out = xf.snapshot_download_with_xet_fallback("org/model") + assert out == "/tmp/snap-dir" + assert captured["repo_id"] == "org/model" + assert captured["prepare_for_http_fn"] is xf._studio_prepare_for_http -def test_cancelled_midattempt_raises_no_fallback(monkeypatch): - fake = _install(monkeypatch, [("cancelled", None)]) - with pytest.raises(RuntimeError, match = "Cancelled"): - xf.hf_hub_download_with_xet_fallback(DL_REPO, FILE, None) - assert len(fake.calls) == 1 +def test_degrades_gracefully_without_shared_helper(monkeypatch): + """On an older unsloth_zoo lacking the shared helper, the shim still imports (Studio + boots) and exposes stub API doing plain HF downloads with the watchdog disabled.""" + import importlib + + class _BlockShared: + def find_spec( + self, + name, + path = None, + target = None, + ): + if name == "unsloth_zoo.hf_xet_fallback": + raise ModuleNotFoundError(f"No module named '{name}'", name = name) + return None + + finder = _BlockShared() + saved_shared = sys.modules.pop("unsloth_zoo.hf_xet_fallback", None) + saved_shim = sys.modules.pop("utils.hf_xet_fallback", None) + sys.meta_path.insert(0, finder) + try: + degraded = importlib.import_module("utils.hf_xet_fallback") + + # Boots without raising and mirrors the shared API surface. + assert issubclass(degraded.DownloadStallError, RuntimeError) + assert degraded.child_should_disable_xet({"disable_xet": True}) is True + assert degraded.get_hf_download_state(["x"]) is None # unmeasurable + event = degraded.start_watchdog(repo_ids = ["x"], on_stall = lambda m: None) + assert hasattr(event, "set") and not event.is_set() # never fires + + # Degraded mode still emits heartbeats so the inactivity deadline is not tripped. + import time as _time + + beats = [] + hb_stop = degraded.start_watchdog( + repo_ids = ["x"], + on_stall = lambda m: None, + on_heartbeat = beats.append, + interval = 0.02, + ) + try: + deadline = _time.monotonic() + 2.0 + while not beats and _time.monotonic() < deadline: + _time.sleep(0.02) + assert beats, "degraded watchdog emitted no heartbeat" + finally: + hb_stop.set() + + # Downloads fall back to plain huggingface_hub (no watchdog, no crash). + called = {} + + def _fake_snapshot(repo_id, **kwargs): + called["repo_id"] = repo_id + return "/snap-dir" + + monkeypatch.setattr(huggingface_hub, "snapshot_download", _fake_snapshot) + assert degraded.snapshot_download_with_xet_fallback("org/model") == "/snap-dir" + assert called["repo_id"] == "org/model" + + # Cancellation still holds: an already-set cancel_event aborts before the HF download. + import threading as _threading + + cancelled = _threading.Event() + cancelled.set() + called.clear() + with pytest.raises(RuntimeError, match = "Cancelled"): + degraded.snapshot_download_with_xet_fallback("org/model", cancel_event = cancelled) + assert "repo_id" not in called, "degraded download ran despite cancellation" + finally: + sys.meta_path.remove(finder) + sys.modules.pop("utils.hf_xet_fallback", None) + if saved_shared is not None: + sys.modules["unsloth_zoo.hf_xet_fallback"] = saved_shared + if saved_shim is not None: + sys.modules["utils.hf_xet_fallback"] = saved_shim -def test_per_file_independent_fallback(monkeypatch): - """A stalled shard falls back; a sibling shard that succeeds does not.""" - monkeypatch.setattr( - "hub.utils.download_registry.prepare_cache_for_transport", lambda *a, **k: None - ) - fake = _install(monkeypatch, [("ok", "/a"), ("stall", None), ("ok", "/b")]) - assert xf.hf_hub_download_with_xet_fallback(DL_REPO, "shardA.gguf", None) == "/a" - assert xf.hf_hub_download_with_xet_fallback(DL_REPO, "shardB.gguf", None) == "/b" - assert [c.disable_xet for c in fake.calls] == [False, False, True] +def test_degrades_when_unsloth_zoo_entirely_absent(): + """When unsloth_zoo is absent entirely, the import raises + ModuleNotFoundError(name='unsloth_zoo') (top-level package). Guard that the shim still + degrades and does not re-raise, breaking every Studio import that pulls it in.""" + import importlib + + class _BlockZoo: + def find_spec( + self, + name, + path = None, + target = None, + ): + # Whole package absent, so ModuleNotFoundError.name is the top-level 'unsloth_zoo'. + if name == "unsloth_zoo" or name.startswith("unsloth_zoo."): + raise ModuleNotFoundError("No module named 'unsloth_zoo'", name = "unsloth_zoo") + return None + + finder = _BlockZoo() + saved = { + k: v + for k, v in list(sys.modules.items()) + if k == "unsloth_zoo" or k.startswith("unsloth_zoo.") + } + for k in saved: + del sys.modules[k] + saved_shim = sys.modules.pop("utils.hf_xet_fallback", None) + sys.meta_path.insert(0, finder) + try: + degraded = importlib.import_module("utils.hf_xet_fallback") + # Boots without raising and exposes the stub API. + assert issubclass(degraded.DownloadStallError, RuntimeError) + assert degraded.get_hf_download_state(["x"]) is None + event = degraded.start_watchdog(repo_ids = ["x"], on_stall = lambda m: None) + assert hasattr(event, "set") and not event.is_set() + finally: + sys.meta_path.remove(finder) + sys.modules.pop("utils.hf_xet_fallback", None) + sys.modules.update(saved) + if saved_shim is not None: + sys.modules["utils.hf_xet_fallback"] = saved_shim -# --------------------------------------------------------------------------- # -# Precondition: HF_HUB_DISABLE_XET is read at import time, so assert its effect -# in a FRESH interpreter (huggingface/huggingface_hub#3266 once ignored it). -# --------------------------------------------------------------------------- # -def _safe_path() -> str: +def test_degrades_when_shared_helper_import_raises_importerror(): + """unsloth_zoo can be installed yet fail to import when torch is missing (llama.cpp/GGUF-only + Studio), raising ImportError not ModuleNotFoundError. The shim must degrade for that too.""" + import importlib + + class _BlockWithImportError: + def find_spec( + self, + name, + path = None, + target = None, + ): + if name == "unsloth_zoo.hf_xet_fallback": + # Mirror a torch-less install: a plain ImportError with no .name. + raise ImportError("Unsloth: Pytorch is not installed.") + return None + + finder = _BlockWithImportError() + saved_shared = sys.modules.pop("unsloth_zoo.hf_xet_fallback", None) + saved_zoo = sys.modules.pop("unsloth_zoo", None) + saved_shim = sys.modules.pop("utils.hf_xet_fallback", None) + sys.meta_path.insert(0, finder) + try: + degraded = importlib.import_module("utils.hf_xet_fallback") + assert issubclass(degraded.DownloadStallError, RuntimeError) + assert degraded.get_hf_download_state(["x"]) is None + event = degraded.start_watchdog(repo_ids = ["x"], on_stall = lambda m: None) + assert hasattr(event, "set") and not event.is_set() + finally: + sys.meta_path.remove(finder) + sys.modules.pop("utils.hf_xet_fallback", None) + if saved_shared is not None: + sys.modules["unsloth_zoo.hf_xet_fallback"] = saved_shared + if saved_zoo is not None: + sys.modules["unsloth_zoo"] = saved_zoo + if saved_shim is not None: + sys.modules["utils.hf_xet_fallback"] = saved_shim + + +def test_retries_under_light_gpu_init_when_import_fails(monkeypatch): + """GPU detection in unsloth_zoo's __init__ raises NotImplementedError on a GPU-less host. The shim + retries under UNSLOTH_ZOO_DISABLE_GPU_INIT=1, restores the env, and degrades if the retry fails.""" + import importlib import os - return os.environ.get("PATH", "") + monkeypatch.delenv("UNSLOTH_ZOO_DISABLE_GPU_INIT", raising = False) + seen_env = [] -def test_disable_xet_constant_set_in_fresh_interpreter(): - code = ( - "from huggingface_hub import constants as c; " - "import sys; sys.exit(0 if c.HF_HUB_DISABLE_XET is True else 17)" - ) - proc = subprocess.run( - [sys.executable, "-c", code], - env = {"HF_HUB_DISABLE_XET": "1", "PATH": _safe_path()}, - capture_output = True, - text = True, - ) - assert proc.returncode == 0, ( - f"HF_HUB_DISABLE_XET=1 did not set constants.HF_HUB_DISABLE_XET=True " - f"(rc={proc.returncode}): {proc.stderr}" - ) + class _GpuGatedBlocker: + def find_spec( + self, + name, + path = None, + target = None, + ): + # Crash is in unsloth_zoo's __init__, so intercept "unsloth_zoo" itself (the parent). + if name == "unsloth_zoo": + # Record the env each attempt sees; raise the no-GPU error both times so the shim + # degrades. + seen_env.append(os.environ.get("UNSLOTH_ZOO_DISABLE_GPU_INIT")) + raise NotImplementedError("Unsloth cannot find any torch accelerator") + return None - -def test_default_leaves_xet_enabled(): - code = ( - "from huggingface_hub import constants as c; " - "import sys; sys.exit(0 if c.HF_HUB_DISABLE_XET is False else 17)" - ) - proc = subprocess.run( - [sys.executable, "-c", code], - env = {"PATH": _safe_path()}, # no HF_HUB_DISABLE_XET - capture_output = True, - text = True, - ) - assert proc.returncode == 0, ( - f"without the env var, constants.HF_HUB_DISABLE_XET was not False " - f"(rc={proc.returncode}): {proc.stderr}" - ) + finder = _GpuGatedBlocker() + saved = { + k: v + for k, v in list(sys.modules.items()) + if k == "unsloth_zoo" or k.startswith("unsloth_zoo.") + } + for k in saved: + del sys.modules[k] + saved_shim = sys.modules.pop("utils.hf_xet_fallback", None) + sys.meta_path.insert(0, finder) + try: + degraded = importlib.import_module("utils.hf_xet_fallback") + # First attempt without the light env, then a retry with it set. + assert seen_env == [None, "1"], seen_env + # Both attempts raised -> Studio still boots in degraded mode. + assert issubclass(degraded.DownloadStallError, RuntimeError) + # The env override must not leak past the import. + assert os.environ.get("UNSLOTH_ZOO_DISABLE_GPU_INIT") is None + finally: + sys.meta_path.remove(finder) + sys.modules.pop("utils.hf_xet_fallback", None) + sys.modules.update(saved) + if saved_shim is not None: + sys.modules["utils.hf_xet_fallback"] = saved_shim diff --git a/studio/backend/tests/test_model_update_robustness.py b/studio/backend/tests/test_model_update_robustness.py index 9cf2a62c39..300eb587b3 100644 --- a/studio/backend/tests/test_model_update_robustness.py +++ b/studio/backend/tests/test_model_update_robustness.py @@ -5,8 +5,8 @@ Covers: * GGUF variant listing computes update_available from the already-fetched sibling metadata instead of a second Hub call. - * hf_hub_download_with_xet_fallback(force_download=True) bypasses the - try_to_load_from_cache cache-first early-return. + * hf_hub_download_with_xet_fallback forwards force_download through the shim to the + shared unsloth_zoo helper (which owns the cache-first early-return and its bypass). The cache "Update" action now runs through the download manager as a normal managed download (so it shows in the Downloads panel with progress + cancel), @@ -341,44 +341,26 @@ def test_cached_model_scan_keeps_local_safetensors_repo(monkeypatch, tmp_path): # ── hf_hub_download_with_xet_fallback force_download bypass (X2/F2) ─── -def test_force_download_bypasses_cache_first_early_return(monkeypatch): - """force_download=True skips the try_to_load_from_cache early-return and - proceeds to the real download path; force_download=False returns the cached - path without ever attempting a download (X2/F2).""" - import huggingface_hub as hf +def test_force_download_is_forwarded_through_the_shim(monkeypatch): + """The shim's contract is to forward force_download unchanged to the shared helper (which owns the + cache-first early-return and bypass). Verify both False and True reach it (X2/F2).""" import utils.hf_xet_fallback as X - cached_path = "/cache/blob/cached.gguf" + seen = [] - # Pretend the blob IS cached on disk (try_to_load_from_cache is imported - # inside the function from huggingface_hub, and os.path.exists must agree). - monkeypatch.setattr(hf, "try_to_load_from_cache", lambda *a, **k: cached_path, raising = False) - monkeypatch.setattr(X.os.path, "exists", lambda p: True, raising = False) + def fake_shared(repo_id, filename, token, **kwargs): + seen.append(kwargs.get("force_download")) + return "/downloaded/path" - attempts = [] + monkeypatch.setattr(X, "_shared_hf_hub_download_with_xet_fallback", fake_shared, raising = True) - def fake_attempt(repo_id, filename, token, **kwargs): - attempts.append( - {"repo_id": repo_id, "filename": filename, "force": kwargs.get("force_download")} - ) - return ("ok", "/freshly/downloaded/path") - - monkeypatch.setattr(X, "_run_download_attempt", fake_attempt, raising = True) - - # force_download=False: cache-first early-return, no download attempt. - out = X.hf_hub_download_with_xet_fallback( + X.hf_hub_download_with_xet_fallback( "unsloth/repo", "model.gguf", token = None, force_download = False ) - assert out == cached_path - assert attempts == [] # never reached the real download - - # force_download=True: bypass the early-return, run the real download. - out2 = X.hf_hub_download_with_xet_fallback( + X.hf_hub_download_with_xet_fallback( "unsloth/repo", "model.gguf", token = None, force_download = True ) - assert out2 == "/freshly/downloaded/path" - assert len(attempts) == 1 - assert attempts[0]["force"] is True + assert seen == [False, True] # the shim forwards force_download to the shared helper unchanged # ── multi-revision GGUF blob comparison and update reclaim ── diff --git a/studio/backend/utils/hf_xet_fallback.py b/studio/backend/utils/hf_xet_fallback.py index 15961ac03a..2dd2247396 100644 --- a/studio/backend/utils/hf_xet_fallback.py +++ b/studio/backend/utils/hf_xet_fallback.py @@ -1,341 +1,204 @@ # SPDX-License-Identifier: AGPL-3.0-only # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 -"""Xet-primary HF downloads with an automatic HTTP fallback on a no-progress stall. +"""Studio shim over the shared ``unsloth_zoo.hf_xet_fallback`` Xet -> HTTP stall fallback. -Xet (``hf_xet``) is the fast default but can hang with no progress and no -exception, and a blocked native thread cannot be killed. Keep Xet primary; fall -back to plain HTTP only when the parent observes a stall. ``HF_HUB_DISABLE_XET`` -is read at import time, so the fallback runs in a fresh ``spawn`` child (not a -thread) that sets the env before importing ``huggingface_hub``. Cached files -short-circuit with no child; deterministic errors (401/403/404/disk-full) and -cancellation propagate without a fallback. Mirrors the safetensors inference -recovery in core/inference/{orchestrator,worker}.py. +Re-exports the shared API and injects Studio's marker-aware cache purge +(``prepare_cache_for_transport``) so the download manager keeps its ``.transport`` +marker semantics on the HTTP retry. """ from __future__ import annotations -import multiprocessing as mp -import os -import queue -import signal -import sys import threading -import time from typing import Any, Callable, Optional -from loggers import get_logger +_shared_import_error = None +try: + import unsloth_zoo.hf_xet_fallback as _shared + _shared_available = True +except Exception as _exc: # noqa: BLE001 - any import failure must degrade, not crash + # unsloth_zoo's __init__ runs torch/GPU detection, which raises on a torch-less/GPU-less Studio + # host. The download helper needs none of it, so retry via the light UNSLOTH_ZOO_DISABLE_GPU_INIT + # path before giving up. + _shared_import_error = _exc + import os as _os -logger = get_logger(__name__) - -_CTX = mp.get_context("spawn") - -# Defaults match the existing inference watchdog and hub shutdown deadline. -DEFAULT_HEARTBEAT_INTERVAL = 30.0 -DEFAULT_STALL_TIMEOUT = 180.0 -DEFAULT_GRACE_PERIOD = 10.0 -_POLL_INTERVAL = 0.5 - - -class DownloadStallError(RuntimeError): - """Raised when no download progress is observed for too long. - - Canonical home; orchestrator.py re-imports it so all paths share one type. - """ - - -def child_should_disable_xet(config: dict) -> bool: - """Single source of truth for the per-worker Xet env flip.""" - return bool(config.get("disable_xet")) - - -def get_hf_download_state( - repo_ids: Optional[list[str]] = None, *, repo_type: str = "model" -) -> Optional[tuple[int, bool]]: - """Return ``(total_on_disk_bytes, has_incomplete)`` for the active HF cache. - - Sparse-aware (st_blocks based) so a sparse Xet/``hf_transfer`` ``.incomplete`` - is not mistaken for full-size progress. ``None`` means the state could not be - measured, so callers skip stall logic for that tick. - """ + _prev_gpu_init = _os.environ.get("UNSLOTH_ZOO_DISABLE_GPU_INIT") + _os.environ["UNSLOTH_ZOO_DISABLE_GPU_INIT"] = "1" try: - from hub.utils.hf_cache_state import ( - blob_bytes_present, - has_active_incomplete_blobs, - hf_cache_root, - iter_active_repo_cache_dirs, - ) + import unsloth_zoo.hf_xet_fallback as _shared + _shared_available = True + _shared_import_error = None + except Exception as _exc2: # noqa: BLE001 - degrade so Studio still boots with plain HF downloads + _shared_import_error = _exc2 + _shared_available = False + finally: + if _prev_gpu_init is None: + _os.environ.pop("UNSLOTH_ZOO_DISABLE_GPU_INIT", None) + else: + _os.environ["UNSLOTH_ZOO_DISABLE_GPU_INIT"] = _prev_gpu_init - if hf_cache_root() is None: - return (0, False) +if _shared_available: + # Bind by assignment so each public name shares one module-level binding with the degraded branch. + DEFAULT_GRACE_PERIOD = _shared.DEFAULT_GRACE_PERIOD + DEFAULT_HEARTBEAT_INTERVAL = _shared.DEFAULT_HEARTBEAT_INTERVAL + DEFAULT_STALL_TIMEOUT = _shared.DEFAULT_STALL_TIMEOUT + DownloadStallError = _shared.DownloadStallError + child_should_disable_xet = _shared.child_should_disable_xet + get_hf_download_state = _shared.get_hf_download_state + start_watchdog = _shared.start_watchdog + _shared_hf_hub_download_with_xet_fallback = _shared.hf_hub_download_with_xet_fallback + _shared_snapshot_download_with_xet_fallback = _shared.snapshot_download_with_xet_fallback +else: + # Degrade instead of crashing Studio: plain HF downloads, stall watchdog disabled. Thin stubs, + # not a second copy of the orchestration; recovery returns once unsloth_zoo is upgraded. + import logging as _logging - total = 0 - has_incomplete = False - for repo_id in repo_ids or []: - # Skip local paths: HF IDs never start with / . ~ or contain "\". - if not repo_id or repo_id.startswith(("/", ".", "~")) or "\\" in repo_id: - continue - for entry in iter_active_repo_cache_dirs(repo_type, repo_id): - blobs_dir = entry / "blobs" - if not blobs_dir.is_dir(): - continue - for blob in blobs_dir.iterdir(): - try: - if blob.is_file(): - total += blob_bytes_present(blob) - except OSError: - pass - if has_active_incomplete_blobs(repo_type, repo_id): - has_incomplete = True - return (total, has_incomplete) - except Exception as e: - logger.debug("Failed to determine HF download state: %s", e) - return None + _logging.getLogger(__name__).warning( + "unsloth_zoo.hf_xet_fallback unavailable (%s); the Xet stall watchdog is " + "disabled. Install/upgrade unsloth_zoo (and its torch dependency) to " + "re-enable automatic Xet -> HTTP download recovery.", + _shared_import_error, + ) + DEFAULT_HEARTBEAT_INTERVAL = 30.0 + DEFAULT_STALL_TIMEOUT = 180.0 + DEFAULT_GRACE_PERIOD = 10.0 -def start_watchdog( - *, - repo_ids: list[str], - on_stall: Callable[[str], None], - repo_type: str = "model", - interval: float = DEFAULT_HEARTBEAT_INTERVAL, - stall_timeout: float = DEFAULT_STALL_TIMEOUT, - xet_disabled: bool = False, - on_heartbeat: Optional[Callable[[str], None]] = None, -) -> threading.Event: - """Start a daemon thread that fires ``on_stall(message)`` exactly once iff a - ``*.incomplete`` is present AND the on-disk size is unchanged for - *stall_timeout* seconds. The timer resets while no ``*.incomplete`` exists, so - post-download init is never misread as a stall. Returns a stop event the - caller sets when the download phase ends. - """ - stop = threading.Event() - transport = "https" if xet_disabled else "xet" - fired = False + class DownloadStallError(RuntimeError): + """Stub mirror so callers' ``except`` clauses resolve; never raised in degraded mode.""" - def _beat() -> None: - nonlocal fired - state = get_hf_download_state(repo_ids, repo_type = repo_type) - last_size = state[0] if state is not None else 0 - last_change = time.monotonic() + def child_should_disable_xet(config: dict) -> bool: + return bool(config.get("disable_xet")) - while not stop.wait(interval): - state = get_hf_download_state(repo_ids, repo_type = repo_type) - now = time.monotonic() + def get_hf_download_state(*args: Any, **kwargs: Any) -> None: + return None # unmeasurable -> the (absent) watchdog never fires - if state is None: - if on_heartbeat is not None: + def start_watchdog( + *, + on_heartbeat: "Optional[Callable[[str], None]]" = None, + interval: float = DEFAULT_HEARTBEAT_INTERVAL, + xet_disabled: bool = False, + **kwargs: Any, + ) -> "threading.Event": + # No stall detection, but keep emitting heartbeats so the orchestrator's inactivity deadline + # is not tripped during a long download. + stop = threading.Event() + if on_heartbeat is None: + return stop + transport = "https" if xet_disabled else "xet" + + def _beat() -> None: + while not stop.wait(interval): + try: on_heartbeat(f"Downloading ({transport} transport)...") - continue + except Exception: + pass - current_size, has_incomplete = state - if current_size != last_size: - last_size = current_size - last_change = now + threading.Thread( + target = _beat, + daemon = True, + name = "hf-xet-degraded-heartbeat", + ).start() + return stop - # Reset unless .incomplete confirms an active download, so model init - # and lock waits are not counted as a stall. - if not has_incomplete: - last_change = now - elif now - last_change >= stall_timeout: - if not fired: - fired = True - on_stall( - f"Download appears stalled ({transport} transport) " - f"-- no progress for {int(now - last_change)}s" - ) - return + def _degraded_cancelled(cancel_event: "Optional[threading.Event]") -> bool: + return cancel_event is not None and cancel_event.is_set() - if on_heartbeat is not None: - on_heartbeat(f"Downloading ({transport} transport)...") + def _shared_hf_hub_download_with_xet_fallback( + repo_id: str, + filename: str, + token: Optional[str], + *, + repo_type: str = "model", + revision: Optional[str] = None, + cache_dir: Optional[str] = None, + force_download: bool = False, + cancel_event: "Optional[threading.Event]" = None, + **_ignored: Any, + ) -> str: + # Keep the cancellation contract: do not start or return a download once cancelled. + if _degraded_cancelled(cancel_event): + raise RuntimeError("Cancelled") - threading.Thread(target = _beat, daemon = True, name = "hf-xet-watchdog").start() - return stop - - -def _download_child_entry( - *, - repo_id: str, - filename: str, - token: Optional[str], - repo_type: str, - disable_xet: bool, - result_queue: Any, - force_download: bool = False, -) -> None: - """Spawn-child entrypoint: download one file and report the result. - - Top-level and picklable. Sets the Xet env BEFORE importing huggingface_hub, - forms its own process group so the parent can kill the whole transfer, and - never logs the token or signed URLs. - """ - # Die with Studio on Linux (this mp child gets no parent-set preexec_fn). - try: - from utils.process_lifetime import bind_current_process_to_parent_lifetime - bind_current_process_to_parent_lifetime() - except Exception: - pass - - if hasattr(os, "setsid"): - try: - os.setsid() - except OSError: - pass - - if disable_xet: - os.environ["HF_HUB_DISABLE_XET"] = "1" - # Keep the HTTP writer sequential and resumable (hf_transfer leaves sparse - # partials a sequential resume cannot safely continue). - os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "0" - os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1") - - # Test-only fault injection (never set in production): stall the Xet attempt - # so the watchdog + HTTP fallback can be exercised against a real repo. - if not disable_xet and os.environ.get("UNSLOTH_HF_XET_FORCE_STALL") == "1": - import time as _t - try: - from huggingface_hub.constants import HF_HUB_CACHE - - blobs = os.path.join(HF_HUB_CACHE, "models--" + repo_id.replace("/", "--"), "blobs") - os.makedirs(blobs, exist_ok = True) - with open(os.path.join(blobs, "xet-force-stall.incomplete"), "wb") as fh: - fh.write(b"\0" * 4096) - except OSError: - pass - while True: - _t.sleep(3600) - - try: from huggingface_hub import hf_hub_download + path = hf_hub_download( repo_id = repo_id, filename = filename, - repo_type = repo_type, token = token, + repo_type = repo_type, + revision = revision, + cache_dir = cache_dir, force_download = force_download, ) - result_queue.put({"ok": True, "path": path}) - except BaseException as e: # noqa: BLE001 - report every failure to the parent - error = f"{type(e).__name__}: {e}" - try: - from hub.utils.download_registry import scrub_secrets - error = scrub_secrets(error, hf_token = token) - except Exception: - pass - result_queue.put({"ok": False, "error": error}) + if _degraded_cancelled(cancel_event): + raise RuntimeError("Cancelled") + return path + def _shared_snapshot_download_with_xet_fallback( + repo_id: str, + *, + revision: Optional[str] = None, + token: Optional[str] = None, + repo_type: str = "model", + cache_dir: Optional[str] = None, + allow_patterns: Optional[Any] = None, + ignore_patterns: Optional[Any] = None, + force_download: bool = False, + cancel_event: "Optional[threading.Event]" = None, + **_ignored: Any, + ) -> str: + if _degraded_cancelled(cancel_event): + raise RuntimeError("Cancelled") -def _terminate_process_group(proc: "mp.process.BaseProcess", grace_period: float) -> None: - """Kill *proc* and its whole process group (Xet may spawn helper procs). + from huggingface_hub import snapshot_download - The child calls ``os.setsid()`` so its pgid equals its pid; signal via - ``os.killpg(pid, ...)`` -- NOT ``getpgid``, which before the child becomes a - group leader resolves to OUR group. SIGTERM, then SIGKILL after *grace_period*. - """ - pid = proc.pid - - def _signal_group(sig: int) -> None: - if pid is not None and hasattr(os, "killpg"): - try: - os.killpg(pid, sig) - return - except (ProcessLookupError, PermissionError, OSError): - pass - # Windows or pre-setsid: best effort on the single process. - try: - proc.terminate() if sig != getattr(signal, "SIGKILL", -9) else proc.kill() - except Exception: - pass - - _signal_group(getattr(signal, "SIGTERM", signal.SIGINT)) - proc.join(timeout = grace_period) - if proc.is_alive(): - _signal_group(getattr(signal, "SIGKILL", signal.SIGTERM)) - proc.join(timeout = 5.0) - - -def _run_download_attempt( - repo_id: str, - filename: str, - token: Optional[str], - *, - repo_type: str, - disable_xet: bool, - cancel_event: Optional[threading.Event], - stall_timeout: float, - interval: float, - grace_period: float, - on_status: Optional[Callable[[str], None]], - force_download: bool = False, -) -> tuple[str, Optional[str]]: - """Run one download in a spawn child supervised by the no-progress watchdog. - - Returns ``("ok", path)``, ``("stall", None)``, ``("cancelled", None)``, or - ``("error", message)``. This is the seam tests monkeypatch to avoid spawning. - """ - result_queue: Any = _CTX.Queue() - proc = _CTX.Process( - target = _download_child_entry, - kwargs = dict( + path = snapshot_download( repo_id = repo_id, - filename = filename, - token = token, repo_type = repo_type, - disable_xet = disable_xet, - result_queue = result_queue, + revision = revision, + token = token, + cache_dir = cache_dir, + allow_patterns = allow_patterns, + ignore_patterns = ignore_patterns, force_download = force_download, - ), - daemon = True, - ) - proc.start() - from utils.process_lifetime import adopt_pid - - adopt_pid(proc.pid) # bind to parent lifetime (Windows job / sweep) - - stalled = threading.Event() - stop_watchdog = start_watchdog( - repo_ids = [repo_id], - on_stall = lambda msg: stalled.set(), - repo_type = repo_type, - interval = interval, - stall_timeout = stall_timeout, - xet_disabled = disable_xet, - on_heartbeat = on_status, - ) - - result: Optional[dict] = None - try: - while proc.is_alive(): - if cancel_event is not None and cancel_event.is_set(): - _terminate_process_group(proc, grace_period) - return ("cancelled", None) - if stalled.is_set(): - _terminate_process_group(proc, grace_period) - return ("stall", None) - try: - result = result_queue.get(timeout = _POLL_INTERVAL) - break - except queue.Empty: - continue - else: - # Process exited; drain any result it enqueued. - try: - result = result_queue.get_nowait() - except queue.Empty: - result = None - finally: - stop_watchdog.set() - proc.join(timeout = grace_period) - - if result is None: - return ( - "error", - f"download process for '{repo_id}/{filename}' exited " - f"(code={proc.exitcode}) without a result", ) - if result.get("ok"): - return ("ok", result["path"]) - return ("error", result.get("error") or "unknown download error") + if _degraded_cancelled(cancel_event): + raise RuntimeError("Cancelled") + return path + + +__all__ = [ + "DEFAULT_GRACE_PERIOD", + "DEFAULT_HEARTBEAT_INTERVAL", + "DEFAULT_STALL_TIMEOUT", + "DownloadStallError", + "child_should_disable_xet", + "get_hf_download_state", + "start_watchdog", + "hf_hub_download_with_xet_fallback", + "snapshot_download_with_xet_fallback", +] + + +def _studio_prepare_for_http(repo_type: str, repo_id: str) -> None: + """Studio's marker-aware purge before an HTTP resume, keeping the download manager's ``.transport`` + accounting consistent (vs unsloth_zoo's generic default). Guarded: a purge failure is logged, + not fatal to the retry.""" + try: + from hub.utils.download_registry import prepare_cache_for_transport + prepare_cache_for_transport(repo_type, repo_id, "http") + except Exception as exc: + try: + from loggers import get_logger + get_logger(__name__).debug( + "Studio prepare_cache_for_transport failed for %s: %s", repo_id, exc + ) + except ModuleNotFoundError as logger_exc: + if logger_exc.name != "loggers": + raise def hf_hub_download_with_xet_fallback( @@ -345,83 +208,32 @@ def hf_hub_download_with_xet_fallback( *, cancel_event: Optional[threading.Event] = None, repo_type: str = "model", + revision: Optional[str] = None, stall_timeout: float = DEFAULT_STALL_TIMEOUT, interval: float = DEFAULT_HEARTBEAT_INTERVAL, grace_period: float = DEFAULT_GRACE_PERIOD, on_status: Optional[Callable[[str], None]] = None, force_download: bool = False, ) -> str: - """Download a single file with Xet primary and HTTP as a stall-only fallback. + """Single-file download via the shared fallback with Studio's marker-aware HTTP-retry prep. + ``force_download`` re-fetches a newer blob over a cached one (Studio's model-update path).""" + return _shared_hf_hub_download_with_xet_fallback( + repo_id, + filename, + token, + cancel_event = cancel_event, + repo_type = repo_type, + revision = revision, + stall_timeout = stall_timeout, + interval = interval, + grace_period = grace_period, + on_status = on_status, + force_download = force_download, + prepare_for_http_fn = _studio_prepare_for_http, + ) - Returns the local cache path. Raises ``RuntimeError("Cancelled")`` if - *cancel_event* is set, re-raises a deterministic child error unchanged (no - fallback), and raises ``DownloadStallError`` only if BOTH transports stall. - When *force_download* is True the cache-first early-return is skipped and the - flag is threaded to ``hf_hub_download`` so a newer remote blob is re-fetched - even if an older blob is already cached. - """ - # Finalized blob already cached: return it with no child and no network. - # Skipped when force_download is set so an update re-fetches a newer blob. - if not force_download: - try: - from huggingface_hub import try_to_load_from_cache - cached = try_to_load_from_cache(repo_id, filename, repo_type = repo_type) - if isinstance(cached, str) and os.path.exists(cached): - return cached - except Exception as e: - logger.debug("Cached probe failed for %s/%s: %s", repo_id, filename, e) - - if cancel_event is not None and cancel_event.is_set(): - raise RuntimeError("Cancelled") - - disable_xet = False - for attempt in range(2): - if disable_xet: - # Purge a non-HTTP partial before resuming over HTTP: an HTTP resume - # over a sparse Xet/hf_transfer partial silently corrupts the blob. - try: - from hub.utils.download_registry import prepare_cache_for_transport - prepare_cache_for_transport(repo_type, repo_id, "http") - except Exception as e: - logger.debug("prepare_cache_for_transport failed for %s: %s", repo_id, e) - - kind, payload = _run_download_attempt( - repo_id, - filename, - token, - repo_type = repo_type, - disable_xet = disable_xet, - cancel_event = cancel_event, - stall_timeout = stall_timeout, - interval = interval, - grace_period = grace_period, - on_status = on_status, - force_download = force_download, - ) - - if kind == "ok": - return payload # type: ignore[return-value] - if kind == "cancelled": - raise RuntimeError("Cancelled") - if kind == "error": - # Deterministic failure: the other transport would fail identically. - raise RuntimeError(payload) - # kind == "stall" - if attempt == 0 and not disable_xet: - logger.warning( - "Download stalled for '%s/%s' -- retrying with HF_HUB_DISABLE_XET=1", - repo_id, - filename, - ) - if on_status is not None: - on_status(f"{repo_id}/{filename}: Xet stalled, retrying over HTTP") - disable_xet = True - continue - raise DownloadStallError( - f"Download stalled for '{repo_id}/{filename}' even with " - f"HF_HUB_DISABLE_XET=1 -- check your network connection" - ) - - # Unreachable: the loop either returns or raises on each attempt. - raise DownloadStallError(f"Download failed for '{repo_id}/{filename}'") +def snapshot_download_with_xet_fallback(repo_id: str, **kwargs: Any) -> str: + """Whole-repo download via the shared fallback with Studio's marker-aware HTTP-retry prep.""" + kwargs.setdefault("prepare_for_http_fn", _studio_prepare_for_http) + return _shared_snapshot_download_with_xet_fallback(repo_id, **kwargs) diff --git a/studio/frontend/src/components/assistant-ui/model-selector/model-update-action.tsx b/studio/frontend/src/components/assistant-ui/model-selector/model-update-action.tsx index d00c812325..db7628777a 100644 --- a/studio/frontend/src/components/assistant-ui/model-selector/model-update-action.tsx +++ b/studio/frontend/src/components/assistant-ui/model-selector/model-update-action.tsx @@ -42,10 +42,8 @@ export function ModelUpdateAction({ }: ModelUpdateActionProps) { const [open, setOpen] = useState(false); - // The update is a managed download (it surfaces in the global Downloads panel - // with progress + cancel). When this exact repo+variant finishes, refresh the - // caller so the "update available" cue clears once the new revision is on - // disk. A ref keeps the subscription stable across renders without resubscribing. + // Refresh the caller when this repo+variant's download finishes so the "update available" cue + // clears. A ref keeps the subscription stable across renders. const onUpdatedRef = useRef(onUpdated); onUpdatedRef.current = onUpdated; useEffect(() => { @@ -60,9 +58,8 @@ export function ModelUpdateAction({ }, [repoId, variant]); const handleConfirm = useCallback(() => { - // Start the background re-download and close the dialog immediately; the - // Downloads panel owns progress + cancel from here. Only a failure to START - // surfaces a toast — a failed download reports itself in the panel. + // Start the re-download and close the dialog; the Downloads panel owns progress + cancel. + // Only a failure to START toasts (a failed download shows in the panel). void Promise.resolve() .then(onConfirm) .catch((err) => { diff --git a/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx b/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx index 20000c82ee..9890c5f574 100644 --- a/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx +++ b/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx @@ -1247,11 +1247,8 @@ export function HubModelPicker({ onEject?: () => void; }) { const gpu = useGpuInfo(); - // The currently-loaded/running model id. We read params.checkpoint from the - // runtime store (backend-mirrored from /api/inference/status.active_model, see - // chat-runtime-store) rather than the dropdown `isSelected` highlight (which is - // just `value === repo_id` and can reflect a staged, not-yet-loaded pick). Used - // to disable the cached-row update action for the model that's live in memory. + // Live model id from the runtime store (backend-mirrored active_model), not the dropdown + // highlight which can be a staged pick. Disables the update action for it. const loadedModelId = useChatRuntimeStore((s) => s.params.checkpoint); // Last-loaded timestamps power the "Recent" sort (vs "Downloaded" = file date). const loadTimes = useModelLoadTimes(value); @@ -1589,11 +1586,8 @@ export function HubModelPicker({ refreshLocalModelsList(); }, [hfToken, refreshLocalModelsList]); - // Updates run as MANAGED downloads (they show in the global Downloads panel - // with manifest-based progress + a working Cancel), instead of a blocking - // call. The worker re-resolves `main` and pulls only changed blobs, so the - // cached copy stays usable until the new revision lands. The row's - // ModelUpdateAction refreshes the list when this repo+variant completes. + // Updates run as managed downloads (Downloads panel: progress + Cancel), not a blocking + // call. The worker pulls only changed blobs, so the cached copy stays usable until done. const startManagedUpdate = useCallback((repoId: string, variant: string, expectedBytes: number) => { return downloadManager .requestStart({ diff --git a/tests/saving/test_quant_method_none_normalization.py b/tests/saving/test_quant_method_none_normalization.py new file mode 100644 index 0000000000..c1c5fd3686 --- /dev/null +++ b/tests/saving/test_quant_method_none_normalization.py @@ -0,0 +1,81 @@ +"""CPU-only regression for the quant-method normalization loops in save.py. + +`unsloth_save_pretrained_gguf` and `save_to_gguf_generic` each normalize the +`quantization_method` list, mapping a ``None`` element to ``"q8_0"``. The mapping +used to call ``quant_method.lower()`` as the first statement of the loop, so a +``None`` element (e.g. ``quantization_method=[None]`` or ``["q4_k_m", None]``) +raised ``AttributeError: 'NoneType' object has no attribute 'lower'`` and the +``elif quant_method is None`` branch was unreachable dead code. + +The loop is inline inside two heavy functions (importing unsloth needs +unsloth_zoo / a GPU), so - like test_is_gpt_oss_detection.py - we extract just the +loop source via ``ast`` and exec it against sample inputs. That exercises the real +source: it fails on the old ordering and passes once ``None`` is handled first. +""" + +from __future__ import annotations + +import ast +from pathlib import Path + +import pytest + +SAVE_PY = Path(__file__).resolve().parents[2] / "unsloth" / "save.py" +SAVE_SRC = SAVE_PY.read_text(encoding = "utf-8") +SAVE_TREE = ast.parse(SAVE_SRC, filename = str(SAVE_PY)) + +# The target functions and the list variable each one appends the normalized method to. +TARGETS = ( + ("unsloth_save_pretrained_gguf", "quantization_methods"), + ("save_to_gguf_generic", "new_quantization_methods"), +) + + +def _func(tree, name): + for node in ast.walk(tree): + if isinstance(node, ast.FunctionDef) and node.name == name: + return node + raise AssertionError(f"function {name!r} not found in {SAVE_PY.name}") + + +def _quant_loop(func_name): + # The quant-normalization `for` loop iterates `quantization_method`; grab its source. + func = _func(SAVE_TREE, func_name) + for node in ast.walk(func): + if ( + isinstance(node, ast.For) + and isinstance(node.iter, ast.Call) + and isinstance(node.iter.func, ast.Name) + and node.iter.func.id == "enumerate" + and isinstance(node.iter.args[0], ast.Name) + and node.iter.args[0].id == "quantization_method" + ): + return node + raise AssertionError(f"quant-normalization loop not found in {func_name}") + + +def _run_loop(func_name, out_var, quantization_method): + # exec just the extracted loop against a given input, returning the appended methods. + loop_src = ast.get_source_segment(SAVE_SRC, _quant_loop(func_name)) + namespace = {out_var: [], "quantization_method": quantization_method} + exec(loop_src, {"__builtins__": __builtins__}, namespace) + return namespace[out_var] + + +@pytest.mark.parametrize("func_name, out_var", TARGETS) +def test_none_element_maps_to_q8_0(func_name, out_var): + # A bare None inside the list must map to q8_0, not raise AttributeError. + assert _run_loop(func_name, out_var, [None]) == ["q8_0"] + + +@pytest.mark.parametrize("func_name, out_var", TARGETS) +def test_none_mixed_with_strings(func_name, out_var): + # None resolves to q8_0 while sibling string methods are still normalized (lowercased). + assert _run_loop(func_name, out_var, ["Q4_K_M", None]) == ["q4_k_m", "q8_0"] + + +@pytest.mark.parametrize("func_name, out_var", TARGETS) +def test_string_methods_unchanged(func_name, out_var): + # The fix must not alter behavior for the ordinary string inputs. + methods = ["not_quantized", "fast_quantized", "quantized", "Q8_0"] + assert _run_loop(func_name, out_var, methods) == ["f16", "q8_0", "q4_k_m", "q8_0"] diff --git a/tests/test_attn_impl_honor_explicit.py b/tests/test_attn_impl_honor_explicit.py new file mode 100644 index 0000000000..3fb7a2208f --- /dev/null +++ b/tests/test_attn_impl_honor_explicit.py @@ -0,0 +1,190 @@ +"""An explicit non-flash attention request must survive the flash disable path. + +When flash attention is disabled for a model, a caller who explicitly asked for +"sdpa" or "flex_attention" should keep that choice instead of being downgraded +to whatever the conservative supports_* fallback would pick. +""" + +import pytest + +from unsloth.models._utils import ( + _disable_flash_attention_if_needed, + resolve_attention_implementation, +) + + +def test_explicit_sdpa_is_honored_even_when_not_marked_supported(): + config = {} + result = _disable_flash_attention_if_needed( + config, + attn_implementation = "sdpa", + supports_sdpa = False, # conservative flag would have skipped sdpa + supports_flex_attention = False, + would_use_flash_attention = True, + disable_reason = "unit test forces flash disabled", + ) + assert result == "sdpa" + assert config.get("_attn_implementation") == "sdpa" + + +def test_explicit_flex_is_honored_when_supported(): + config = {} + result = _disable_flash_attention_if_needed( + config, + attn_implementation = "flex_attention", + supports_sdpa = True, + supports_flex_attention = True, + would_use_flash_attention = True, + disable_reason = "unit test forces flash disabled", + ) + assert result == "flex_attention" + assert config.get("_attn_implementation") == "flex_attention" + + +def test_explicit_flex_falls_back_when_not_supported(): + # flex_attention is False for known-broken/excluded configs (e.g. gpt_oss), + # so an explicit flex request must not select that backend - it falls back. + config = {} + result = _disable_flash_attention_if_needed( + config, + attn_implementation = "flex_attention", + supports_sdpa = True, + supports_flex_attention = False, + would_use_flash_attention = True, + disable_reason = "unit test forces flash disabled", + ) + assert result == "sdpa" + + +def test_synthesized_config_sdpa_is_not_treated_as_explicit(): + # The language loader seeds the config with attn_implementation="sdpa"; when the + # caller passes nothing, that synthesized value must not override the flex fallback + # for a model that supports flex but not sdpa. + config = {"attn_implementation": "sdpa"} + result = _disable_flash_attention_if_needed( + config, + attn_implementation = None, + supports_sdpa = False, + supports_flex_attention = True, + would_use_flash_attention = False, + disable_reason = "unit test forces flash disabled", + ) + assert result == "flex_attention" + + +def test_no_disable_reason_returns_request_untouched(): + result = _disable_flash_attention_if_needed( + {}, + attn_implementation = "flash_attention_2", + disable_reason = None, + ) + assert result == "flash_attention_2" + + +def test_flash_request_still_falls_back_when_disabled(): + config = {} + result = _disable_flash_attention_if_needed( + config, + attn_implementation = "flash_attention_2", + supports_sdpa = True, + would_use_flash_attention = True, + disable_reason = "unit test forces flash disabled", + ) + assert result == "sdpa" + + +def test_resolver_honors_explicit_sdpa_when_not_supported_and_flash_disabled(): + # End-to-end through the public resolver: an explicit sdpa request with a + # flash-disabled config (oversized head dim) and supports_sdpa=False must not be + # rewritten to eager by the resolver's own not-supports_sdpa guard. + config = {"model_type": "test", "head_dim": 512} # head_dim > 256 disables flash + result = resolve_attention_implementation( + model_class = None, + config = config, + requested_attn_implementation = "sdpa", + supports_sdpa = False, + ) + assert result == "sdpa" + assert config.get("_attn_implementation") == "sdpa" + + +def test_resolver_downgrades_non_explicit_sdpa_when_not_supported(): + # No explicit request: the model resolution seeds sdpa/eager and the guard must + # still downgrade a synthesized sdpa to eager for a model that cannot run it. + config = {"model_type": "test", "attn_implementation": "sdpa"} + result = resolve_attention_implementation( + model_class = None, + config = config, + requested_attn_implementation = None, + supports_sdpa = False, + ) + assert result == "eager" + + +def test_resolver_downgrades_explicit_sdpa_for_sdpa_excluded_model(): + # gpt_oss is in _SDPA_EXCLUDED_MODELS (sdpa is known-broken) and _FLASH_EXCLUDED_MODELS + # (flash disabled). Honoring an explicit sdpa request must not re-enable that broken + # backend: it downgrades to eager, mirroring how an explicit flex request falls back + # for _FLEX_EXCLUDED_MODELS. supports_sdpa=True proves the exclusion overrides even a + # model that otherwise advertises SDPA support. + config = {"model_type": "gpt_oss"} + result = resolve_attention_implementation( + model_class = None, + config = config, + requested_attn_implementation = "sdpa", + supports_sdpa = True, + ) + assert result == "eager" + assert config.get("_attn_implementation") == "eager" + + +@pytest.mark.parametrize("model_type", ["gemma3", "gemma3_text"]) +def test_resolver_downgrades_explicit_sdpa_for_disable_sdpa_model(model_type): + # gemma3 / gemma3_text are in DISABLE_SDPA_MODEL_NAMES: the loader forces + # supports_sdpa=False because their bundled SDPA modules are wrong. An explicit + # sdpa request with flash disabled must NOT re-enable that known-wrong path - it + # downgrades to eager, exactly like _SDPA_EXCLUDED_MODELS (gpt_oss). head_dim>256 + # disables flash to mirror the real flash-disabled scenario. + config = {"model_type": model_type, "head_dim": 512} + result = resolve_attention_implementation( + model_class = None, + config = config, + requested_attn_implementation = "sdpa", + supports_sdpa = False, + ) + assert result == "eager" + assert config.get("_attn_implementation") == "eager" + + +def test_resolver_does_not_overmatch_gemma3n_for_explicit_sdpa(): + # The "gemma3," trailing-comma guard must not match gemma3n: gemma3n is not in + # DISABLE_SDPA_MODEL_NAMES, so it stays a conservative (not known-wrong) model and an + # explicit sdpa request is still honored. Proves the substring match neither over- nor + # under-matches. + config = {"model_type": "gemma3n", "head_dim": 512} + result = resolve_attention_implementation( + model_class = None, + config = config, + requested_attn_implementation = "sdpa", + supports_sdpa = False, + ) + assert result == "sdpa" + assert config.get("_attn_implementation") == "sdpa" + + +def test_resolver_downgrades_synthesized_sdpa_for_disable_sdpa_model(): + # A synthesized/default sdpa (requested is None; the value came from config) on a + # DISABLE_SDPA_MODEL_NAMES model must still downgrade to eager. + config = {"model_type": "gemma3", "attn_implementation": "sdpa"} + result = resolve_attention_implementation( + model_class = None, + config = config, + requested_attn_implementation = None, + supports_sdpa = False, + ) + assert result == "eager" + + +if __name__ == "__main__": + import sys + sys.exit(pytest.main([__file__, "-q"])) diff --git a/tests/test_fp8_tiny_e8m0.py b/tests/test_fp8_tiny_e8m0.py new file mode 100644 index 0000000000..cf49c8c92f --- /dev/null +++ b/tests/test_fp8_tiny_e8m0.py @@ -0,0 +1,123 @@ +"""FP8 block-quant linear must handle tiny / non-tileable weights and e8m0 scales. + +Two things break the triton block path: + * a hidden dim not divisible by the activation block size (tiny test models), + * float8_e8m0fnu weight scales, which have no triton dtype mapping. +The forward falls back to a torch-native blockwise dequant + bf16 matmul; this +test checks that fallback runs finite forward + backward and matches a plain +dequant reference. +""" + +import pytest +import torch + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason = "needs CUDA") + + +def _reference(X, weight, scale, block): + # Expand the per-block scale to full weight shape and dequantize. + m, n = weight.shape + s = scale.to(torch.float32) + s = s.repeat_interleave(block[0], 0)[:m].repeat_interleave(block[1], 1)[:, :n] + W = (weight.to(torch.float32) * s).to(X.dtype) + return X @ W.T + + +def test_tiny_non_tileable_forward_backward_matches_reference(): + from unsloth.kernels.fp8 import FP8BlockQuantLinear + + torch.manual_seed(0) + dev = "cuda" + block = [128, 128] + m, n = 8, 8 # non-tileable, in-dim % 128 != 0 + weight = torch.randn(m, n, device = dev, dtype = torch.bfloat16) # (out=m, in=n) + scale = torch.rand(1, 1, device = dev, dtype = torch.float32) + 0.5 + X = torch.randn(4, n, device = dev, dtype = torch.bfloat16, requires_grad = True) + + out = FP8BlockQuantLinear.apply(X, weight, scale) + assert torch.isfinite(out).all(), "forward produced non-finite values" + + ref = _reference(X.detach(), weight, scale, block) + torch.testing.assert_close(out, ref, atol = 5e-2, rtol = 5e-2) + + out.sum().backward() + assert X.grad is not None and torch.isfinite(X.grad).all(), "backward non-finite" + + +def test_e8m0_scale_is_upcast_and_runs(): + from unsloth.kernels.fp8 import FP8BlockQuantLinear + + if not hasattr(torch, "float8_e8m0fnu"): + pytest.skip("torch build lacks float8_e8m0fnu") + + dev = "cuda" + m, n = 8, 8 + weight = torch.randn(m, n, device = dev, dtype = torch.bfloat16) + scale = (torch.rand(1, 1, device = dev) + 1.0).to(torch.float8_e8m0fnu) + X = torch.randn(4, n, device = dev, dtype = torch.bfloat16, requires_grad = True) + + out = FP8BlockQuantLinear.apply(X, weight, scale) + assert torch.isfinite(out).all() + out.sum().backward() + assert torch.isfinite(X.grad).all() + + +def test_rectangular_block_dequant_matches_reference(): + # Rectangular blocks (block_size[0] != block_size[1]) that tile evenly used to + # route through the triton weight_dequant kernel, which uses a single BLOCK_SIZE + # for both axes and mis-indexes the column scale. Verify the torch expansion path + # now matches the reference for a 64x256 weight with block [64, 128] (scale 1x2). + from unsloth.kernels.fp8 import _blockwise_weight_dequant_any_shape + + torch.manual_seed(0) + dev = "cuda" + block = [64, 128] + m, n = 64, 256 # evenly tiled: 64 % 64 == 0, 256 % 128 == 0 + weight = torch.randn(m, n, device = dev, dtype = torch.bfloat16) + # Distinct per-block column scales expose column mis-indexing. + scale = torch.tensor([[0.5, 3.0]], device = dev, dtype = torch.float32) + + W_deq = _blockwise_weight_dequant_any_shape(weight, scale, block, torch.bfloat16) + + s = scale.repeat_interleave(block[0], 0)[:m].repeat_interleave(block[1], 1)[:, :n] + ref = (weight.to(torch.float32) * s).to(torch.bfloat16) + torch.testing.assert_close(W_deq, ref, atol = 5e-3, rtol = 5e-3) + + +def test_e8m0_scale_preserves_non_default_block_size_attr(): + # An e8m0 scale carrying a non-default block_size attribute must keep it across + # the float32 upcast in forward; otherwise the lookup falls back to [128, 128] + # and a compatible layout is wrongly rejected as incompatible. + from unsloth.kernels.fp8 import FP8BlockQuantLinear + + if not hasattr(torch, "float8_e8m0fnu"): + pytest.skip("torch build lacks float8_e8m0fnu") + + torch.manual_seed(0) + dev = "cuda" + block = [64, 64] + # in-dim 96 is not divisible by block[1]=64 -> forward takes the torch dequant + # fallback (no fp8 matmul kernel). Scale shape (2, 2) validates for [64, 64] but + # not [128, 128] (which expects (1, 1)). + m, n = 128, 96 + weight = torch.randn(m, n, device = dev, dtype = torch.bfloat16) # no block_size attr + scale_f = torch.rand(2, 2, device = dev) + 1.0 + scale = scale_f.to(torch.float8_e8m0fnu) + scale.block_size = block # attribute lives on the scale, not the weight + X = torch.randn(4, n, device = dev, dtype = torch.bfloat16, requires_grad = True) + + # With [128, 128] this raises "not compatible with block size"; success proves + # the [64, 64] attribute survived the e8m0 -> float32 upcast. + out = FP8BlockQuantLinear.apply(X, weight, scale) + assert torch.isfinite(out).all() + + ref = _reference(X.detach(), weight, scale.to(torch.float32), block) + torch.testing.assert_close(out, ref, atol = 5e-2, rtol = 5e-2) + + out.sum().backward() + assert X.grad is not None and torch.isfinite(X.grad).all() + + +if __name__ == "__main__": + import sys + sys.exit(pytest.main([__file__, "-q"])) diff --git a/tests/test_moe_lora_targets.py b/tests/test_moe_lora_targets.py index 994d39f261..7f9b9a0485 100644 --- a/tests/test_moe_lora_targets.py +++ b/tests/test_moe_lora_targets.py @@ -49,3 +49,190 @@ def test_explicit_dotted_module_target_does_not_discover_moe_parameters(): ) is None ) + + +@pytest.mark.parametrize( + "target_modules", + [ + # Attention-only auto-regex lists every projection leaf (incl. gate/up/down) + # but its path segment is attention-only, so experts must NOT be targeted. + r"(?:\bmodel\.layers\.[\d]{1,}\.(?:self_attn|attention|attn|mixer)\.(?:q_proj|k_proj|v_proj|o_proj|gate_proj|up_proj|down_proj))", + ".*self_attn.*proj", + # An mlp path alternative with attention-only leaves is still attention-only. + r"model\.layers\.\d+\.(?:mlp|self_attn)\.(?:q_proj|k_proj|v_proj|o_proj)", + ], +) +def test_attention_only_regex_does_not_discover_moe_parameters(target_modules): + from unsloth.models._utils import get_moe_target_parameters + assert get_moe_target_parameters(_FakeMoeModel(), target_modules) is None + + +def test_single_leaf_regex_targets_only_that_projection(): + from unsloth.models._utils import get_moe_target_parameters + assert get_moe_target_parameters(_FakeMoeModel(), ".*experts.*down_proj") == [ + "mlp.experts.down_proj", + ] + assert get_moe_target_parameters(_FakeMoeModel(), ".*mlp.*gate_proj") == [ + "mlp.experts.gate_up_proj", + ] + + +def test_auto_regex_mlp_tag_block_discovers_moe_on_fused_models(): + # get_peft_regex on a fused-expert model lists only attention Linears as + # leaves; the mlp tag block is the remaining signal of MLP finetune intent. + from unsloth.models._utils import get_moe_target_parameters + both_auto = ( + r"(?:\bmodel\.layers\.[\d]{1,}\." + r"(?:self_attn|attention|attn|mixer|mlp|feed_forward|ffn|dense|mixer)\." + r"(?:(?:q_proj|k_proj|v_proj|o_proj)))" + ) + assert get_moe_target_parameters(_FakeMoeModel(), both_auto) == [ + "mlp.experts.gate_up_proj", + "mlp.experts.down_proj", + ] + + +def test_explicit_attention_only_list_does_not_discover_moe_parameters(): + # An explicit attention-only leaf list names no MLP projection, so experts + # must never be targeted. get_peft_model routes this ORIGINAL list (not the + # scoped regex) into detection precisely because family scoping makes + # get_peft_regex emit its full "mlp|feed_forward|ffn|dense" component block + # even for an attention-only request (see the regex below), which the + # string fallback cannot distinguish from the fused-expert auto regex. + from unsloth.models._utils import get_moe_target_parameters + + attn_only_list = ["q_proj", "k_proj", "v_proj", "o_proj"] + assert get_moe_target_parameters(_FakeMoeModel(), attn_only_list) is None + assert get_moe_target_parameters(_FakeMoeModel(), tuple(attn_only_list)) is None + + # The regex get_peft_regex emits for that same attention-only list under a + # vision-off family scope carries the mlp component block, so the string + # path would wrongly enable experts -- hence detection must use the list. + scoped_regex = ( + r"(?:.*?(?:language|text).*?" + r"(?:self_attn|attention|attn|mixer|mlp|feed_forward|ffn|dense|mixer).*?" + r"(?:q_proj|k_proj|v_proj|o_proj))" + ) + assert get_moe_target_parameters(_FakeMoeModel(), scoped_regex) == [ + "mlp.experts.gate_up_proj", + "mlp.experts.down_proj", + ] + + +def test_frozen_mlp_full_list_does_not_discover_moe_parameters(): + # Regression: an explicit list that names MLP leaves together with + # finetune_mlp_modules=False must NOT train experts. get_peft_regex scopes + # the MLP leaves out (its emitted regex carries no mlp tag block), so + # detection has to key on that SCOPED regex -- keying on the original list + # would let its gate/up/down leaves silently re-enable the frozen experts. + from unsloth.models._utils import ( + _select_moe_detection_targets, + get_moe_target_parameters, + ) + + original_list = [ + "q_proj", + "k_proj", + "v_proj", + "o_proj", + "gate_proj", + "up_proj", + "down_proj", + ] + # Representative of what get_peft_regex emits for that list under + # finetune_mlp_modules=False: attention-only path, no mlp component block. + scoped_regex = ( + r"(?:.*?(?:language|text).*?" + r"(?:self_attn|attention|attn|mixer).*?" + r"(?:q_proj|k_proj|v_proj|o_proj))" + ) + selected = _select_moe_detection_targets( + original_list, + scoped_regex, + finetune_mlp_modules = False, + finetune_language_layers = True, + ) + assert selected is scoped_regex + assert get_moe_target_parameters(_FakeMoeModel(), selected) is None + + +def test_frozen_language_full_list_does_not_discover_moe_parameters(): + # Vision-only request (finetune_language_layers=False) with a full leaf list + # must not reach the language-model experts either. + from unsloth.models._utils import ( + _select_moe_detection_targets, + get_moe_target_parameters, + ) + + original_list = ["q_proj", "gate_proj", "up_proj", "down_proj"] + scoped_regex = ( + r"(?:.*?(?:vision|visual|image).*?" + r"(?:self_attn|attention|attn|mixer).*?" + r"(?:q_proj|k_proj|v_proj|o_proj))" + ) + selected = _select_moe_detection_targets( + original_list, + scoped_regex, + finetune_mlp_modules = True, + finetune_language_layers = False, + ) + assert selected is scoped_regex + assert get_moe_target_parameters(_FakeMoeModel(), selected) is None + + +def test_in_scope_mlp_full_list_still_discovers_moe_parameters(): + # With MLP and language both in scope, an explicit list that names MLP + # leaves SHOULD enable the experts (unchanged behavior): the original list + # is preferred and carries the gate/up/down intent. + from unsloth.models._utils import ( + _select_moe_detection_targets, + get_moe_target_parameters, + ) + + original_list = [ + "q_proj", + "k_proj", + "v_proj", + "o_proj", + "gate_proj", + "up_proj", + "down_proj", + ] + scoped_regex = r".*self_attn.*proj" # unused: original list is preferred + selected = _select_moe_detection_targets( + original_list, + scoped_regex, + finetune_mlp_modules = True, + finetune_language_layers = True, + ) + assert selected is original_list + assert get_moe_target_parameters(_FakeMoeModel(), selected) == [ + "mlp.experts.gate_up_proj", + "mlp.experts.down_proj", + ] + + +def test_attention_only_list_prefers_original_when_in_scope(): + # The case the PR originally fixed: an attention-only list routed through + # get_peft_regex under a family scope (e.g. vision-off) still keeps experts + # off, because with MLP+language in scope detection uses the original + # attention-only list rather than the regex's spurious mlp component block. + from unsloth.models._utils import ( + _select_moe_detection_targets, + get_moe_target_parameters, + ) + + attn_only_list = ["q_proj", "k_proj", "v_proj", "o_proj"] + scoped_regex = ( # carries the spurious mlp block get_peft_regex always adds + r"(?:.*?(?:language|text).*?" + r"(?:self_attn|attention|attn|mixer|mlp|feed_forward|ffn|dense).*?" + r"(?:q_proj|k_proj|v_proj|o_proj))" + ) + selected = _select_moe_detection_targets( + attn_only_list, + scoped_regex, + finetune_mlp_modules = True, + finetune_language_layers = True, + ) + assert selected is attn_only_list + assert get_moe_target_parameters(_FakeMoeModel(), selected) is None diff --git a/tests/test_prefetch_snapshot_scope.py b/tests/test_prefetch_snapshot_scope.py new file mode 100644 index 0000000000..c7ec4f2c34 --- /dev/null +++ b/tests/test_prefetch_snapshot_scope.py @@ -0,0 +1,916 @@ +# Unsloth Zoo - Utilities for Unsloth +# 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 Affero 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 Affero General Public License for more details. +# +# You should have received a copy of the GNU Affero General Public License +# along with this program. If not, see . + +"""Pure-CPU, no-network unit tests for prefetch snapshot scoping in unsloth/models/_utils.py. + +maybe_prefetch_hf_snapshot warms the HF cache before the in-process load. The warm must cover at +least what the load reads (else the missing file falls to an unprotected in-process Xet fetch) but +not pull weights the load never reads. These tests lock the allow/ignore patterns each mode hands +snapshot_download_with_xet_fallback. The zoo downloader is monkeypatched to capture its kwargs. +""" + +import fnmatch +import sys +import types + +import pytest + +from unsloth.models import _utils as U + + +def _filter(names, allow_patterns, ignore_patterns): + """Mirror HF filter_repo_objects: keep on allow match (or None), drop on ignore match.""" + kept = [] + for name in names: + if allow_patterns is not None and not any(fnmatch.fnmatch(name, p) for p in allow_patterns): + continue + if ignore_patterns and any(fnmatch.fnmatch(name, p) for p in ignore_patterns): + continue + kept.append(name) + return kept + + +@pytest.fixture +def capture(monkeypatch): + """Run maybe_prefetch_hf_snapshot with a fake repo, capturing the patterns forwarded to a + fake injected zoo downloader (independent of the installed unsloth_zoo). Offline env cleared.""" + monkeypatch.delenv("HF_HUB_OFFLINE", raising = False) + monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False) + + state = {} + + def fake_download(repo_id, **kw): + state["repo_id"] = repo_id + state["allow_patterns"] = kw.get("allow_patterns") + state["ignore_patterns"] = kw.get("ignore_patterns") + state["variant"] = kw.get("variant") + return "/tmp/fake-snapshot" + + fake_module = types.ModuleType("unsloth_zoo.hf_xet_fallback") + fake_module.snapshot_download_with_xet_fallback = fake_download + fake_module.DownloadStallError = type("DownloadStallError", (RuntimeError,), {}) + monkeypatch.setitem(sys.modules, "unsloth_zoo.hf_xet_fallback", fake_module) + + # Neutralize the model_info network call by default; tests exercising format selection + # install their own. + import huggingface_hub + + class _NoNetworkApi: + def model_info(self, *a, **k): + raise RuntimeError("no network in test") + + monkeypatch.setattr(huggingface_hub, "HfApi", _NoNetworkApi) + + def run(**call_kwargs): + state.clear() + ok = U.maybe_prefetch_hf_snapshot("some-org/some-repo", **call_kwargs) + return ok, state + + return run + + +# Representative repo listing: root weights + aux, subdir, adapter, checkpoint, merged weights. +_SAMPLE_FILES = [ + "config.json", + "tokenizer.json", + "tokenizer_config.json", + "model-00001-of-00002.safetensors", + "model-00002-of-00002.safetensors", + "model.safetensors.index.json", + "pytorch_model.bin", + "fp16/model.safetensors", + "experimental/model-00001-of-00002.safetensors", + "checkpoint-500/model.safetensors", + "adapter_config.json", + "adapter_model.safetensors", +] + + +def test_weights_at_root_excludes_subdir_weights(capture): + """A root load ignores subdir weights (fp16/, experimental/, checkpoint-500/) but keeps root weights.""" + ok, st = capture(weights_at_root = True, use_safetensors = True) + assert ok is True + assert st["allow_patterns"] is None + ig = st["ignore_patterns"] + assert "*/*.safetensors" in ig and "*/*.bin" in ig + kept = _filter(_SAMPLE_FILES, st["allow_patterns"], ig) + assert "model-00001-of-00002.safetensors" in kept + assert "model.safetensors.index.json" in kept + assert "config.json" in kept + assert "fp16/model.safetensors" not in kept + assert "experimental/model-00001-of-00002.safetensors" not in kept + assert "checkpoint-500/model.safetensors" not in kept + + +def test_adapter_only_excludes_merged_weights(capture): + """An adapter warm keeps adapter files + root aux, not merged full-model weights.""" + ok, st = capture(adapter_only = True) + assert ok is True + assert st["ignore_patterns"] is None + allow = st["allow_patterns"] + assert "adapter_config.json" in allow and "adapter_model*" in allow + kept = _filter(_SAMPLE_FILES, allow, st["ignore_patterns"]) + assert "adapter_config.json" in kept + assert "adapter_model.safetensors" in kept + assert "config.json" in kept and "tokenizer.json" in kept + assert "model-00001-of-00002.safetensors" not in kept + assert "pytorch_model.bin" not in kept + assert "fp16/model.safetensors" not in kept + + +def test_adapter_only_warms_sharded_adapter(capture): + """A sharded adapter is still covered by the adapter_model* glob.""" + _, st = capture(adapter_only = True) + sharded = [ + "adapter_config.json", + "adapter_model-00001-of-00002.safetensors", + "adapter_model-00002-of-00002.safetensors", + "adapter_model.safetensors.index.json", + ] + kept = _filter(sharded, st["allow_patterns"], st["ignore_patterns"]) + assert set(kept) == set(sharded) + + +def test_tokenizer_only_warms_only_aux_files(capture): + """A tokenizer-only repo warms tokenizer/config/vocab files, never weights.""" + _, st = capture(tokenizer_only = True) + assert st["ignore_patterns"] is None + assert st["allow_patterns"] == list(U._ROOT_AUX_PREFETCH_PATTERNS) + kept = _filter(_SAMPLE_FILES, st["allow_patterns"], st["ignore_patterns"]) + assert "tokenizer.json" in kept and "config.json" in kept + assert "model-00001-of-00002.safetensors" not in kept + assert "adapter_model.safetensors" not in kept + + +def test_aux_warm_covers_arbitrary_remote_code_modules(capture): + """The aux warm must cover any *.py, since trust_remote_code auto_map names modules freely.""" + _, st = capture(tokenizer_only = True) + allow = st["allow_patterns"] + assert "*.py" in allow + remote_code = [ + "config.json", + "modeling.py", + "tokenization.py", + "my_custom_code.py", + "configuration_foo.py", + ] + kept = _filter(remote_code, allow, st["ignore_patterns"]) + for name in ("modeling.py", "tokenization.py", "my_custom_code.py", "configuration_foo.py"): + assert name in kept, name + + +def test_subfolder_warms_subfolder_plus_root_aux(capture): + """A subfolder load warms that subfolder's weights plus root aux; other subdirs/root weights skipped.""" + _, st = capture(subfolder = "fp16") + allow = st["allow_patterns"] + assert "fp16/*" in allow + assert all(p in allow for p in U._ROOT_AUX_PREFETCH_PATTERNS) + kept = _filter(_SAMPLE_FILES, allow, st["ignore_patterns"]) + assert "fp16/model.safetensors" in kept + assert "config.json" in kept + assert "experimental/model-00001-of-00002.safetensors" not in kept + + +def test_subfolder_takes_precedence_over_weights_at_root(capture): + """When a subfolder is requested the subfolder branch wins over weights_at_root.""" + _, st = capture(subfolder = "fp16", weights_at_root = True) + assert "fp16/*" in st["allow_patterns"] + kept = _filter(_SAMPLE_FILES, st["allow_patterns"], st["ignore_patterns"]) + assert "fp16/model.safetensors" in kept + + +def test_local_dir_is_not_warmed(capture, tmp_path): + """A local directory path skips the warm (returns False).""" + d = tmp_path / "local-model" + d.mkdir() + ok = U.maybe_prefetch_hf_snapshot(str(d), weights_at_root = True) + assert ok is False + + +def _install_fake_model_info(monkeypatch, filenames): + """Make HfApi().model_info(...).siblings report filenames, with no network.""" + import huggingface_hub + + class _Sib: + def __init__(self, name): + self.rfilename = name + + class _Info: + def __init__(self, names): + self.siblings = [_Sib(n) for n in names] + + class _Api: + def model_info(self, *a, **k): + return _Info(filenames) + + monkeypatch.setattr(huggingface_hub, "HfApi", _Api) + + +# ----- Finding P: variant-aware weight-format selection ----- + + +def test_variant_keeps_bin_when_only_default_safetensors(monkeypatch): + """A default model.safetensors must not prove a variant .bin redundant; without a variant it does.""" + _install_fake_model_info(monkeypatch, ["model.safetensors", "pytorch_model.fp16.bin"]) + ig = U._prefetch_ignore_patterns("org/repo", variant = "fp16", weights_at_root = True) + assert "*.bin" not in ig + ig_default = U._prefetch_ignore_patterns("org/repo", weights_at_root = True) + assert "*.bin" in ig_default + + +def test_variant_drops_bin_when_variant_safetensors_present(monkeypatch): + """A variant-matching safetensors makes the variant .bin redundant, so .bin is dropped.""" + _install_fake_model_info(monkeypatch, ["model.fp16.safetensors", "pytorch_model.fp16.bin"]) + ig = U._prefetch_ignore_patterns("org/repo", variant = "fp16", weights_at_root = True) + assert "*.bin" in ig + + +def test_no_variant_keeps_bin_when_only_variant_safetensors(monkeypatch): + """For a no-variant load, only a canonical safetensors (not a lone variant) makes .bin redundant.""" + _install_fake_model_info(monkeypatch, ["model.fp16.safetensors", "pytorch_model.bin"]) + ig = U._prefetch_ignore_patterns("org/repo", weights_at_root = True) + assert "*.bin" not in ig + _install_fake_model_info(monkeypatch, ["model.safetensors", "pytorch_model.bin"]) + ig2 = U._prefetch_ignore_patterns("org/repo", weights_at_root = True) + assert "*.bin" in ig2 + + +def test_variant_keeps_bin_for_noncanonical_sidecar(monkeypatch): + """A non-canonical variant sidecar must not prove the variant .bin redundant; a canonical one does.""" + _install_fake_model_info( + monkeypatch, ["consolidated.fp16.safetensors", "pytorch_model.fp16.bin"] + ) + ig = U._prefetch_ignore_patterns("org/repo", variant = "fp16", weights_at_root = True) + assert "*.bin" not in ig + _install_fake_model_info(monkeypatch, ["model.fp16.safetensors", "pytorch_model.fp16.bin"]) + ig2 = U._prefetch_ignore_patterns("org/repo", variant = "fp16", weights_at_root = True) + assert "*.bin" in ig2 + + +def test_is_canonical_model_weight_safetensors(): + """The canonical detector matches only non-variant model-weight safetensors names.""" + assert U._is_canonical_model_weight_safetensors("model.safetensors") is True + assert U._is_canonical_model_weight_safetensors("model-00001-of-00002.safetensors") is True + assert U._is_canonical_model_weight_safetensors("model.safetensors.index.json") is True + assert U._is_canonical_model_weight_safetensors("model.fp16.safetensors") is False + assert ( + U._is_canonical_model_weight_safetensors("model.fp16-00001-of-00002.safetensors") is False + ) + assert U._is_canonical_model_weight_safetensors("adapter_model.safetensors") is False + + +def test_st_prefetch_resolves_env_cache_and_runs_after_validation(): + """The ST prefetch must resolve SENTENCE_TRANSFORMERS_HOME and run after load-mode validation.""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "sentence_transformer.py") + with open(src_path, "r", encoding = "utf-8") as f: + src = f.read() + tree = ast.parse(src) + prefetch_calls = [ + n + for n in ast.walk(tree) + if isinstance(n, ast.Call) + and isinstance(n.func, ast.Name) + and n.func.id == "maybe_prefetch_hf_snapshot" + ] + assert len(prefetch_calls) == 1, "expected exactly one ST prefetch call" + call = prefetch_calls[0] + # cache_dir kwarg resolves SENTENCE_TRANSFORMERS_HOME. + cache_dir_kw = next((kw for kw in call.keywords if kw.arg == "cache_dir"), None) + assert cache_dir_kw is not None, "ST prefetch must pass cache_dir" + assert "SENTENCE_TRANSFORMERS_HOME" in ast.dump( + cache_dir_kw.value + ), "ST prefetch cache_dir must resolve SENTENCE_TRANSFORMERS_HOME" + # Load-mode validation runs before the prefetch (fewer source lines = earlier). + val_lineno = src[: src.index("Can only load in 4bit or 8bit or 16bit")].count("\n") + assert val_lineno < call.lineno, "load-mode validation must precede the ST prefetch" + + +def test_st_cache_resolutions_honor_explicit_hf_cache_dir(): + """Every ST cache resolution falling back to SENTENCE_TRANSFORMERS_HOME must first honor an explicit HF cache_dir.""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "sentence_transformer.py") + with open(src_path, "r", encoding = "utf-8") as f: + tree = ast.parse(f.read()) + resolutions = [ + kw + for kw in ast.walk(tree) + if isinstance(kw, ast.keyword) + and kw.arg == "cache_dir" + and "SENTENCE_TRANSFORMERS_HOME" in ast.dump(kw.value) + ] + assert resolutions, "expected cache_dir resolutions referencing SENTENCE_TRANSFORMERS_HOME" + for kw in resolutions: + assert "'cache_dir'" in ast.dump( + kw.value + ), "an ST cache_dir resolution must read an explicit kwargs.get('cache_dir') first" + + +def test_st_native_loads_map_hf_cache_dir_to_cache_folder(): + """Native SentenceTransformer loads take cache_folder, so an explicit HF cache_dir must be mapped onto it.""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "sentence_transformer.py") + with open(src_path, "r", encoding = "utf-8") as f: + src = f.read() + tree = ast.parse(src) + # Every native SentenceTransformer(...) forwarding cache_folder must read cache_dir. + st_calls = [ + n + for n in ast.walk(tree) + if isinstance(n, ast.Call) + and isinstance(n.func, ast.Name) + and n.func.id == "SentenceTransformer" + ] + cache_folder_kws = [kw for call in st_calls for kw in call.keywords if kw.arg == "cache_folder"] + assert cache_folder_kws, "expected a native SentenceTransformer call forwarding cache_folder" + for kw in cache_folder_kws: + assert "'cache_dir'" in ast.dump( + kw.value + ), "a native SentenceTransformer cache_folder must map the explicit HF cache_dir first" + # for_inference feeds cache_folder via st_kwargs; both native branches map cache_dir -> cache_folder. + normalized = "".join(src.split()) + assert ( + 'st_kwargs["cache_folder"]=' in normalized + ), "for_inference must set st_kwargs cache_folder" + assert ( + normalized.count('kwargs.get("cache_dir")orkwargs.get("cache_folder")') >= 2 + ), "both native ST branches (for_inference, fast-encoder) must map cache_dir -> cache_folder" + + +def test_vision_warms_vllm_tokenizer_after_remap(): + """On the vLLM path the tokenizer warm is deferred until after the fast_inference_setup remap.""" + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "vision.py") + with open(src_path, "r", encoding = "utf-8") as f: + src = f.read() + guard = "if _vllm_owns_weights and isinstance(tokenizer_name" + assert guard in src, "expected a vLLM-gated tokenizer warm" + assert src.index(guard) > src.index( + "fast_inference_setup(" + ), "the vLLM tokenizer warm must run after the fast_inference_setup remap" + + +def test_diffusion_forwards_variant_to_real_load(): + """FastDiffusionModel must forward variant to the real model_cls.from_pretrained load, not just the prefetch.""" + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "diffusion.py") + with open(src_path, "r", encoding = "utf-8") as f: + src = f.read() + assert ( + 'load_kwargs["variant"] = kwargs["variant"]' in src + ), "the diffusion load must forward variant to model_cls.from_pretrained" + + +def test_vision_prefetch_runs_after_load_mode_validation(): + """The FastBaseModel (vision) prefetch must run after the load-mode validation.""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "vision.py") + with open(src_path, "r", encoding = "utf-8") as f: + src = f.read() + tree = ast.parse(src) + prefetch_calls = [ + n + for n in ast.walk(tree) + if isinstance(n, ast.Call) + and isinstance(n.func, ast.Name) + and n.func.id == "maybe_prefetch_hf_snapshot" + ] + assert prefetch_calls, "expected a vision prefetch call" + first_prefetch = min(call.lineno for call in prefetch_calls) + val_lineno = src[: src.index("Can only load in 4bit or 8bit or 16bit")].count("\n") + assert val_lineno < first_prefetch, "load-mode validation must precede the vision prefetch" + + +def test_llama_prefetch_skips_only_real_vllm_loads(): + """The llama prefetch's fast_inference skip must be gated on num_labels is None (a classification load still downloads).""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "llama.py") + with open(src_path, "r", encoding = "utf-8") as f: + tree = ast.parse(f.read()) + gated = False + for n in ast.walk(tree): + if not ( + isinstance(n, ast.Call) + and isinstance(n.func, ast.Name) + and n.func.id == "maybe_prefetch_hf_snapshot" + ): + continue + fi_kw = next((kw for kw in n.keywords if kw.arg == "fast_inference"), None) + if fi_kw is None: + continue + dumped = ast.dump(fi_kw.value) + if "fast_inference" in dumped and "num_labels" in dumped: + gated = True + assert gated, "llama prefetch fast_inference must be gated on num_labels is None" + + +def test_st_fallback_module_loads_resolve_env_cache(): + """Fallback module loads deriving cache_dir from cache_folder must also fall back to SENTENCE_TRANSFORMERS_HOME.""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "sentence_transformer.py") + with open(src_path, "r", encoding = "utf-8") as f: + src = f.read() + tree = ast.parse(src) + + # Fallback sites (cache_dir derived from cache_folder) must resolve SENTENCE_TRANSFORMERS_HOME. + checked = 0 + for node in ast.walk(tree): + if not (isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute)): + continue + if node.func.attr not in ("_module_path", "_load_modules"): + continue + cache_dir_kw = next((kw for kw in node.keywords if kw.arg == "cache_dir"), None) + if cache_dir_kw is None: + continue + dumped = ast.dump(cache_dir_kw.value) + if "cache_folder" not in dumped: + continue # internal pass-through, not a resolution site + checked += 1 + assert ( + "SENTENCE_TRANSFORMERS_HOME" in dumped + ), f"{node.func.attr} cache_dir resolves cache_folder but not SENTENCE_TRANSFORMERS_HOME" + assert ( + checked >= 2 + ), "expected the fallback _module_path and _load_modules calls to resolve the env cache" + + +def test_st_fallback_module_loads_forward_revision(): + """The fallback module loads must forward revision so module files match the revision-pinned weights. + Guards: (a) helpers accept revision, (b) every download primitive forwards it, (c) _load_modules + threads it into internal calls, (d) the from_pretrained fallback sites forward it.""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "sentence_transformer.py") + with open(src_path, "r", encoding = "utf-8") as f: + tree = ast.parse(f.read()) + + funcs = { + n.name: n + for n in ast.walk(tree) + if isinstance(n, ast.FunctionDef) + and n.name in ("_module_path", "_read_pooling_mode", "_load_modules") + } + assert set(funcs) == {"_module_path", "_read_pooling_mode", "_load_modules"} + + # (a) each helper takes a revision parameter. + for name, fn in funcs.items(): + arg_names = {a.arg for a in fn.args.args + fn.args.kwonlyargs} + assert "revision" in arg_names, f"{name} must accept a revision argument" + + # (b) every download primitive inside the helpers forwards revision. + downloads = 0 + for name, fn in funcs.items(): + for node in ast.walk(fn): + if not (isinstance(node, ast.Call) and isinstance(node.func, ast.Name)): + continue + if node.func.id not in ("hf_hub_download", "load_dir_path"): + continue + downloads += 1 + assert any( + kw.arg == "revision" for kw in node.keywords + ), f"{node.func.id} in {name} must forward revision" + assert downloads >= 3, "expected the module-download primitives to be revision-guarded" + + # (c) _load_modules threads revision into its internal _module_path / _read_pooling_mode calls. + internal = 0 + for node in ast.walk(funcs["_load_modules"]): + if not (isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute)): + continue + if node.func.attr not in ("_module_path", "_read_pooling_mode"): + continue + internal += 1 + assert any( + kw.arg == "revision" for kw in node.keywords + ), f"_load_modules must forward revision to {node.func.attr}" + assert internal >= 2, "expected _load_modules to call _module_path and _read_pooling_mode" + + # (d) the from_pretrained fallback _module_path / _load_modules sites forward revision. + checked = 0 + for node in ast.walk(tree): + if not (isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute)): + continue + if node.func.attr not in ("_module_path", "_load_modules"): + continue + cache_dir_kw = next((kw for kw in node.keywords if kw.arg == "cache_dir"), None) + if cache_dir_kw is None or "cache_folder" not in ast.dump(cache_dir_kw.value): + continue # internal pass-through, not a fallback site + checked += 1 + rev_kw = next((kw for kw in node.keywords if kw.arg == "revision"), None) + assert rev_kw is not None and "revision" in ast.dump( + rev_kw.value + ), f"{node.func.attr} fallback call must forward revision" + assert ( + checked >= 2 + ), "expected the fallback _module_path and _load_modules calls to forward revision" + + +def test_st_fallback_model_load_resolves_env_cache(): + """from_pretrained must resolve the warmed ST cache into kwargs['cache_dir'] before the FastModel weight load.""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "sentence_transformer.py") + with open(src_path, "r", encoding = "utf-8") as f: + tree = ast.parse(f.read()) + + def _resolves_st_cache(value_node): + # Resolution may be inline or in the assignment to an intermediate variable the value references. + dumped = ast.dump(value_node) + if "cache_folder" in dumped and "SENTENCE_TRANSFORMERS_HOME" in dumped: + return True + if isinstance(value_node, ast.Name): + for n in ast.walk(tree): + if isinstance(n, ast.Assign) and any( + isinstance(t, ast.Name) and t.id == value_node.id for t in n.targets + ): + d = ast.dump(n.value) + if "cache_folder" in d and "SENTENCE_TRANSFORMERS_HOME" in d: + return True + return False + + resolved_lines = [] + for node in ast.walk(tree): + if not isinstance(node, ast.Assign): + continue + for tgt in node.targets: + if ( + isinstance(tgt, ast.Subscript) + and isinstance(tgt.value, ast.Name) + and tgt.value.id == "kwargs" + and isinstance(tgt.slice, ast.Constant) + and tgt.slice.value == "cache_dir" + and _resolves_st_cache(node.value) + ): + resolved_lines.append(node.lineno) + assert resolved_lines, "from_pretrained must resolve the ST cache into kwargs['cache_dir']" + + fastmodel_calls = [ + n.lineno + for n in ast.walk(tree) + if isinstance(n, ast.Call) + and isinstance(n.func, ast.Attribute) + and n.func.attr == "from_pretrained" + and isinstance(n.func.value, ast.Name) + and n.func.value.id == "FastModel" + ] + assert fastmodel_calls, "expected a FastModel.from_pretrained call" + assert min(resolved_lines) < min( + fastmodel_calls + ), "kwargs['cache_dir'] must be resolved before the fallback FastModel weight load" + + +def test_canonical_variant_model_weight_matches_transformers_names(): + """The variant safetensors detector matches only canonical variant names, rejecting sidecars and wrong variants.""" + f = U._is_canonical_variant_model_weight_safetensors + assert f("model.fp16.safetensors", "fp16") is True + assert f("model.fp16-00001-of-00002.safetensors", "fp16") is True + assert f("model-00001-of-00002.fp16.safetensors", "fp16") is True + assert f("model.safetensors.index.fp16.json", "fp16") is True + assert f("consolidated.fp16.safetensors", "fp16") is False + assert f("model.safetensors", "fp16") is False + assert f("model-00001-of-00002.safetensors", "fp16") is False + assert f("model.bf16.safetensors", "fp16") is False + + +def test_variant_is_forwarded_to_downloader(capture): + """maybe_prefetch_hf_snapshot must forward variant to the downloader (absent a variant, nothing is forwarded).""" + _, st = capture(weights_at_root = True, use_safetensors = True, variant = "fp16") + assert st["variant"] == "fp16" + _, st = capture(weights_at_root = True, use_safetensors = True) + assert st["variant"] is None + + +def test_variant_drops_bin_for_sharded_variant_safetensors(monkeypatch): + """A sharded variant safetensors is recognized, so its redundant variant .bin is dropped.""" + _install_fake_model_info( + monkeypatch, + [ + "model.fp16-00001-of-00002.safetensors", + "model.fp16-00002-of-00002.safetensors", + "pytorch_model.fp16-00001-of-00002.bin", + ], + ) + ig = U._prefetch_ignore_patterns("org/repo", variant = "fp16", weights_at_root = True) + assert "*.bin" in ig + + +def test_tokenizer_only_warms_extra_vocab_files(capture): + """tokenizer_only must warm SentencePiece / vocab / processor files, including a named jinja template.""" + _, st = capture(tokenizer_only = True) + allow = st["allow_patterns"] + for name in ( + "spm.model", + "normalizer.json", + "video_preprocessor_config.json", + "tokenizer.model.v3", + ): + assert name in allow, name + sample = [ + "spm.model", + "normalizer.json", + "video_preprocessor_config.json", + "tokenizer.model.v3", + "additional_chat_templates/custom.jinja", + ] + kept = _filter(sample, allow, st["ignore_patterns"]) + assert set(kept) == set(sample) + + +def test_format_probe_runs_even_when_config_cached(capture, monkeypatch): + """A cached config.json must not skip the weight-format probe; model_info still drops the redundant .bin.""" + import huggingface_hub + + # Pretend config.json is cached (the AutoConfig side effect); this must not gate the probe. + monkeypatch.setattr( + huggingface_hub, "try_to_load_from_cache", lambda *a, **k: "/cache/config.json" + ) + _install_fake_model_info(monkeypatch, ["model.safetensors", "pytorch_model.bin"]) + _, st = capture(weights_at_root = True) + ig = st["ignore_patterns"] or [] + assert "*.bin" in ig + + +def test_optimizer_safetensors_does_not_drop_bin(monkeypatch): + """An optimizer.safetensors sidecar must not count as model safetensors, so the real .bin weights are kept.""" + _install_fake_model_info(monkeypatch, ["pytorch_model.bin", "optimizer.safetensors"]) + ig = U._prefetch_ignore_patterns("org/repo", weights_at_root = True) + assert "*.bin" not in ig + + +def test_model_safetensors_still_drops_bin(monkeypatch): + """Control for the optimizer case: a real model.safetensors next to pytorch_model.bin still drops the .bin.""" + _install_fake_model_info( + monkeypatch, ["model.safetensors", "pytorch_model.bin", "optimizer.safetensors"] + ) + ig = U._prefetch_ignore_patterns("org/repo", weights_at_root = True) + assert "*.bin" in ig + + +def test_whole_multi_component_snapshot_keeps_subdir_bin(monkeypatch): + """A whole multi-component snapshot must not drop *.bin (it would strip a subdir module's weight); a root load still does.""" + _install_fake_model_info(monkeypatch, ["model.safetensors", "1_Dense/pytorch_model.bin"]) + ig = U._prefetch_ignore_patterns("org/repo", weights_at_root = False) + assert "*.bin" not in ig + ig_root = U._prefetch_ignore_patterns("org/repo", weights_at_root = True) + assert "*.bin" in ig_root + + +def test_is_model_weight_safetensors_classification(): + """Real model weights count; adapter / trainer-state sidecars do not.""" + assert U._is_model_weight_safetensors("model.safetensors") is True + assert U._is_model_weight_safetensors("model-00001-of-00002.safetensors") is True + assert U._is_model_weight_safetensors("model.safetensors.index.json") is True + assert U._is_model_weight_safetensors("consolidated.safetensors") is True + assert U._is_model_weight_safetensors("adapter_model.safetensors") is False + assert U._is_model_weight_safetensors("optimizer.safetensors") is False + assert U._is_model_weight_safetensors("scheduler.safetensors") is False + assert U._is_model_weight_safetensors("rng_state_0.safetensors") is False + + +def test_tokenizer_only_warms_slow_sentencepiece_vocab(capture): + """tokenizer_only must warm the slow-tokenizer SentencePiece / BPE vocab files AutoTokenizer fetches first.""" + _, st = capture(tokenizer_only = True) + allow = st["allow_patterns"] + for name in ( + "sentencepiece.bpe.model", + "source.spm", + "target.spm", + "bpe.codes", + "vocab.bpe", + "sentencepiece.model", + "vocab-src.json", + "vocab-tgt.json", + ): + assert name in allow, name + + +def test_adapter_safetensors_check_scoped_to_root(monkeypatch): + """_adapter_repo_has_safetensors must only count a root adapter_model*.safetensors, not a subdir one.""" + import huggingface_hub + + class _Sib: + def __init__(self, name): + self.rfilename = name + + class _Api: + def __init__(self, names): + self._names = names + + def model_info(self, *a, **k): + return type("MI", (), {"siblings": [_Sib(n) for n in self._names]})() + + # Subdir safetensors only -> not reported present. + monkeypatch.setattr( + huggingface_hub, + "HfApi", + lambda: _Api( + ["adapter_config.json", "adapter_model.bin", "checkpoint-5/adapter_model.safetensors"] + ), + ) + assert U._adapter_repo_has_safetensors("org/repo") is False + # Root safetensors -> reported present. + monkeypatch.setattr( + huggingface_hub, + "HfApi", + lambda: _Api(["adapter_config.json", "adapter_model.safetensors"]), + ) + assert U._adapter_repo_has_safetensors("org/repo") is True + + +def test_gguf_file_warm_keeps_gguf(capture): + """A gguf_file load allow-lists that GGUF while not pulling other quants the repo publishes.""" + _, st = capture(weights_at_root = True, gguf_file = "model-Q4_K_M.gguf") + allow = st["allow_patterns"] + ig = st["ignore_patterns"] + assert allow is not None and "model-Q4_K_M.gguf" in allow + sample = [ + "model-Q4_K_M.gguf", + "model-Q8_0.gguf", + "config.json", + "tokenizer.json", + ] + kept = _filter(sample, allow, ig) + assert "model-Q4_K_M.gguf" in kept + assert "config.json" in kept + assert "model-Q8_0.gguf" not in kept + + +# ----- Finding Q: adapter weight-format selection ----- + + +def test_adapter_only_prefers_safetensors_over_bin(capture, monkeypatch): + """A mixed-format adapter repo warms only the safetensors PeftModel reads, not both formats.""" + _install_fake_model_info( + monkeypatch, ["adapter_config.json", "adapter_model.safetensors", "adapter_model.bin"] + ) + _, st = capture(adapter_only = True) + ig = st["ignore_patterns"] + assert ig is not None and "adapter_model*.bin" in ig + kept = _filter( + ["adapter_config.json", "adapter_model.safetensors", "adapter_model.bin"], + st["allow_patterns"], + ig, + ) + assert "adapter_model.safetensors" in kept + assert "adapter_model.bin" not in kept + + +def test_adapter_only_bin_only_keeps_bin(capture, monkeypatch): + """A .bin-only adapter repo must keep adapter_model.bin (no safetensors found -> both formats eligible).""" + _install_fake_model_info(monkeypatch, ["adapter_config.json", "adapter_model.bin"]) + _, st = capture(adapter_only = True) + kept = _filter( + ["adapter_config.json", "adapter_model.bin"], st["allow_patterns"], st["ignore_patterns"] + ) + assert "adapter_model.bin" in kept + + +def test_adapter_only_explicit_use_safetensors_false_keeps_bin(capture): + """An explicit use_safetensors=False forces the .bin form without a model_info call.""" + _, st = capture(adapter_only = True, use_safetensors = False) + ig = st["ignore_patterns"] + assert ig is not None and "adapter_model*.safetensors" in ig + kept = _filter( + ["adapter_config.json", "adapter_model.safetensors", "adapter_model.bin"], + st["allow_patterns"], + ig, + ) + assert "adapter_model.bin" in kept + assert "adapter_model.safetensors" not in kept + + +def test_gguf_file_with_subfolder_warms_subfolder_path(capture): + """gguf_file + subfolder: the warm allow-lists /, not the bare root name.""" + _, st = capture(weights_at_root = True, gguf_file = "model-Q4_K_M.gguf", subfolder = "gguf") + allow = st["allow_patterns"] + assert "gguf/model-Q4_K_M.gguf" in allow + kept = _filter(["gguf/model-Q4_K_M.gguf", "config.json"], allow, st["ignore_patterns"]) + assert "gguf/model-Q4_K_M.gguf" in kept and "config.json" in kept + + +def test_from_tf_root_load_ignores_nested_h5(capture): + """A from_tf root load keeps the root .h5 but drops nested .h5 / .msgpack checkpoints.""" + _, st = capture(weights_at_root = True, from_tf = True) + ig = st["ignore_patterns"] + assert "*/*.h5" in ig and "*/*.msgpack" in ig + kept = _filter(["model.h5", "checkpoint-1/model.h5", "config.json"], st["allow_patterns"], ig) + assert "model.h5" in kept + assert "checkpoint-1/model.h5" not in kept + + +def test_sentence_transformer_from_pretrained_is_prefetch_wired(): + """from_pretrained must call maybe_prefetch_hf_snapshot as an unconditional top-level statement before any return.""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "sentence_transformer.py") + with open(src_path, "r", encoding = "utf-8") as f: + tree = ast.parse(f.read()) + cls = next( + n for n in tree.body if isinstance(n, ast.ClassDef) and n.name == "FastSentenceTransformer" + ) + fp = next(n for n in cls.body if isinstance(n, ast.FunctionDef) and n.name == "from_pretrained") + + def _prefetch_call(node): + # a bare call statement, or one whose return is captured (e.g. _st_prefetched = ...) + value = node.value if isinstance(node, (ast.Expr, ast.Assign)) else None + if ( + isinstance(value, ast.Call) + and isinstance(value.func, ast.Name) + and value.func.id == "maybe_prefetch_hf_snapshot" + ): + return value + return None + + prefetch_pos = next((i for i, n in enumerate(fp.body) if _prefetch_call(n)), None) + return_pos = next((i for i, n in enumerate(fp.body) if isinstance(n, ast.Return)), len(fp.body)) + assert ( + prefetch_pos is not None + ), "from_pretrained must call maybe_prefetch_hf_snapshot at top level" + assert prefetch_pos < return_pos, "prefetch must run before any top-level return" + # local_files_only must be forwarded so an offline load does not start a Hub download. + prefetch_call = _prefetch_call(fp.body[prefetch_pos]) + assert "local_files_only" in { + kw.arg for kw in prefetch_call.keywords + }, "prefetch must forward local_files_only" + + +def test_st_module_download_forwards_cache_folder(): + """_load_modules must forward the custom cache_folder into load_dir_path so per-module subdirs read the warmed cache.""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "sentence_transformer.py") + with open(src_path, "r", encoding = "utf-8") as f: + tree = ast.parse(f.read()) + calls = [ + n + for n in ast.walk(tree) + if isinstance(n, ast.Call) and isinstance(n.func, ast.Name) and n.func.id == "load_dir_path" + ] + assert calls, "expected a load_dir_path call in sentence_transformer.py" + assert all( + "cache_folder" in {kw.arg for kw in c.keywords} for c in calls + ), "every load_dir_path call must forward cache_folder" + + +def test_st_native_sentence_transformer_calls_forward_cache_folder(): + """Every native SentenceTransformer(model_name, ...) load must forward cache_folder; a modules-based build needs none.""" + import ast + import os + + src_path = os.path.join(os.path.dirname(U.__file__), "sentence_transformer.py") + with open(src_path, "r", encoding = "utf-8") as f: + tree = ast.parse(f.read()) + weight_loading_calls = [] + for n in ast.walk(tree): + if not ( + isinstance(n, ast.Call) + and isinstance(n.func, ast.Name) + and n.func.id == "SentenceTransformer" + ): + continue + kw_names = {kw.arg for kw in n.keywords} + # A modules-based build downloads nothing; only a repo-name load reads the cache. + if "modules" in kw_names: + continue + weight_loading_calls.append(n) + assert ( + weight_loading_calls + ), "expected a repo-name SentenceTransformer load in sentence_transformer.py" + # cache_folder is forwarded explicitly or via a **kwargs unpacking (kw.arg == None). + for c in weight_loading_calls: + kw_names = {kw.arg for kw in c.keywords} + forwards = "cache_folder" in kw_names or None in kw_names + assert forwards, ( + "a repo-name SentenceTransformer load must forward cache_folder " + f"(explicitly or via **kwargs) at line {c.lineno}" + ) diff --git a/tests/test_synthetic_chunk_data.py b/tests/test_synthetic_chunk_data.py index b9167d214f..abc2c01443 100644 --- a/tests/test_synthetic_chunk_data.py +++ b/tests/test_synthetic_chunk_data.py @@ -104,10 +104,36 @@ def test_chunk_data_rejects_overlap_not_smaller_than_chunk(): os.unlink(path) +def test_chunk_data_uninitialized_error_names_real_class(): + # Without max_seq_length the guard tells the user which method to call first. + # The message must name the real class (SyntheticDataKit) so copying it works; + # a misspelling would raise NameError when the user follows it verbatim. + kit = SyntheticDataKit.__new__(SyntheticDataKit) + kit.tokenizer = _MockTokenizer() # max_seq_length intentionally unset + with tempfile.NamedTemporaryFile("w", suffix = ".txt", delete = False) as f: + f.write("word " * 50) + path = f.name + try: + try: + kit.chunk_data(filename = path) + raise AssertionError("expected RuntimeError when max_seq_length is unset") + except RuntimeError as e: + msg = str(e) + assert ( + "SyntheticDataKit.from_pretrained" in msg + ), f"error must name SyntheticDataKit.from_pretrained, got: {msg}" + assert ( + "SynthetidDataKit" not in msg + ), f"error must not misspell the class name, got: {msg}" + finally: + os.unlink(path) + + if __name__ == "__main__": test_chunk_data_keeps_single_chunk_document() test_chunk_data_still_splits_long_document() test_chunk_data_empty_document_yields_no_chunks() test_chunk_data_short_document_is_not_split_into_fragments() test_chunk_data_rejects_overlap_not_smaller_than_chunk() + test_chunk_data_uninitialized_error_names_real_class() print("OK") diff --git a/unsloth/dataprep/synthetic.py b/unsloth/dataprep/synthetic.py index e009e3caab..0c8f9451c4 100644 --- a/unsloth/dataprep/synthetic.py +++ b/unsloth/dataprep/synthetic.py @@ -410,7 +410,7 @@ class SyntheticDataKit: assert os.path.exists(filename) assert hasattr(self, "tokenizer") if not hasattr(self, "max_seq_length"): - raise RuntimeError("Please use SynthetidDataKit.from_pretrained(...) first!") + raise RuntimeError("Please use SyntheticDataKit.from_pretrained(...) first!") if not hasattr(self, "overlap") or not hasattr(self, "max_generation_tokens"): raise RuntimeError("Please use prepare_qa_generation first!") diff --git a/unsloth/kernels/fp8.py b/unsloth/kernels/fp8.py index ca608fa01b..935ffbb447 100644 --- a/unsloth/kernels/fp8.py +++ b/unsloth/kernels/fp8.py @@ -68,7 +68,9 @@ def weight_dequant_kernel(x_ptr, s_ptr, y_ptr, M, N, BLOCK_SIZE: tl.constexpr): n = tl.cdiv(N, BLOCK_SIZE) offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) offs_n = pid_n * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - offs = offs_m[:, None] * N + offs_n[None, :] + # tl.arange is int32, so offs_m * N overflows for tensors with more than + # 2**31 elements (e.g. flattened MoE expert stacks); index in int64. + offs = offs_m[:, None].to(tl.int64) * N + offs_n[None, :].to(tl.int64) mask = (offs_m[:, None] < M) & (offs_n[None, :] < N) x = tl.load(x_ptr + offs, mask = mask).to(tl.float32) s = tl.load(s_ptr + pid_m * n + pid_n) @@ -327,11 +329,42 @@ fp8_block_matmul = ( ) +def _blockwise_weight_dequant_any_shape(weight, weight_scale, block_size, out_dtype): + """Blockwise fp8 weight dequant for any shape: triton when the weight tiles + evenly into block_size, else a torch-native per-block scale expansion.""" + m, n = weight.shape + if weight_scale.dtype not in (torch.float32, torch.float16, torch.bfloat16): + weight_scale = weight_scale.to(torch.float32) # e.g. float8_e8m0fnu scales break triton + if weight_scale.numel() == 1: + # Per-tensor scale: the normal forward stashes the un-expanded scalar, + # which repeat_interleave cannot grow to (m, n). Scale directly. + return (weight.to(torch.float32) * weight_scale.float()).to(out_dtype) + if m % block_size[0] != 0 or n % block_size[1] != 0 or block_size[0] != block_size[1]: + # Uneven tiling, or rectangular blocks. The triton kernel uses a single + # BLOCK_SIZE for both axes and derives the column scale stride from it, so + # it mis-indexes the scale when block_size[0] != block_size[1]. Expand the + # per-block scales in torch, which handles both dimensions independently. + s_full = weight_scale.repeat_interleave(block_size[0], 0)[:m] + s_full = s_full.repeat_interleave(block_size[1], 1)[:, :n] + return (weight.to(torch.float32) * s_full).to(out_dtype) + # Even tiling with square blocks: block-quant dequant with the real block size + # (weight_dequant would silently default to 128 and dequantize wrongly). + return weight_dequant_block(weight, weight_scale, block_size = block_size[0], dtype = out_dtype) + + class FP8BlockQuantLinear(torch.autograd.Function): @staticmethod def forward(ctx, X, weight, weight_scale): m, n = weight.shape + if weight_scale.dtype not in (torch.float32, torch.float16, torch.bfloat16): + # Upcast (e.g. e8m0) returns a fresh tensor and drops any Python + # attribute, so carry block_size across the cast for the lookup below. + _scale_block_size = getattr(weight_scale, "block_size", None) + weight_scale = weight_scale.to(torch.float32) # e8m0 scales break triton dtype mapping + if _scale_block_size is not None: + weight_scale.block_size = _scale_block_size + # Original scale, saved for backward before any transformation original_weight_scale = weight_scale @@ -360,6 +393,18 @@ class FP8BlockQuantLinear(torch.autograd.Function): if not weight.is_contiguous(): weight = weight.contiguous() + if X.shape[-1] % block_size[1] != 0: + # Hidden dim not divisible by the activation block: dequant + plain matmul. + # Use the original (un-expanded) scale so a scalar per-tensor scale keeps + # the fast scalar path in both forward and backward. + W_deq = _blockwise_weight_dequant_any_shape( + weight, original_weight_scale, block_size, X.dtype + ) + ctx.weight = weight + ctx.weight_scale = original_weight_scale + ctx.block_size = block_size + return torch_matmul(X, W_deq.T).to(X.dtype) + qinput, scale = act_quant(X, block_size[1]) output = fp8_block_matmul( qinput, @@ -371,11 +416,14 @@ class FP8BlockQuantLinear(torch.autograd.Function): ) ctx.weight = weight ctx.weight_scale = original_weight_scale # Save original for backward + ctx.block_size = block_size return output.to(X.dtype) @staticmethod def backward(ctx, grad_output): - W_deq = weight_dequant(ctx.weight, ctx.weight_scale) + W_deq = _blockwise_weight_dequant_any_shape( + ctx.weight, ctx.weight_scale, ctx.block_size, grad_output.dtype + ) grad_X = torch_matmul(grad_output, W_deq) del W_deq return grad_X, None, None diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index d8adf51d5c..a85d460f9a 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -83,8 +83,10 @@ __all__ = [ "verify_fp8_support_if_applicable", "_get_inference_mode_context_manager", "hf_login", + "maybe_prefetch_hf_snapshot", "is_moe_model", "get_moe_target_parameters", + "_select_moe_detection_targets", "make_fast_generate_wrapper", "_mark_unsloth_disable_data_parallel", "_patch_transformers_trainer_data_parallel", @@ -421,6 +423,18 @@ def apply_unsloth_gradient_checkpointing(use_gradient_checkpointing, max_seq_len _FLEX_EXCLUDED_MODELS = ("gpt_oss", "mllama", "nemotron_h", "modernbert") _FLEX_PREFERRED_MODELS = ("gemma3", "gemma3_text", "shieldgemma2") _SDPA_EXCLUDED_MODELS = ("gpt_oss",) +# The loader (loader.py) forces supports_sdpa=False for these because their bundled +# SDPA modules are wrong. Kept here, not in loader.py, so _is_sdpa_excluded can honor +# them without a loader -> _utils import cycle (loader.py already imports from _utils +# and re-exports this name for callers like sentence_transformer.py). Entries are matched +# as substrings against a comma-joined model_types string ending in a comma, so "gemma3," +# matches a distinct "gemma3" entry but not "gemma3n", and "gemma3_text" matches the +# EmbeddingGemma text model. +DISABLE_SDPA_MODEL_NAMES = [ + "gemma3,", # Add comma bc gemma3 will match gemma3n + "gemma3_text", # Gemma3TextModel (EmbeddingGemma) - substring match, keep underscore + "gpt_oss", +] _FLASH_EXCLUDED_MODELS = ("gpt_oss",) _EAGER_ONLY_PREFIXES = ("gemma3n",) _FLASH_ATTENTION_MAX_HEAD_DIM = 256 @@ -431,8 +445,23 @@ def _is_flex_excluded(model_type): return model_type in _FLEX_EXCLUDED_MODELS +def _is_sdpa_disabled_by_name(model_type): + # Mirror the loader's DISABLE_SDPA_MODEL_NAMES check: loader.py builds + # model_types_all = ",".join(model_types) + "," and tests `name in model_types_all`. + # Rebuild the same trailing-comma form for a single model_type so the match is + # identical (e.g. "gemma3," matches "gemma3" but not "gemma3n", and "gemma3_text" + # still matches "gemma3_text"). + model_types_all = model_type.lower() + "," + return any(name.lower() in model_types_all for name in DISABLE_SDPA_MODEL_NAMES) + + def _is_sdpa_excluded(model_type): - return model_type in _SDPA_EXCLUDED_MODELS + # SDPA is known-broken for these models, so an explicit sdpa request must not + # re-enable it. Two sources: _SDPA_EXCLUDED_MODELS (resolver-level, e.g. gpt_oss) + # and DISABLE_SDPA_MODEL_NAMES (loader-level, e.g. gemma3 / gemma3_text, which the + # loader also forces to supports_sdpa=False). + lowered = model_type.lower() + return lowered in _SDPA_EXCLUDED_MODELS or _is_sdpa_disabled_by_name(lowered) def _is_flash_excluded(model_type): @@ -608,6 +637,12 @@ def _disable_flash_attention_if_needed( if disable_reason is None: return attn_implementation + # Only an implementation passed by the caller counts as an explicit request. + # Values read from the config are synthesized by the loaders (the language path + # seeds the config with attn_implementation="sdpa") or come from Transformers + # defaults, so they must not be treated as a deliberate user choice. + explicit_request = attn_implementation + requested_attn_implementation = attn_implementation if requested_attn_implementation is None: requested_attn_implementation = _config_get(config, "_attn_implementation", None) @@ -617,6 +652,20 @@ def _disable_flash_attention_if_needed( if requested_attn_implementation == "eager": return _set_attn_impl(config, "eager") + model_type = _config_get(config, "model_type", "") + + # The disable reason is flash-specific: honor an explicit non-flash request from + # the caller instead of downgrading it. SDPA is honored unless the model's SDPA is + # known-broken - _SDPA_EXCLUDED_MODELS (e.g. gpt_oss) or DISABLE_SDPA_MODEL_NAMES + # (e.g. gemma3 / gemma3_text); flex_attention + # is honored only when it is actually usable, since supports_flex_attention already + # rejects the excluded/broken/unavailable configs. This keeps an explicit request + # from selecting a backend the repo marks as wrong. + if explicit_request == "sdpa" and not _is_sdpa_excluded(model_type.lower()): + return _set_attn_impl(config, "sdpa") + if explicit_request == "flex_attention" and supports_flex_attention: + return _set_attn_impl(config, "flex_attention") + if supports_sdpa: fallback_attn_implementation = "sdpa" elif supports_flex_attention: @@ -629,7 +678,6 @@ def _disable_flash_attention_if_needed( if _is_flash_attention_requested(requested_attn_implementation) else "flash_attention_2" ) - model_type = _config_get(config, "model_type", "") warning_key = ( model_type, logged_attn_implementation, @@ -843,7 +891,19 @@ def resolve_attention_implementation( final_attn_impl = requested_attn_implementation _set_attn_impl(config, final_attn_impl) - if not supports_sdpa and final_attn_impl == "sdpa": + # A caller who explicitly passes requested_attn_implementation="sdpa" keeps it even + # on a conservatively unsupported model, mirroring _disable_flash_attention_if_needed + # which honors an explicit sdpa request. The exception is a model whose SDPA is + # known-broken - _SDPA_EXCLUDED_MODELS (e.g. gpt_oss) or DISABLE_SDPA_MODEL_NAMES + # (e.g. gemma3 / gemma3_text, which the loader also forces to supports_sdpa=False): + # an explicit request must not re-enable it, so it still downgrades to eager, just + # like flex falls back for _FLEX_EXCLUDED_MODELS. A synthesized/default sdpa + # (requested is None, so the value came from the model resolution above or the + # config) also downgrades. + honor_explicit_sdpa = requested_attn_implementation == "sdpa" and not _is_sdpa_excluded( + model_type + ) + if not supports_sdpa and final_attn_impl == "sdpa" and not honor_explicit_sdpa: print( f"Unsloth: {(model_type_name or 'model').title()} does not support SDPA - switching to fast eager." ) @@ -905,6 +965,411 @@ logging.getLogger("transformers.tokenization_utils_base").setLevel(logging.CRITI TORCHAO_MSG = "Error: torchao not found, please install with `pip install torchao`" +# Artifacts a Transformers/PEFT load never reads (ONNX/TF/Flax/CoreML/GGUF/training state), skipped +# when prewarming so a mixed-format repo is not pulled in full. +_PREFETCH_IGNORE_PATTERNS = ( + "*.onnx", + "onnx/*", + "*.h5", + "*.msgpack", + "*.tflite", + "coreml/*", + "*.mlpackage/*", + "*.mlmodel", + "*.gguf", + # Training / checkpoint formats from_pretrained never reads. + "*.pt", + "*.pth", + "*.ckpt", + "optimizer.*", + "scheduler.*", + "rng_state*", + "trainer_state.json", + "events.out.tfevents*", + "checkpoint-*/*", +) + + +# Repo-root tokenizer / config / processor files from_pretrained reads from root even when weights +# load from a subfolder. Exact names (no wildcard) so they match only root-level files. +_ROOT_AUX_PREFETCH_PATTERNS = ( + "config.json", + "generation_config.json", + "tokenizer_config.json", + "tokenizer.json", + "tokenizer.model", + "special_tokens_map.json", + "added_tokens.json", + "vocab.json", + "vocab.txt", + "merges.txt", + "spiece.model", + # More VOCAB_FILES_NAMES the slow tokenizer may fetch (DeBERTa-v2, Whisper, Mistral, XLM-R/mBART, Marian, FSMT/XLM, GPT-2). + "spm.model", + "normalizer.json", + "tokenizer.model.v3", + "sentencepiece.bpe.model", + "source.spm", + "target.spm", + "bpe.codes", + "vocab.bpe", + # More VOCAB_FILES_NAMES (RemBERT, FSMT) a distinct-tokenizer-repo warm must cache too. + "sentencepiece.model", + "vocab-src.json", + "vocab-tgt.json", + "chat_template.jinja", + "chat_template.json", + # chat_template="" fetches additional_chat_templates/.jinja. + "additional_chat_templates/*.jinja", + "preprocessor_config.json", + "processor_config.json", + "video_preprocessor_config.json", # Qwen2.5-VL-style video processors + # trust_remote_code auto_map can name any module, so warm every *.py (tiny; none in a non-remote repo). + "*.py", + "*.tiktoken", # tiktoken vocab (e.g. Qwen's qwen.tiktoken) +) + + +# Files a PEFT adapter load reads: config + weights (glob covers sharded adapters). Any merged +# full-model weights the repo also ships match none of these. +_ADAPTER_PREFETCH_PATTERNS = ( + "adapter_config.json", + "adapter_model*", +) + + +# Weight files in a SUBDIRECTORY. A bare root load reads only root weights, so ignoring these drops +# alternate-precision/experimental dirs (fp16/, experimental/). "*/*" spans "/" (HF fnmatch), so nested +# weights match while root "model.safetensors" is kept. Only applied when weights_at_root (diffusion +# keeps weights in subfolders). +_SUBDIR_WEIGHT_IGNORE_PATTERNS = ( + "*/*.safetensors", + "*/*.bin", + "*/*.h5", + "*/*.msgpack", + "*/*.pt", + "*/*.pth", +) + + +def _in_requested_load_scope(filename, subfolder): + """True if *filename* is in the location being loaded (*subfolder*, else root). Scopes the ".bin is + redundant when safetensors exist" test so a .bin-only subfolder keeps its .bin.""" + filename = filename.replace("\\", "/") + if isinstance(subfolder, str) and subfolder.strip("/"): + return filename.startswith(subfolder.strip("/") + "/") + return "/" not in filename # root load: no directory component + + +# .safetensors training-state files that are NOT model weights (e.g. optimizer.safetensors next to a +# real pytorch_model.bin); counting them as "model safetensors present" would drop the needed .bin. +_NON_MODEL_WEIGHT_STEMS = frozenset( + { + "optimizer", + "scheduler", + "scaler", + "rng_state", + "training_args", + } +) + + +def _is_model_weight_safetensors(filename): + """True if *filename* is a model-weights safetensors, not a PEFT adapter/sidecar + (adapter_model.safetensors) or trainer-state (optimizer.safetensors). Only a real one proves the + .bin redundant; counting a sidecar would wrongly drop the needed .bin (fetched then without Xet fallback).""" + name = filename.replace("\\", "/").rsplit("/", 1)[-1] + if not name.endswith((".safetensors", ".safetensors.index.json")): + return False + if name.startswith("adapter_"): + return False + # Stem before first dot: "optimizer.safetensors" -> "optimizer" (real shards kept); rng_state via prefix. + stem = name.split(".", 1)[0].lower() + if stem in _NON_MODEL_WEIGHT_STEMS or stem.startswith("rng_state"): + return False + return True + + +def _is_canonical_variant_model_weight_safetensors(filename, variant): + """True for a canonical model-weights safetensors carrying the requested *variant*, in the forms + transformers reads (single, either numbered-shard layout, or the index). Strict (base must be + "model"): a sidecar like consolidated..safetensors does not prove the variant .bin redundant.""" + base = filename.replace("\\", "/").rsplit("/", 1)[-1] + v = re.escape(variant) + return bool( + re.match( + rf"^(?:model\.{v}\.safetensors" + rf"|model\.{v}-\d{{5}}-of-\d{{5}}\.safetensors" + rf"|model-\d{{5}}-of-\d{{5}}\.{v}\.safetensors" + rf"|model\.safetensors\.index\.{v}\.json)$", + base, + ) + ) + + +_CANONICAL_MODEL_WEIGHT_SAFETENSORS_RE = re.compile( + r"^(?:model\.safetensors|model-\d{5}-of-\d{5}\.safetensors|model\.safetensors\.index\.json)$" +) + + +def _is_canonical_model_weight_safetensors(filename): + """True for a canonical (non-variant) model-weights safetensors a default load reads (model.safetensors, + a numbered shard, or the index). Strict: an unrecognized name keeps both formats, so a variant-only + safetensors + pytorch_model.bin repo never has its .bin dropped for a no-variant load.""" + name = filename.replace("\\", "/").rsplit("/", 1)[-1] + return bool(_CANONICAL_MODEL_WEIGHT_SAFETENSORS_RE.match(name)) + + +def _adapter_repo_has_safetensors( + model_name, + *, + token = None, + revision = None, +): + """Best-effort: does the adapter repo ship a root safetensors adapter weight (making the .bin + redundant)? Scoped to root adapter_model* files; any failure returns False.""" + try: + from huggingface_hub import HfApi + siblings = HfApi().model_info(model_name, revision = revision, token = token).siblings or [] + return any( + "/" not in sibling.rfilename.replace("\\", "/") # root only + and sibling.rfilename.startswith("adapter_model") + and sibling.rfilename.endswith(".safetensors") + for sibling in siblings + ) + except Exception: + return False + + +def _prefetch_ignore_patterns( + model_name, + *, + token = None, + revision = None, + subfolder = None, + use_safetensors = None, + from_tf = False, + from_flax = False, + variant = None, + weights_at_root = False, +): + """ignore_patterns for the prewarm snapshot: the static skip list, minus the checkpoint guard when + loading from a checkpoint-* subfolder, minus the weight format the load will not read. use_safetensors + is a format allowlist (True -> skip *.bin, False -> skip *.safetensors); auto (None) skips *.bin only + when in-scope safetensors are shipped. from_tf/from_flax keep *.h5/*.msgpack. + + Suppressed for a whole multi-component snapshot (weights_at_root=False, no subfolder: ST/diffusers + repos with per-subfolder weights, each in its own format), since "*" spans "/" so dropping "*.bin" + would strip a module's only weight.""" + # Keep checkpoint-*/* under a checkpoint-* subfolder; keep *.h5 / *.msgpack under from_tf/flax. + ignore_patterns = [ + pattern + for pattern in _PREFETCH_IGNORE_PATTERNS + if not ( + ( + pattern == "checkpoint-*/*" + and isinstance(subfolder, str) + and subfolder.startswith("checkpoint-") + ) + or (from_tf and pattern == "*.h5") + or (from_flax and pattern == "*.msgpack") + ) + ] + # Drop the format the load will not read (the other doubles the download); skipped for a whole + # multi-component snapshot (see docstring). + whole_multi_component = not weights_at_root and not ( + isinstance(subfolder, str) and subfolder.strip("/") + ) + if whole_multi_component: + pass + elif from_tf or from_flax: + # TF / Flax loads never read the PyTorch formats; drop safetensors and .bin. + ignore_patterns.extend( + ( + "*.safetensors", + "*.safetensors.index.json", + "*.bin", + "*.bin.index.json", + ) + ) + elif use_safetensors is True: + # Explicit safetensors: load never reads .bin (no model_info call needed). + ignore_patterns.extend(("*.bin", "*.bin.index.json")) + elif use_safetensors is False: + # Explicit .bin: load never reads safetensors. + ignore_patterns.extend(("*.safetensors", "*.safetensors.index.json")) + else: + # Auto: skip .bin only once in-scope safetensors are confirmed (best-effort; any failure keeps both). + try: + from huggingface_hub import HfApi + + siblings = ( + HfApi() + .model_info( + model_name, + revision = revision, + token = token, + ) + .siblings + or [] + ) + # Count only in-scope model-weights safetensors (not adapters/sidecars): variant-matching if + # a variant is requested, else canonical, proving the .bin redundant. + has_safetensors = any( + _is_model_weight_safetensors(sibling.rfilename) + and _in_requested_load_scope(sibling.rfilename, subfolder) + and ( + _is_canonical_variant_model_weight_safetensors(sibling.rfilename, variant) + if variant + else _is_canonical_model_weight_safetensors(sibling.rfilename) + ) + for sibling in siblings + ) + if has_safetensors: + ignore_patterns.extend(("*.bin", "*.bin.index.json")) + except Exception: + pass + return ignore_patterns + + +def maybe_prefetch_hf_snapshot( + model_name, + token = None, + *, + revision = None, + cache_dir = None, + local_files_only = False, + fast_inference = False, + subfolder = None, + force_download = False, + use_safetensors = None, + from_tf = False, + from_flax = False, + tokenizer_only = False, + adapter_only = False, + weights_at_root = False, + variant = None, + gguf_file = None, +): + """Warm the HF cache for a remote repo before the in-process load. + + Xet can hang on a blob with no progress or exception, and a blocked native Xet thread cannot be + killed in-process. So pull the snapshot first in a killable subprocess that falls back Xet -> HTTP + on a stall (unsloth_zoo.hf_xet_fallback), making from_pretrained a cache hit. + + Returns True iff warmed (caller can clear force_download), else False (skipped: local/offline/ + local_files_only/fast_inference/old unsloth_zoo, or failed). Only a both-transports-stalled + DownloadStallError is raised; other failures are left for from_pretrained to surface. + """ + try: + from unsloth_zoo.hf_xet_fallback import ( + snapshot_download_with_xet_fallback, + DownloadStallError, + ) + except Exception: + return False # older unsloth_zoo without the helper: load normally + + if not isinstance(model_name, str) or not model_name: + return False + # Local path: nothing to download. Expand ~ first (os.path.exists does not). + model_path = os.path.expanduser(model_name) + if os.path.isdir(model_path) or os.path.exists(model_path): + return False + # Looks local but not yet on disk (e.g. an uncreated output dir): not a Hub repo id, so leave it + # for from_pretrained rather than download it. + if ( + os.path.isabs(model_path) + or model_name.startswith(("~", "./", "../", ".\\", "..\\")) + or "\\" in model_name + ): + return False + if local_files_only: # cache-only: never reach out + return False + if any( + os.environ.get(flag, "0").lower() in ("1", "true", "yes", "on") + for flag in ("HF_HUB_OFFLINE", "TRANSFORMERS_OFFLINE") + ): + return False + if fast_inference: # vLLM has its own download path + return False + + # tokenizer-only / adapter-only warms allow-list exact files below, so the weight-format ignore + # list (and its auto-branch model_info call) is skipped. + ignore_patterns = ( + None + if tokenizer_only or adapter_only or gguf_file + else _prefetch_ignore_patterns( + model_name, + token = token, + revision = revision, + subfolder = subfolder, + use_safetensors = use_safetensors, + from_tf = from_tf, + from_flax = from_flax, + variant = variant, + weights_at_root = weights_at_root, + ) + ) + # Narrow the warm to what the load reads (skip extra checkpoints/precisions); every branch still warms + # root tokenizer/config/custom-code so those never fall in-process. + allow_patterns = None + if gguf_file: + # gguf_file=NAME reads exactly that GGUF, but the static ignore list drops *.gguf; so warm just + # that file (plus root aux), under / if set. + _gguf_path = ( + f"{subfolder.strip('/')}/{gguf_file}" + if isinstance(subfolder, str) and subfolder.strip("/") + else gguf_file + ) + allow_patterns = [_gguf_path, *_ROOT_AUX_PREFETCH_PATTERNS] + elif tokenizer_only: + # A distinct tokenizer repo: warm only tokenizer / config / vocab files, never its weights. + allow_patterns = list(_ROOT_AUX_PREFETCH_PATTERNS) + elif adapter_only: + # A PEFT adapter load reads only adapter_config.json + adapter_model.* (plus root aux), not any + # merged weights the repo may also publish. + allow_patterns = [*_ADAPTER_PREFETCH_PATTERNS, *_ROOT_AUX_PREFETCH_PATTERNS] + # PeftModel reads one format (safetensors when present): explicit use_safetensors wins, else + # prefer safetensors when shipped (best-effort; any failure keeps both). + if use_safetensors is False: + ignore_patterns = [ + "adapter_model*.safetensors", + "adapter_model*.safetensors.index.json", + ] + elif use_safetensors is True or _adapter_repo_has_safetensors( + model_name, token = token, revision = revision + ): + ignore_patterns = ["adapter_model*.bin", "adapter_model*.bin.index.json"] + elif isinstance(subfolder, str) and subfolder.strip("/"): + # subfolder=X: load resolves every weight under X/, so warm that subfolder (plus root aux). + allow_patterns = [f"{subfolder.strip('/')}/*", *_ROOT_AUX_PREFETCH_PATTERNS] + elif weights_at_root: + # A bare load reads only root weights: drop subdir weights (fp16/, checkpoint dirs) while keeping + # subdir configs. Diffusion leaves weights_at_root False. + ignore_patterns = [*(ignore_patterns or []), *_SUBDIR_WEIGHT_IGNORE_PATTERNS] + try: + snapshot_download_with_xet_fallback( + model_name, + token = token, + revision = revision, + cache_dir = cache_dir, + allow_patterns = allow_patterns, + ignore_patterns = ignore_patterns, + force_download = force_download, + variant = variant, + ) + return True + except DownloadStallError: + # Both transports stalled: surface a clear network error, not a silent in-process hang. + raise + except Exception as exception: + logger.warning_once( + f"Unsloth: Could not pre-download {model_name} " + f"({type(exception).__name__}: {exception}); continuing with the normal load." + ) + return False + + # Ignore logging messages class HideLoggingMessage(logging.Filter): __slots__ = ("text",) @@ -3511,8 +3976,25 @@ def _moe_target_set_from_string(target_modules: str) -> set[str]: return {target_modules} is_regex = re.search(r"[*+?()[\]{}|\\^$]", target_modules) is not None - targets_mlp = "mlp" in target_modules or "ffn" in target_modules - if is_regex and "proj" in target_modules and targets_mlp: + # Key detection on the mlp/ffn/experts path segment (absent from an + # attention-only regex), never on q/k/v/o leaves alone. + targets_mlp_path = any( + tag in target_modules for tag in ("mlp", "ffn", "feed_forward", "experts") + ) + if not is_regex or not targets_mlp_path: + return set() + # Explicit expert leaves scope the target set to exactly those leaves. + named = {name for name in _MOE_BROAD_MLP_TARGETS if name in target_modules} + if named: + return named + # A generic projection under an mlp path (e.g. ".*mlp.*proj"): any proj + # occurrence that is not an attention leaf name. + if re.search(r"(? Optional[List[str return None +def _select_moe_detection_targets( + original_target_modules, + scoped_target_modules, + finetune_mlp_modules = True, + finetune_language_layers = True, +): + """Pick what get_moe_target_parameters keys expert detection on. + + Prefer the caller's ORIGINAL explicit leaf list over the scoped regex so an + attention-only request is not pushed into the experts by get_peft_regex's + ``mlp|feed_forward|ffn|dense`` component block (which the string fallback + cannot tell apart from a fused-expert auto regex). + + But only when the MLP and language families are BOTH still in scope. If the + caller scoped MLP or language OFF (``finetune_mlp_modules=False`` or + ``finetune_language_layers=False``) the scoped regex already drops the MoE + experts, and reusing the original list -- which may still name gate/up/down + leaves -- would wrongly re-introduce them. In that case honor the scoped + result so the frozen-MLP / vision-only request is respected. + """ + if original_target_modules is not None and finetune_mlp_modules and finetune_language_layers: + return original_target_modules + return scoped_target_modules + + def make_fast_generate_wrapper(original_generate): """ Creates a wrapper around model.generate that checks for incorrect diff --git a/unsloth/models/cohere.py b/unsloth/models/cohere.py index 0b7f3ab973..cb367d451e 100644 --- a/unsloth/models/cohere.py +++ b/unsloth/models/cohere.py @@ -22,6 +22,7 @@ from ..utils.attention_dispatch import ( AttentionContext, run_attention, select_attention_backend, + resolve_prefix_seg_info, ) try: @@ -151,6 +152,9 @@ def CohereAttention_fast_forward( "softmax_scale": getattr(self, "softmax_scale", None), }, ) + # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward; misuse + # (KV cache / padding mask) raises. None => byte-identical default. + _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask) context = AttentionContext( bsz = bsz, q_len = q_len, @@ -161,6 +165,7 @@ def CohereAttention_fast_forward( seq_info = seq_info, attention_mask = attention_mask, causal_mask = causal_mask, + prefix_seg_info = _pg_seg, ) A = run_attention(config = attention_config, context = context, Q = Q, K = K, V = V) diff --git a/unsloth/models/diffusion.py b/unsloth/models/diffusion.py index 12596b432e..955bf55987 100644 --- a/unsloth/models/diffusion.py +++ b/unsloth/models/diffusion.py @@ -24,7 +24,7 @@ import os import torch from transformers import AutoConfig, AutoProcessor, AutoTokenizer -from ._utils import is_bfloat16_supported +from ._utils import is_bfloat16_supported, maybe_prefetch_hf_snapshot from .llama import logger __all__ = ["FastDiffusionModel", "DIFFUSION_MODEL_TYPES", "is_diffusion_model_type"] @@ -79,7 +79,14 @@ def _resolve_diffusion_model_class(config): ) -def _load_diffusion_config(model_name, token, trust_remote_code, revision, local_files_only): +def _load_diffusion_config( + model_name, + token, + trust_remote_code, + revision, + local_files_only, + cache_dir = None, +): """Load the config, aliasing the legacy ``diffusion_gemma`` model_type to the ``diffusion_gemma4`` classes current transformers ships. AutoConfig raises on the legacy type; catch that, rewrite the type/arch names in-memory, and rebuild.""" @@ -90,6 +97,7 @@ def _load_diffusion_config(model_name, token, trust_remote_code, revision, local trust_remote_code = trust_remote_code, revision = revision, local_files_only = local_files_only, + cache_dir = cache_dir, ) except ValueError as e: if "diffusion_gemma" not in str(e): @@ -103,6 +111,7 @@ def _load_diffusion_config(model_name, token, trust_remote_code, revision, local token = token, revision = revision, local_files_only = local_files_only, + cache_dir = cache_dir, ) with open(cfg_path, encoding = "utf-8") as f: cd = json.load(f) @@ -152,12 +161,16 @@ class FastDiffusionModel: os.environ.get("HF_HUB_OFFLINE", "0") == "1" or os.environ.get("TRANSFORMERS_OFFLINE", "0") == "1" ) + + cache_dir = kwargs.get("cache_dir") + config = _load_diffusion_config( model_name, token, trust_remote_code, revision, local_files_only, + cache_dir = cache_dir, ) model_type = getattr(config, "model_type", None) if not is_diffusion_model_type(model_type): @@ -168,6 +181,21 @@ class FastDiffusionModel: model_cls = _resolve_diffusion_model_class(config) + # Prefetch the whole repo root so the weight load is a cache hit. No subfolder: the pipeline + # loads every component subfolder, so narrowing would leave unet/vae/text_encoder to Xet. + maybe_prefetch_hf_snapshot( + model_name, + token = token, + revision = revision, + cache_dir = cache_dir, + local_files_only = local_files_only, + fast_inference = False, + force_download = kwargs.get("force_download", False), + use_safetensors = kwargs.get("use_safetensors"), + # Forward variant (e.g. "fp16") so the warm keeps variant weights. + variant = kwargs.get("variant"), + ) + load_kwargs = dict( dtype = dtype, device_map = device_map, @@ -176,7 +204,14 @@ class FastDiffusionModel: attn_implementation = attn_implementation, revision = revision, local_files_only = local_files_only, + cache_dir = cache_dir, ) + # Match the load's weight format to the warm (None/auto already matches). + if kwargs.get("use_safetensors") is not None: + load_kwargs["use_safetensors"] = kwargs["use_safetensors"] + # Forward variant to the real load so it reads the warmed variant weights. + if kwargs.get("variant") is not None: + load_kwargs["variant"] = kwargs["variant"] # Optional bitsandbytes quant. The MoE experts (3D Parameters) are not nn.Linear so bnb skips # them; only attention + dense MLP Linears quantize, lm_head/embeddings stay full precision. @@ -222,6 +257,7 @@ class FastDiffusionModel: trust_remote_code = trust_remote_code, revision = revision, local_files_only = local_files_only, + cache_dir = cache_dir, ) except Exception: tokenizer = AutoTokenizer.from_pretrained( @@ -230,6 +266,7 @@ class FastDiffusionModel: trust_remote_code = trust_remote_code, revision = revision, local_files_only = local_files_only, + cache_dir = cache_dir, ) return model, tokenizer diff --git a/unsloth/models/gemma2.py b/unsloth/models/gemma2.py index 68b9ebe22f..4a0531db78 100644 --- a/unsloth/models/gemma2.py +++ b/unsloth/models/gemma2.py @@ -22,6 +22,7 @@ from ..utils.attention_dispatch import ( AttentionContext, run_attention, select_attention_backend, + resolve_prefix_seg_info, SDPA, ) from .gemma import ( @@ -168,6 +169,11 @@ def Gemma2Attention_fast_forward( }, ) + # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward; misuse + # (KV cache / padding mask) raises. None => byte-identical default. gemma2 is + # sliding-window and softcapped: the engage gate caps spans at the window and + # excludes softcap models entirely, so PG never engages here. + _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask) context = AttentionContext( bsz = bsz, q_len = q_len, @@ -179,6 +185,7 @@ def Gemma2Attention_fast_forward( attention_mask = attention_mask, causal_mask = causal_mask, sliding_window = sliding_window, + prefix_seg_info = _pg_seg, ) A = run_attention(config = attention_config, context = context, Q = Q, K = K, V = V) diff --git a/unsloth/models/granite.py b/unsloth/models/granite.py index f5b0f57aa6..4dedf642eb 100644 --- a/unsloth/models/granite.py +++ b/unsloth/models/granite.py @@ -23,6 +23,7 @@ from ..utils.attention_dispatch import ( AttentionContext, run_attention, select_attention_backend, + resolve_prefix_seg_info, SDPA, ) from .llama import ( @@ -159,6 +160,9 @@ def GraniteAttention_fast_forward( }, ) + # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward; misuse + # (KV cache / padding mask) raises. None => byte-identical default. + _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask) context = AttentionContext( bsz = bsz, q_len = q_len, @@ -169,6 +173,7 @@ def GraniteAttention_fast_forward( seq_info = seq_info, attention_mask = attention_mask, causal_mask = causal_mask, + prefix_seg_info = _pg_seg, ) A = run_attention(config = attention_config, context = context, Q = Q, K = K, V = V) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index bb7289dfa8..564be09578 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -39,6 +39,7 @@ from ..utils.attention_dispatch import ( run_attention, SDPA, select_attention_backend, + resolve_prefix_seg_info, ) from torch.nn.functional import scaled_dot_product_attention from transformers import __version__ as transformers_version @@ -738,6 +739,10 @@ def LlamaAttention_fast_forward( flash_dense_kwargs = {"causal": True}, flash_varlen_kwargs = {"dropout_p": 0.0, "causal": True}, ) + # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward (same route + # as packed_seq_lengths); misuse (KV cache / padding mask) raises. None => byte-identical + # default. Reuse of this forward also carries the branch to qwen2 & gemma. + _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask) context = AttentionContext( bsz = bsz, q_len = q_len, @@ -748,6 +753,7 @@ def LlamaAttention_fast_forward( seq_info = seq_info, attention_mask = attention_mask, causal_mask = causal_mask, + prefix_seg_info = _pg_seg, ) A = run_attention(config = config, context = context, Q = Q, K = K, V = V) @@ -895,8 +901,10 @@ def LlamaModel_fast_forward( seq_length_with_past = seq_length # Fix out of bounds tokenization unless we were given packed metadata - allow_overlength = getattr(self, "_unsloth_allow_packed_overlength", False) or ( - "packed_seq_lengths" in kwargs + allow_overlength = ( + getattr(self, "_unsloth_allow_packed_overlength", False) + or ("packed_seq_lengths" in kwargs) + or ("prefix_seg_info" in kwargs and kwargs["prefix_seg_info"] is not None) ) if hasattr(self, "max_seq_length") and not allow_overlength: if seq_length > self.max_seq_length: @@ -2420,6 +2428,73 @@ class FastLlamaModel: preferred_attn_impl = resolve_attention_implementation(model_function, model_config) + # Prefetch the repo (killable child) so the weight load is a cache hit. Runs after the + # AutoConfig/model-class check so an unsupported repo fails on its small config fetch. No + # revision: the load resolves model_name (maybe a remapped prequant repo) on its default branch. + _prefetched = maybe_prefetch_hf_snapshot( + model_name, + token = token, + cache_dir = kwargs.get("cache_dir"), + local_files_only = kwargs.get("local_files_only", False), + # Skip the warm only for a real vLLM load; a num_labels classification load still goes + # in-process below, so it must be warmed even under fast_inference. + fast_inference = fast_inference and num_labels is None, + subfolder = kwargs.get("subfolder"), + force_download = kwargs.get("force_download", False), + use_safetensors = kwargs.get("use_safetensors"), + from_tf = kwargs.get("from_tf", False), + from_flax = kwargs.get("from_flax", False), + # Bare load reads only ROOT weights; skip subdir weights. Ignored when a subfolder is set. + weights_at_root = True, + variant = kwargs.get("variant"), # forward so the warm keeps the variant .bin + gguf_file = kwargs.get( + "gguf_file" + ), # forward so the warm fetches the GGUF (else ignored) + ) + # Child did the forced download; clear the flag so the load reuses the warm cache. + if _prefetched and kwargs.get("force_download", False): + kwargs["force_download"] = False + + # Tokenizer always loads in-process. Resolve the cache_dir the tokenizer load will actually + # use, mirroring load_correct_tokenizer: without an explicit cache_dir, Colab/Kaggle route to + # a special tokenizer cache (huggingface_tokenizers_cache / Kaggle tmp), NOT the HF-default + # cache the base snapshot warmed. So the base warm does not cover the tokenizer there. + from ..tokenizer_utils import ( + IS_COLAB_ENVIRONMENT, + IS_KAGGLE_ENVIRONMENT, + KAGGLE_TMP, + ) + + _tokenizer_repo = ( + tokenizer_name if (isinstance(tokenizer_name, str) and tokenizer_name) else model_name + ) + _tokenizer_cache_dir = kwargs.get("cache_dir") + if _tokenizer_cache_dir is None: + if IS_COLAB_ENVIRONMENT: + _tokenizer_cache_dir = "huggingface_tokenizers_cache" + elif IS_KAGGLE_ENVIRONMENT: + _tokenizer_cache_dir = os.path.join(KAGGLE_TMP, "huggingface_tokenizers_cache") + # Warm the tokenizer repo into the cache the load will use whenever the base warm did not + # cover it: a distinct tokenizer repo, fast_inference (base warm skipped), or a tokenizer + # cache_dir that differs from the base-warm cache_dir (Colab/Kaggle special cache). + _warm_tokenizer_repo = ( + isinstance(_tokenizer_repo, str) + and bool(_tokenizer_repo) + and ( + _tokenizer_repo != model_name + or fast_inference + or _tokenizer_cache_dir != kwargs.get("cache_dir") + ) + ) + if _warm_tokenizer_repo: + maybe_prefetch_hf_snapshot( + _tokenizer_repo, + token = token, + cache_dir = _tokenizer_cache_dir, + local_files_only = kwargs.get("local_files_only", False), + tokenizer_only = True, + ) + has_rope_scaling = False try: with open(inspect.getfile(model_function), "r", encoding = "utf-8") as file: @@ -2672,6 +2747,10 @@ class FastLlamaModel: # Counteract saved tokenizers tokenizer_name = model_name if tokenizer_name is None else tokenizer_name + # Route the tokenizer load to the custom cache_dir the prefetch warmed. + _tokenizer_cache_kwargs = {} + if kwargs.get("cache_dir") is not None: + _tokenizer_cache_kwargs["cache_dir"] = kwargs["cache_dir"] tokenizer = load_correct_tokenizer( tokenizer_name = tokenizer_name, model_max_length = max_position_embeddings, @@ -2679,6 +2758,7 @@ class FastLlamaModel: token = token, trust_remote_code = trust_remote_code, fix_tokenizer = fix_tokenizer, + **_tokenizer_cache_kwargs, ) model, tokenizer = patch_tokenizer(model, tokenizer) @@ -2805,6 +2885,7 @@ class FastLlamaModel: model_max_length = max_position_embeddings, padding_side = "right", token = token, + cache_dir = kwargs.get("cache_dir"), ) patch_saving_functions(tokenizer) @@ -3743,4 +3824,17 @@ class FastLlamaModel: from .rl import PatchFastRL +# Auto-enable grouped-GEMM MoE (tf<5 ModuleList experts) on built / PEFT'd models. Wrap the +# loader leaves before PatchFastRL so downstream patchers see the wrapped versions. Guarded. +try: + from unsloth_zoo.temporary_patches.moe_grouped_modulelist import wrap_loader_for_grouped_moe + FastLlamaModel.from_pretrained = staticmethod( + wrap_loader_for_grouped_moe(FastLlamaModel.from_pretrained) + ) + FastLlamaModel.get_peft_model = staticmethod( + wrap_loader_for_grouped_moe(FastLlamaModel.get_peft_model) + ) +except Exception: + pass + PatchFastRL(FastLanguageModel = FastLlamaModel) diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index 562afdd645..ba23197861 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -21,6 +21,10 @@ from ._utils import ( USE_MODELSCOPE, get_transformers_model_type, hf_login, + # Single source of truth is _utils.py; re-exported here so callers doing + # `from unsloth.models.loader import DISABLE_SDPA_MODEL_NAMES` keep working and so + # _is_sdpa_excluded (in _utils) can honor it without a loader -> _utils cycle. + DISABLE_SDPA_MODEL_NAMES, ) from .granite import FastGraniteModel from .llama import FastLlamaModel, logger @@ -106,24 +110,31 @@ from ._utils import ( _is_family_text_decoder, _apply_text_only_key_mapping, set_task_config_attr, + maybe_prefetch_hf_snapshot, ) -# Single source of truth is unsloth_zoo.model_lists. Re-exported so callers -# doing `from unsloth.models.loader import FORCE_FLOAT32` keep working. -# Fallback list mirrors zoo for users who upgrade unsloth without upgrading -# unsloth_zoo (so this module never fails at import). +# Source of truth is unsloth_zoo.model_lists. Re-exported so callers doing +# `from unsloth.models.loader import FORCE_FLOAT32` keep working. The fallback +# list is also unioned in so a newer unsloth still forces float32 for these +# archs when paired with an older unsloth_zoo that predates them (upgrade skew). +_FORCE_FLOAT32_FALLBACK = [ + "gemma3,", # Add comma bc gemma3 will match gemma3n + "gemma3text", # Gemma3TextModel (EmbeddingGemma, standalone text-only Gemma3) + "gemma3n", + "gemma4", # Gemma4 (gemma4 / gemma4_text): float16 NaNs grad norms in the backward + "glm4_moe", # GLM-4.x MoE (glm4_moe / glm4_moe_lite): float16 NaNs grad norms + "gpt_oss", + "qwen3_5", # Qwen3.5 GDN layers produce NaN grad norms in float16 training + "qwen3_moe", # Qwen3-MoE (Qwen3-30B-A3B): float16 NaNs grad norms in the backward +] try: - from unsloth_zoo import FORCE_FLOAT32 # noqa: F401 + from unsloth_zoo import FORCE_FLOAT32 as _ZOO_FORCE_FLOAT32 + FORCE_FLOAT32 = list(_ZOO_FORCE_FLOAT32) except ImportError: - global FORCE_FLOAT32 - # Forces float32 precision since float16 goes to infinity - FORCE_FLOAT32 = [ - "gemma3,", # Add comma bc gemma3 will match gemma3n - "gemma3text", # Gemma3TextModel (EmbeddingGemma, standalone text-only Gemma3) - "gemma3n", - "gpt_oss", - "qwen3_5", # Qwen3.5 GDN layers produce NaN grad norms in float16 training - ] + FORCE_FLOAT32 = [] +for _mt in _FORCE_FLOAT32_FALLBACK: + if not any(_mt in _entry for _entry in FORCE_FLOAT32): + FORCE_FLOAT32.append(_mt) global DISABLE_COMPILE_MODEL_NAMES # Must be alphabetically sorted for each entry @@ -195,15 +206,45 @@ DISABLE_COMPILE_MODEL_NAMES = [ "granite,llava_next", # Granite-vision 3 ] -global DISABLE_SDPA_MODEL_NAMES -# Disables some SDPA modules since it's wrong -DISABLE_SDPA_MODEL_NAMES = [ - "gemma3,", # Add comma bc gemma3 will match gemma3n - "gemma3_text", # Gemma3TextModel (EmbeddingGemma) - substring match, keep underscore - "gpt_oss", -] +# Architectures with gated-deltanet (linear attention) layers. Unsloth bundles the +# flash-linear-attention Triton kernels (unsloth_zoo/_vendored/fla), so no install is +# needed; transformers uses the much slower pure PyTorch path only when they can't be enabled. +FLA_MODEL_TYPE_PREFIXES = ("qwen3_next", "qwen3_5", "kimi_linear", "olmo_hybrid") +_fla_advised = False +def _maybe_advise_fla_install(model_types): + """One-time note when a gated-deltanet model loads without the fast kernels. + + The kernels ship with Unsloth (no install needed); this fires only when they + could not be enabled on this platform (e.g. no CUDA, torch < 2.7 or + triton < 3.3), i.e. exactly when transformers uses the slow pure PyTorch path. + """ + global _fla_advised + if _fla_advised: + return + if model_types is None: + return + if isinstance(model_types, str): + model_types = [model_types] # a lone string would otherwise iterate chars + try: + if not any( + isinstance(t, str) and t.startswith(FLA_MODEL_TYPE_PREFIXES) for t in model_types + ): + return + from transformers.utils.import_utils import is_flash_linear_attention_available + if is_flash_linear_attention_available(): + return # bundled (or user-installed) fast kernels are active + except Exception: + return + _fla_advised = True + print( + "Unsloth: This model uses gated-deltanet linear attention layers. Unsloth\n" + "bundles the flash-linear-attention kernels, but they could not be enabled\n" + "on this setup (they need CUDA with torch >= 2.7 and triton >= 3.3), so\n" + "transformers will use a slower pure PyTorch path." + ) + def _fix_rope_inv_freq(model): """Fix inv_freq corruption caused by transformers v5 meta-device loading. @@ -469,8 +510,10 @@ class FastLanguageModel(FastLlamaModel): ("-unsloth-bnb-4bit", "-bnb-4bit") ): model_name = _strip_unsloth_bnb_4bit_suffix(model_name) - # Change -BF16 to all False for 4bit, 8bit etc - if model_name.lower().endswith("-bf16"): + # '-bf16' hub repos load bf16; a local dir keeps the requested quant unless 16bit is set + if model_name.lower().endswith("-bf16") and ( + load_in_16bit or not os.path.isdir(os.path.expanduser(model_name)) + ): load_in_4bit = False load_in_8bit = False load_in_fp8 = False @@ -628,8 +671,10 @@ class FastLanguageModel(FastLlamaModel): ("-unsloth-bnb-4bit", "-bnb-4bit") ): model_name = _strip_unsloth_bnb_4bit_suffix(model_name) - # Change -BF16 to all False for 4bit, 8bit etc - if model_name.lower().endswith("-bf16"): + # '-bf16' hub repos load bf16; a local dir keeps the requested quant unless 16bit is set + if model_name.lower().endswith("-bf16") and ( + load_in_16bit or not os.path.isdir(os.path.expanduser(model_name)) + ): load_in_4bit = False load_in_8bit = False load_in_fp8 = False @@ -865,6 +910,28 @@ class FastLanguageModel(FastLlamaModel): if is_peft: # From https://github.com/huggingface/peft/issues/184 # Now add PEFT adapters + # Warm the adapter repo: PeftModel downloads it in-process and can hang on Xet. + _prefetched = maybe_prefetch_hf_snapshot( + old_model_name, + token = token, + revision = revision, + cache_dir = kwargs.get("cache_dir"), + local_files_only = local_files_only, + # Adapter always loads in-process via PeftModel, so warm it even under fast_inference. + fast_inference = False, + force_download = kwargs.get("force_download", False), + # Leave use_safetensors auto (inheriting base format could skip a safetensors-only + # adapter). adapter_only restricts the warm to the adapter files + root aux. + adapter_only = True, + ) + # Child did the forced download; clear the flag so the load reuses the warm cache. + if _prefetched and kwargs.get("force_download", False): + kwargs["force_download"] = False + # Forward cache_dir so the load reads the warmed adapter. No subfolder (that targets the + # base checkpoint; adapters live at the root). + peft_load_kwargs = {} + if kwargs.get("cache_dir") is not None: + peft_load_kwargs["cache_dir"] = kwargs["cache_dir"] model = PeftModel.from_pretrained( model, old_model_name, @@ -873,9 +940,19 @@ class FastLanguageModel(FastLlamaModel): local_files_only = local_files_only, is_trainable = True, trust_remote_code = trust_remote_code, + **peft_load_kwargs, ) # Patch it as well! model = dispatch_model.patch_peft_model(model, use_gradient_checkpointing) + # Re-evaluate grouped MoE now the adapter is attached: an expert-LoRA block falls back + # to the original loop, an attention-only adapter keeps the grouped path. Guarded. + try: + from unsloth_zoo.temporary_patches.moe_grouped_modulelist import ( + auto_enable_grouped_moe, + ) + auto_enable_grouped_moe(model) + except Exception: + pass # optional speedup; never block model loading # Patch Tiled MLP # to turn on set UNSLOTH_TILED_MLP to "arctic", "target", or "target:{GB}"" @@ -1116,8 +1193,10 @@ class FastModel(FastBaseModel): ("-unsloth-bnb-4bit", "-bnb-4bit") ): model_name = _strip_unsloth_bnb_4bit_suffix(model_name) - # Change -BF16 to all False for 4bit, 8bit etc - if model_name.lower().endswith("-bf16"): + # '-bf16' hub repos load bf16; a local dir keeps the requested quant unless 16bit is set + if model_name.lower().endswith("-bf16") and ( + load_in_16bit or not os.path.isdir(os.path.expanduser(model_name)) + ): load_in_4bit = False load_in_8bit = False load_in_fp8 = False @@ -1263,6 +1342,7 @@ class FastModel(FastBaseModel): trust_remote_code = trust_remote_code, ) model_types_all = ",".join(model_types) + "," + _maybe_advise_fla_install(model_types) # ---- Text-diffusion models (e.g. DiffusionGemma) take a transformers-only slow path. ---- # These use a custom block-diffusion `generate` and a novel backbone, so we skip Unsloth's @@ -1474,8 +1554,10 @@ class FastModel(FastBaseModel): ("-unsloth-bnb-4bit", "-bnb-4bit") ): model_name = _strip_unsloth_bnb_4bit_suffix(model_name) - # Change -BF16 to all False for 4bit, 8bit etc - if model_name.lower().endswith("-bf16"): + # '-bf16' hub repos load bf16; a local dir keeps the requested quant unless 16bit is set + if model_name.lower().endswith("-bf16") and ( + load_in_16bit or not os.path.isdir(os.path.expanduser(model_name)) + ): load_in_4bit = False load_in_8bit = False load_in_fp8 = False @@ -1790,6 +1872,28 @@ class FastModel(FastBaseModel): _LoraModel._create_and_replace = _patched_car + # Warm the adapter repo: PeftModel downloads it in-process and can hang on Xet. + _prefetched = maybe_prefetch_hf_snapshot( + old_model_name, + token = token, + revision = revision, + cache_dir = kwargs.get("cache_dir"), + local_files_only = local_files_only, + # Adapter always loads in-process via PeftModel, so warm it even under fast_inference. + fast_inference = False, + force_download = kwargs.get("force_download", False), + # Leave use_safetensors auto (inheriting base format could skip a safetensors-only + # adapter). adapter_only restricts the warm to the adapter files + root aux. + adapter_only = True, + ) + # Child did the forced download; clear the flag so the load reuses the warm cache. + if _prefetched and kwargs.get("force_download", False): + kwargs["force_download"] = False + # Forward cache_dir so the load reads the warmed adapter. No subfolder (that targets the + # base checkpoint; adapters live at the root). + peft_load_kwargs = {} + if kwargs.get("cache_dir") is not None: + peft_load_kwargs["cache_dir"] = kwargs["cache_dir"] try: model = PeftModel.from_pretrained( model, @@ -1799,6 +1903,7 @@ class FastModel(FastBaseModel): local_files_only = local_files_only, is_trainable = True, trust_remote_code = trust_remote_code, + **peft_load_kwargs, ) finally: # Always restore original PEFT method, even if loading fails @@ -1809,6 +1914,15 @@ class FastModel(FastBaseModel): model = FastBaseModel.post_patch_model( model, use_gradient_checkpointing, trust_remote_code = trust_remote_code ) + # Re-evaluate grouped MoE now the adapter is attached: an expert-LoRA block falls back + # to the original loop, an attention-only adapter keeps the grouped path. Guarded. + try: + from unsloth_zoo.temporary_patches.moe_grouped_modulelist import ( + auto_enable_grouped_moe, + ) + auto_enable_grouped_moe(model) + except Exception: + pass # optional speedup; never block model loading # Apply QAT if specified if qat_scheme is not None: diff --git a/unsloth/models/mistral.py b/unsloth/models/mistral.py index df2a4de5bd..4350565fe2 100644 --- a/unsloth/models/mistral.py +++ b/unsloth/models/mistral.py @@ -27,6 +27,7 @@ from ..utils.attention_dispatch import ( run_attention, SDPA, select_attention_backend, + resolve_prefix_seg_info, ) from .llama import ( LlamaRotaryEmbedding, @@ -124,6 +125,9 @@ def MistralAttention_fast_forward( "softmax_scale": getattr(self, "softmax_scale", None), }, ) + # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward; misuse + # (KV cache / padding mask) raises. None => byte-identical default. + _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask) context = AttentionContext( bsz = bsz, q_len = q_len, @@ -134,6 +138,7 @@ def MistralAttention_fast_forward( seq_info = seq_info, attention_mask = attention_mask, causal_mask = causal_mask, + prefix_seg_info = _pg_seg, ) A = run_attention(config = attention_config, context = context, Q = Q, K = K, V = V) @@ -161,7 +166,13 @@ def MistralForCausalLM_fast_forward( *args, **kwargs, ) -> Union[Tuple, CausalLMOutputWithPast]: - if causal_mask is None and past_key_values is None: + # PrefixGrouper brings its own mask: a synthesized causal attention_mask would trip + # resolve_prefix_seg_info on the no-xFormers path and force a fallback. + if ( + causal_mask is None + and past_key_values is None + and kwargs.get("prefix_seg_info", None) is None + ): bsz, q_len = input_ids.shape sliding_window = getattr(self.config, "sliding_window", None) diff --git a/unsloth/models/qwen3.py b/unsloth/models/qwen3.py index e28e72d3ea..0d05a2d538 100644 --- a/unsloth/models/qwen3.py +++ b/unsloth/models/qwen3.py @@ -23,6 +23,7 @@ from ..utils.attention_dispatch import ( run_attention, SDPA, select_attention_backend, + resolve_prefix_seg_info, ) from .llama import ( LlamaRotaryEmbedding, @@ -146,6 +147,9 @@ def Qwen3Attention_fast_forward( "softmax_scale": getattr(self, "softmax_scale", None), }, ) + # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward; misuse + # (KV cache / padding mask) raises. None => byte-identical default. + _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask) context = AttentionContext( bsz = bsz, q_len = q_len, @@ -156,6 +160,7 @@ def Qwen3Attention_fast_forward( seq_info = seq_info, attention_mask = attention_mask, causal_mask = causal_mask, + prefix_seg_info = _pg_seg, ) A = run_attention(config = attention_config, context = context, Q = Q, K = K, V = V) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 3be614cf4a..098950de08 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -29,7 +29,9 @@ from collections import defaultdict from unsloth_zoo.rl_replacements import ( RL_REPLACEMENTS, left_pack_padding, + create_completion_attention_mask, chunked_selective_log_softmax, + chunked_hidden_states_selective_log_softmax, _unsloth_get_mm_token_id, _unsloth_fix_mm_token_type_ids, ) @@ -48,7 +50,41 @@ from ..device_type import ( ALLOW_PREQUANTIZED_MODELS, ) import textwrap -from ._utils import _get_inference_mode_context_manager +from ._utils import _get_inference_mode_context_manager, UNSLOTH_ENABLE_LOGGING + +# One-time GRPO sequence-packing gates; mirrored into the generated trainer cache via RL_PRE_ITEMS. +UNSLOTH_GRPO_SEQ_PACKING_ON = os.environ.get("UNSLOTH_GRPO_SEQ_PACKING", "1").lower() not in ( + "0", + "false", + "no", + "off", +) +# Packing needs zoo#840's masked-column guard in grpo_compute_loss (installed zoo is fixed per-process). +try: + UNSLOTH_ZOO_HAS_MASKED_COL_GUARD = "torch.where(_keep, new" in inspect.getsource( + RL_REPLACEMENTS["grpo_compute_loss"] + ) +except Exception: + UNSLOTH_ZOO_HAS_MASKED_COL_GUARD = False +# One-time PrefixGrouper gate; any import failure degrades to "PrefixGrouper off". +_pg_build_layout = _pg_enabled_fn = _pg_verify_on = _pg_tol_ok = _PG_TOL_KILL = None +UNSLOTH_GRPO_PREFIX_GROUPER_ON = os.environ.get("UNSLOTH_GRPO_PREFIX_GROUPER", "1").lower() not in ( + "0", + "false", + "no", + "off", +) +if UNSLOTH_GRPO_PREFIX_GROUPER_ON: + try: + from ..utils.prefix_grouper import ( + build_group_layout as _pg_build_layout, + prefix_grouper_enabled as _pg_enabled_fn, + verify_on as _pg_verify_on, + tol_ok as _pg_tol_ok, + TOL_KILL as _PG_TOL_KILL, + ) + except Exception: + UNSLOTH_GRPO_PREFIX_GROUPER_ON = False RL_EXTRA_ARGS = defaultdict(list) RL_FUNCTIONS = defaultdict(list) @@ -1359,6 +1395,439 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): ) os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" + # ---- Sequence packing (default-on; disable with UNSLOTH_GRPO_SEQ_PACKING=0) ---- + # One varlen [1, sum L] forward replaces the padded [B, Lmax] loop (also fixes the + # left-pad RoPE error). Self-verified against the per-row forward, re-checked as T + # grows; falls back if a backend ignores packed_seq_lengths. + logprobs = None + + # ---- PrefixGrouper (GRPO shared-prompt dedup; default ON, exact + self-verified) ---- + # G completions per prompt share the prefix; the packed path forwards it G times, + # PrefixGrouper stores it once (FlexAttention shared-prefix mask), cutting the trunk + # forward from G*(P+R) to P+G*R tokens. Gated by UNSLOTH_GRPO_PREFIX_GROUPER (needs + # seq-packing), tok_r auto-gate, and first-use self-verify vs the packed path + # (mismatch => fall back + mark unsafe), so a mask/isolation regression cannot ship + # silently. When off / ungrouped / unverified, the packed path below runs as before. + _pg_result = None + _pg_use = False + _pg_skip_pk = False # once a shape is PG-verified, skip the full-row forward + _pg_forward_fn = None # deferred PG forward (runs at the verify site below) + _pg_num_gen = getattr(self, "num_generations", None) + # Env gate hoisted to module level (mirrored via RL_PRE_ITEMS). Skip PG under vLLM + # (fast_inference=True): the rollout dominates the step, so PG saves little and its + # first-use self-verify is net overhead. + _pg_engage = ( + UNSLOTH_GRPO_PREFIX_GROUPER_ON + and not getattr(self, "use_vllm", False) + and not getattr(unwrapped_model, "_unsloth_prefix_grouper_nograd_disabled", False) + ) + if _pg_engage: + try: + # Skip softcap models (the flex kernel never applies attn_logit_softcapping) + # and hybrid SSM / MoE models: only the threaded attention forwards get the + # shared-prefix isolation, so a Mamba or MoE decoder that does not forward + # prefix_seg_info would leak suffixes across completions. PG also rides on + # sequence packing, so it needs the same zoo masked-column guard. + _pg_cfg = getattr(unwrapped_model, "config", None) + _pg_engage = ( + _pg_enabled_fn() + and UNSLOTH_ZOO_HAS_MASKED_COL_GUARD + and pixel_values is None + and token_type_ids is None + and mm_token_type_ids is None + and _pg_num_gen is not None + and _pg_num_gen >= 2 + and not getattr(_pg_cfg, "attn_logit_softcapping", None) + # normal backends apply config.attention_dropout in training; the flex + # path is deterministic, so skip PG when it is set. + and not getattr(_pg_cfg, "attention_dropout", 0) + and not any( + getattr(_pg_cfg, _pg_a, None) is not None + for _pg_a in ( + "mamba_d_ssm", + "mamba_d_state", + "mamba_expand", + "num_experts", + "num_local_experts", + "n_routed_experts", + "moe_intermediate_size", + ) + ) + ) + except Exception: + _pg_engage = False + if _pg_engage: + try: + _pg_pad = self.processing_class.pad_token_id + # cap the PG span (P+max(R)) at the sliding window, like the packed _pk_sw guard. + _pg_sw = getattr( + getattr(unwrapped_model, "config", None), "sliding_window", None + ) + if not (isinstance(_pg_sw, int) and _pg_sw > 0): + _pg_sw = None + _pg_layout = _pg_build_layout( + input_ids, + logits_to_keep, + _pg_pad, + _pg_num_gen, + left_pad_tokens_per_prompt, + max_segment_cap = _pg_sw, + ) + _pg_unsafe = getattr( + unwrapped_model, "_unsloth_prefix_grouper_nograd_unsafe", None + ) + if _pg_unsafe is None: + _pg_unsafe = set() + if _pg_layout is not None and _pg_layout.signature not in _pg_unsafe: + _pg_sig = _pg_layout.signature + _pg_verified = getattr( + unwrapped_model, "_unsloth_prefix_grouper_nograd_verified", None + ) + if _pg_verified is None: + _pg_verified = set() + _pg_chunks = max(1, total_rows * multiplier) + + def _pg_run_forward(_pg_layout = _pg_layout, _pg_chunks = _pg_chunks): + with _get_inference_mode_context_manager(model): + with torch.amp.autocast( + device_type = "cuda", dtype = self._autocast_dtype + ): + _pg_hidden = unwrapped_model( + input_ids = _pg_layout.flat_ids, + position_ids = _pg_layout.position_ids, + prefix_seg_info = _pg_layout.prefix_seg_info, + use_cache = False, + ).logits + _pg_r = _pg_layout.extract_logps( + _pg_hidden, + lm_head, + chunked_hidden_states_selective_log_softmax, + _pg_chunks, + logit_scale_multiply, + logit_scale_divide, + logit_softcapping, + temperature, + ) + _pg_hidden = None # release before any verify forward + device_synchronize() + # clip to the loss window [B, logits_to_keep+max_left_pad] + _pg_w = logits_to_keep + max_left_pad + if _pg_r.shape[1] > _pg_w: + _pg_r = _pg_r[:, -_pg_w:] + return _pg_r + + # trust only within the verified envelope: re-verify when T or the + # longest segment grows, like the packed path + _pg_T = int(_pg_layout.flat_ids.shape[1]) + _pg_maxseg = int(_pg_layout.position_ids.max()) + 1 + _pg_env = ( + _pg_verified.get(_pg_sig) if isinstance(_pg_verified, dict) else None + ) + if (not _pg_verify_on()) or ( + _pg_env is not None and _pg_T <= _pg_env[0] and _pg_maxseg <= _pg_env[1] + ): + # trusted shape: run PG now and skip the full-row forward below + _pg_result = _pg_run_forward() + _pg_use = True + _pg_skip_pk = True + else: + # unverified shape: defer the forward until the packed reference + # exists (verify site below), so a declined packed path never wastes + # a whole-batch PG forward + _pg_forward_fn = _pg_run_forward + except Exception as _pg_err: + _pg_result = None + _pg_use = False + _pg_skip_pk = False + _pg_forward_fn = None + # A FlexAttention/Triton compile failure or OOM here is GPU-wide, not + # layout-specific, so retrying the same PG forward every step just re-pays + # the failure. Persistently disable PG (mirrors the seq-packing handler + # setting _unsloth_seq_packing_nograd_ok = False); the packed/padded path + # below still produces the exact result. + unwrapped_model._unsloth_prefix_grouper_nograd_disabled = True + if isinstance(_pg_err, torch.cuda.OutOfMemoryError): + torch.cuda.empty_cache() + os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" + if UNSLOTH_ENABLE_LOGGING: + print( + f"[Unsloth] GRPO PrefixGrouper (no-grad) disabled (fell back to packed): {_pg_err!r}", + flush = True, + ) + + # ---- Sequence packing (default-on; disable with UNSLOTH_GRPO_SEQ_PACKING=0) ---- + # One varlen [1, sum L] block-diagonal forward replaces the padded [B, Lmax] loop + # (exact per-row result; also fixes the padded path's left-pad RoPE error). + # Self-verified vs the per-row forward, re-checked as T grows; falls back if a + # backend ignores packed_seq_lengths. lm_head runs on completion positions only. + _pk_result = None + _pk_use = False + _pk_enabled = UNSLOTH_GRPO_SEQ_PACKING_ON + # Without zoo#840's masked-column guard, zeroed prompt/pad columns turn NaN in exp(). + _pk_enabled = _pk_enabled and UNSLOTH_ZOO_HAS_MASKED_COL_GUARD + _pk_ok = getattr(unwrapped_model, "_unsloth_seq_packing_nograd_ok", None) + if ( + _pk_enabled + and not _pg_skip_pk + and pixel_values is None + and token_type_ids is None + and mm_token_type_ids is None + and _pk_ok is not False + ): + try: + _pk_pad = self.processing_class.pad_token_id + _pk_keep = input_ids != _pk_pad + _pk_len = _pk_keep.sum(dim = 1) + _pk_len_cpu = _pk_len.tolist() # single GPU->CPU sync, reused below + _pk_nz_cpu = [_n for _n in _pk_len_cpu if _n > 0] + _pk_flat = input_ids[_pk_keep].unsqueeze(0) + _pk_T = _pk_flat.shape[1] + _pk_L = input_ids.shape[1] + _pk_W = logits_to_keep + max_left_pad + _pk_maxseg = max(_pk_nz_cpu) if _pk_nz_cpu else 0 + # sliding-window models lose the per-sequence local window in a packed stream + _pk_sw = getattr( + getattr(unwrapped_model, "config", None), "sliding_window", None + ) + _pk_sw_ok = not (isinstance(_pk_sw, int) and _pk_sw > 0 and _pk_maxseg > _pk_sw) + # per-row completion mask (same as the loss); prompt-only rows count as inactive + _pk_cmask = create_completion_attention_mask( + input_ids[:, -_pk_W:], left_pad_tokens_per_prompt, max_left_pad, _pk_pad + ) + _pk_active = int(_pk_cmask.any(dim = 1).sum()) + # skip the packed forward entirely at known-unsafe lengths (avoids a wasted pass / OOM) + _pk_unsafe = getattr( + unwrapped_model, "_unsloth_seq_packing_nograd_unsafe_T", None + ) + # cap the flattened forward at one padded [batch_size, seq_len] mini-batch's + # token budget; anything larger uses the chunked padded loop + _pk_cap = batch_size * seq_len + if ( + _pk_T >= 2 + and _pk_T <= _pk_cap + and len(_pk_nz_cpu) > 0 + and _pk_sw_ok + and not (_pk_unsafe is not None and _pk_T >= _pk_unsafe) + and (_pk_ok is True or _pk_active >= 2) + ): + # reset 0-based position_ids per segment + _pk_pos = (_pk_keep.cumsum(dim = 1) - 1)[_pk_keep].unsqueeze(0) + _pk_chunks = max(1, total_rows * multiplier) + _pk_nz_idx = _pk_keep.nonzero( + as_tuple = False + ) # [T, 2] = (row, col), row-major + _pk_within = _pk_nz_idx[1:, 0] == _pk_nz_idx[:-1, 0] # [T-1] + # per-row completion start after left-packing (matches create_completion_attention_mask) + _pk_cstart = (_pk_L - logits_to_keep) - left_pad_tokens_per_prompt # [rows] + _pk_ctgt = (_pk_nz_idx[1:, 1] >= _pk_cstart[_pk_nz_idx[1:, 0]]) & _pk_within + with _get_inference_mode_context_manager(model): + with torch.amp.autocast(device_type = "cuda", dtype = self._autocast_dtype): + # use_cache=False: a KV cache silently disables varlen packing + _pk_hidden = unwrapped_model( + input_ids = _pk_flat, + position_ids = _pk_pos, + packed_seq_lengths = torch.tensor( + _pk_nz_cpu, dtype = torch.int32, device = input_ids.device + ), + use_cache = False, + ).logits + _pk_sel = chunked_hidden_states_selective_log_softmax( + _pk_hidden[0, :-1, :][_pk_ctgt].unsqueeze(0), + lm_head, + _pk_flat[0, 1:][_pk_ctgt].unsqueeze(0), + _pk_chunks, + logit_scale_multiply, + logit_scale_divide, + logit_softcapping, + temperature, + )[0] + # GPT-OSS offload race guard (matches the padded loop) + device_synchronize() + # scatter each logprob back to its (row, col) so [:, -_pk_W:] matches padded + _pk_tgt = (_pk_nz_idx[1:, 0] * _pk_L + _pk_nz_idx[1:, 1])[_pk_ctgt] + _pk_result = ( + torch.zeros( + total_rows * _pk_L, + dtype = torch.float32, + device = input_ids.device, + ) + .index_put((_pk_tgt,), _pk_sel.to(torch.float32)) + .view(total_rows, _pk_L)[:, -_pk_W:] + ) + # re-verify when T or the longest segment grows past what was verified + # (a LongRoPE cache switch can change the result) + _pk_vT = int( + getattr(unwrapped_model, "_unsloth_seq_packing_nograd_verified_T", 0) + ) + _pk_vS = int( + getattr(unwrapped_model, "_unsloth_seq_packing_nograd_verified_seg", 0) + ) + # debug: hand-edit this condition to force re-verify every step + if _pk_ok is True and _pk_T <= _pk_vT and _pk_maxseg <= _pk_vS: + _pk_use = True # already verified for this shape + else: + # verify against the per-row forward (ground truth) + _pk_ref = torch.zeros_like(_pk_result) + with _get_inference_mode_context_manager(model): + with torch.amp.autocast( + device_type = "cuda", dtype = self._autocast_dtype + ): + for _pk_i in range(total_rows): + _pk_ni = _pk_len_cpu[_pk_i] + if _pk_ni < 2: + continue + _pk_rmask = _pk_keep[_pk_i] + _pk_real = input_ids[_pk_i][_pk_rmask].unsqueeze(0) + _pk_rpos = torch.arange( + _pk_ni, device = input_ids.device + ).unsqueeze(0) + _pk_rh = unwrapped_model( + input_ids = _pk_real, + position_ids = _pk_rpos, + use_cache = False, + ).logits + _pk_rsel = chunked_hidden_states_selective_log_softmax( + _pk_rh[:, :-1, :], + lm_head, + _pk_real[:, 1:], + 1, + logit_scale_multiply, + logit_scale_divide, + logit_softcapping, + temperature, + )[0] + _pk_rcols = _pk_rmask.nonzero(as_tuple = False).squeeze(1)[ + 1: + ] - (_pk_L - _pk_W) + _pk_rkeep = _pk_rcols >= 0 + _pk_ref[_pk_i, _pk_rcols[_pk_rkeep]] = _pk_rsel[ + _pk_rkeep + ].to(torch.float32) + device_synchronize() + # compare over the loss-mask region only + _pk_cm = _pk_cmask.float() + _pk_diff = float(((_pk_result - _pk_ref).abs() * _pk_cm).max()) + if UNSLOTH_ENABLE_LOGGING: + print( + f"[Unsloth] GRPO seq-packing (no-grad) verify: T={_pk_T} maxseg={_pk_maxseg} packed-vs-perrow max|d|={_pk_diff:.4f}", + flush = True, + ) + # kernel-noise floor ~0.25; cross-sample contamination is >= 2.4 + if _pk_diff < 7e-1: + unwrapped_model._unsloth_seq_packing_nograd_ok = True + # widen the trusted shape only when >= 2 completion rows exercised + # cross-sample packing; single-row passes prove nothing + if _pk_active >= 2: + unwrapped_model._unsloth_seq_packing_nograd_verified_T = max( + _pk_vT, _pk_T + ) + unwrapped_model._unsloth_seq_packing_nograd_verified_seg = max( + _pk_vS, _pk_maxseg + ) + _pk_ok = True + _pk_use = True + else: + _pk_use = False + if _pk_diff >= 1.5: + # contamination (attention ignores the packed mask): disable packing + unwrapped_model._unsloth_seq_packing_nograd_ok = False + else: + # likely a length boundary (LongRoPE): mark unsafe, keep smaller shapes + unwrapped_model._unsloth_seq_packing_nograd_unsafe_T = ( + _pk_T if _pk_unsafe is None else min(_pk_unsafe, _pk_T) + ) + if UNSLOTH_ENABLE_LOGGING: + print( + f"[Unsloth] GRPO seq-packing (no-grad) fell back at T={_pk_T} (diff={_pk_diff:.3f})", + flush = True, + ) + except Exception as _pk_err: + # any failure: drop intermediates, use the padded loop, do not retry + _pk_hidden = None + _pk_sel = None + _pk_result = None + _pk_use = False + if isinstance(_pk_err, torch.cuda.OutOfMemoryError): + torch.cuda.empty_cache() + unwrapped_model._unsloth_seq_packing_nograd_ok = False + if UNSLOTH_ENABLE_LOGGING: + print( + f"[Unsloth] GRPO sequence-packing (no-grad) disabled (fell back to padded): {_pk_err!r}", + flush = True, + ) + # ---- PrefixGrouper first-use self-verify (no-grad) ---- + # Compare the untrusted PG result to the full-row packed result (itself verified vs + # per-row) over the completion mask: < tol_ok -> trust the structure; >= TOL_KILL -> + # unsafe forever; borderline -> fall back this shape. + if _pg_forward_fn is not None and not _pg_use: + if _pk_use and _pk_result is not None: + try: + # deferred PG forward, run only now that the packed reference exists + _pg_result = _pg_forward_fn() + _pg_W2 = logits_to_keep + max_left_pad + _pg_cm = create_completion_attention_mask( + input_ids[:, -_pg_W2:], + left_pad_tokens_per_prompt, + max_left_pad, + self.processing_class.pad_token_id, + ).float() + _pg_a = _pg_result[:, -_pg_W2:].float() + _pg_b = _pk_result[:, -_pg_W2:].float() + _pg_diff = float(((_pg_a - _pg_b).abs() * _pg_cm).max()) + if UNSLOTH_ENABLE_LOGGING: + print( + f"[Unsloth] GRPO PrefixGrouper (no-grad) verify: sig={_pg_layout.signature} " + f"shared-prefix vs full-row-packed max|d|={_pg_diff:.4f}", + flush = True, + ) + if _pg_diff < _pg_tol_ok(): + _pg_v = getattr( + unwrapped_model, "_unsloth_prefix_grouper_nograd_verified", None + ) + if not isinstance(_pg_v, dict): + _pg_v = {} + _pg_vT = int(_pg_layout.flat_ids.shape[1]) + _pg_vS = int(_pg_layout.position_ids.max()) + 1 + _pg_old = _pg_v.get(_pg_layout.signature, (0, 0)) + _pg_v[_pg_layout.signature] = ( + max(_pg_vT, _pg_old[0]), + max(_pg_vS, _pg_old[1]), + ) + unwrapped_model._unsloth_prefix_grouper_nograd_verified = _pg_v + _pg_use = True + else: + _pg_u = getattr( + unwrapped_model, "_unsloth_prefix_grouper_nograd_unsafe", None + ) + if _pg_u is None: + _pg_u = set() + if _pg_diff >= _PG_TOL_KILL: + _pg_u.add(_pg_layout.signature) + unwrapped_model._unsloth_prefix_grouper_nograd_unsafe = _pg_u + _pg_use = False + except Exception as _pg_err3: + _pg_result = None + _pg_use = False + if isinstance(_pg_err3, torch.cuda.OutOfMemoryError): + torch.cuda.empty_cache() + os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" + if UNSLOTH_ENABLE_LOGGING: + print( + f"[Unsloth] GRPO PrefixGrouper (no-grad) verify failed (fell back to packed): {_pg_err3!r}", + flush = True, + ) + # else: no packed reference (packing off/failed) -> cannot verify; fall back. + + if _pg_use and _pg_result is not None: + logprobs = _pg_result # PrefixGrouper verified/trusted -> skip the loop + zipped_inputs = [] + elif _pk_use and _pk_result is not None: + logprobs = _pk_result # verified -> skip the loop + zipped_inputs = [] + else: + # free packed intermediates before running the padded loop + _pk_hidden = _pk_sel = _pk_result = _pk_ref = None + with _get_inference_mode_context_manager(model): for ( input_ids_chunk, @@ -1443,7 +1912,8 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): # However, it seems that this line does not slow down or disrupt models. device_synchronize() all_logprobs_list.append(logprobs_chunk) - logprobs = torch.cat(all_logprobs_list, dim = 0) + if logprobs is None: # padded fallback when packing was not used + logprobs = torch.cat(all_logprobs_list, dim = 0) entropies = None os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "0" @@ -1523,6 +1993,34 @@ RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(grpo_accumulated_loss)) RL_PRE_ITEMS["grpo_trainer"].append(grpo_compute_loss_slow) RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(grpo_update_SamplingParams)) RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(_get_inference_mode_context_manager)) +# inspect.getsource inlines function bodies but not module imports, so constants the inlined +# grpo functions reference (e.g. UNSLOTH_ENABLE_LOGGING) must be redefined in the generated cache. +RL_PRE_ITEMS["grpo_trainer"].append( + "import os as _unsloth_os\n" + "UNSLOTH_ENABLE_LOGGING = _unsloth_os.environ.get('UNSLOTH_ENABLE_LOGGING', '0') in ('1', 'True', 'true')\n" +) +# Sequence-packing gates, same values as the module-top constants. +RL_PRE_ITEMS["grpo_trainer"].append( + "UNSLOTH_GRPO_SEQ_PACKING_ON = _unsloth_os.environ.get('UNSLOTH_GRPO_SEQ_PACKING', '1').lower() not in ('0', 'false', 'no', 'off')\n" +) +RL_PRE_ITEMS["grpo_trainer"].append( + "try:\n" + " import inspect as _unsloth_inspect\n" + " from unsloth_zoo.rl_replacements import RL_REPLACEMENTS as _unsloth_zoo_RL\n" + " UNSLOTH_ZOO_HAS_MASKED_COL_GUARD = 'torch.where(_keep, new' in _unsloth_inspect.getsource(_unsloth_zoo_RL['grpo_compute_loss'])\n" + "except Exception:\n" + " UNSLOTH_ZOO_HAS_MASKED_COL_GUARD = False\n" +) +# PrefixGrouper gate, same shape as the module-top constants. +RL_PRE_ITEMS["grpo_trainer"].append( + "_pg_build_layout = _pg_enabled_fn = _pg_verify_on = _pg_tol_ok = _PG_TOL_KILL = None\n" + "UNSLOTH_GRPO_PREFIX_GROUPER_ON = _unsloth_os.environ.get('UNSLOTH_GRPO_PREFIX_GROUPER', '1').lower() not in ('0', 'false', 'no', 'off')\n" + "if UNSLOTH_GRPO_PREFIX_GROUPER_ON:\n" + " try:\n" + " from unsloth.utils.prefix_grouper import build_group_layout as _pg_build_layout, prefix_grouper_enabled as _pg_enabled_fn, verify_on as _pg_verify_on, tol_ok as _pg_tol_ok, TOL_KILL as _PG_TOL_KILL\n" + " except Exception:\n" + " UNSLOTH_GRPO_PREFIX_GROUPER_ON = False\n" +) # Edit _get_per_token_logps to handle mixed precision diff --git a/unsloth/models/sentence_transformer.py b/unsloth/models/sentence_transformer.py index 7e43442bfd..c1172faa94 100644 --- a/unsloth/models/sentence_transformer.py +++ b/unsloth/models/sentence_transformer.py @@ -19,6 +19,7 @@ from ._utils import ( SUPPORTS_BFLOAT16, resolve_model_class, resolve_encoder_attention_implementation, + maybe_prefetch_hf_snapshot, ) import inspect import json @@ -541,7 +542,12 @@ class FastSentenceTransformer(FastModel): return transformer_module @staticmethod - def _read_pooling_mode(model_name, token): + def _read_pooling_mode( + model_name, + token, + cache_dir = None, + revision = None, + ): """Read the pooling mode from modules.json, else return "mean".""" try: if os.path.exists(model_name) and os.path.exists( @@ -549,7 +555,13 @@ class FastSentenceTransformer(FastModel): ): modules_json_path = os.path.join(model_name, "modules.json") else: - modules_json_path = hf_hub_download(model_name, "modules.json", token = token) + modules_json_path = hf_hub_download( + model_name, + "modules.json", + token = token, + cache_dir = cache_dir, + revision = revision, + ) with open(modules_json_path, "r", encoding = "utf-8") as f: modules_config = json.load(f) @@ -571,6 +583,8 @@ class FastSentenceTransformer(FastModel): model_name, os.path.join(pooling_path, "config.json"), token = token, + cache_dir = cache_dir, + revision = revision, ) break @@ -950,7 +964,12 @@ class FastSentenceTransformer(FastModel): f.write(content) @staticmethod - def _module_path(model_name, token = None): + def _module_path( + model_name, + token = None, + cache_dir = None, + revision = None, + ): """Return the path to the modules.json file, or None.""" try: if os.path.exists(model_name) and os.path.isdir(model_name): @@ -958,7 +977,13 @@ class FastSentenceTransformer(FastModel): return path if os.path.exists(path) else None else: try: - return hf_hub_download(model_name, "modules.json", token = token) + return hf_hub_download( + model_name, + "modules.json", + token = token, + cache_dir = cache_dir, + revision = revision, + ) except: return None except: @@ -1135,6 +1160,8 @@ class FastSentenceTransformer(FastModel): max_seq_length, pooling_mode, trust_remote_code = False, + cache_dir = None, + revision = None, ) -> tuple[OrderedDict, bool]: """Load modules from modules.json, else fall back to hard-coded modules. @@ -1145,7 +1172,9 @@ class FastSentenceTransformer(FastModel): from sentence_transformers.models import Pooling, Normalize modules = OrderedDict() - modules_json_path = FastSentenceTransformer._module_path(model_name, token) + modules_json_path = FastSentenceTransformer._module_path( + model_name, token, cache_dir = cache_dir, revision = revision + ) if modules_json_path: with open(modules_json_path, encoding = "utf8") as f: @@ -1171,7 +1200,13 @@ class FastSentenceTransformer(FastModel): load_path = os.path.join(model_name, module_path) else: try: - load_path = load_dir_path(model_name, module_path, token = token) + load_path = load_dir_path( + model_name, + module_path, + token = token, + cache_folder = cache_dir, + revision = revision, + ) except Exception as e: print(f"Unsloth Warning: Could not download module {module_path}: {e}") continue @@ -1198,7 +1233,9 @@ class FastSentenceTransformer(FastModel): hidden_size = getattr(model.config, "hidden_size", 768) if pooling_mode == "mean": - pooling_mode = FastSentenceTransformer._read_pooling_mode(model_name, token) + pooling_mode = FastSentenceTransformer._read_pooling_mode( + model_name, token, cache_dir = cache_dir, revision = revision + ) modules["1"] = Pooling(word_embedding_dimension = hidden_size, pooling_mode = pooling_mode) modules["2"] = Normalize() @@ -1386,6 +1423,45 @@ class FastSentenceTransformer(FastModel): "Run `pip install sentence-transformers` to install it." ) + # Validate the load modes BEFORE the prefetch so a bad config fails without downloading weights. + # Guard on not for_inference: that branch below never used these flags. + if not for_inference: + # sanity check, thanks Etherl: + if full_finetuning and (load_in_4bit or load_in_8bit): + print( + "Unsloth: You selected full finetuning support, but 4bit / 8bit is enabled - disabling LoRA / QLoRA." + ) + load_in_4bit = False + load_in_8bit = False + load_in_fp8 = False + load_in_16bit = False + + if int(load_in_4bit) + int(load_in_8bit) + int(load_in_16bit) >= 2: + raise RuntimeError( + "Unsloth: Can only load in 4bit or 8bit or 16bit, not a combination!\n" + "Also, we by default set `load_in_16bit = True`.\n" + "If you want 4bit LoRA finetuning, set `load_in_16bit = False` and `load_in_4bit = True`\n" + "If you want 8bit finetuning, set both `load_in_16bit = False` and `load_in_8bit = True`" + ) + + # Prefetch so the ST load below is a cache hit. weights_at_root stays False (ST component + # weights live in per-module subfolders). Resolve the same cache the load uses: HF cache_dir, + # else cache_folder, else SENTENCE_TRANSFORMERS_HOME, else default -- a wrong cache misses the warm. + _st_prefetched = maybe_prefetch_hf_snapshot( + model_name, + token = token, + revision = revision, + cache_dir = kwargs.get("cache_dir") + or kwargs.get("cache_folder") + or os.environ.get("SENTENCE_TRANSFORMERS_HOME"), + local_files_only = kwargs.get("local_files_only", False), + # Forward force_download so the refresh happens in the killable child, then clear it so the + # in-process ST load reuses the warm cache instead of re-downloading over unguarded Xet. + force_download = kwargs.get("force_download", False), + ) + if _st_prefetched and kwargs.get("force_download", False): + kwargs["force_download"] = False + # if for_inference == True, skip Unsloth optimizations to avoid torch compile issues if for_inference: st_device = device_map @@ -1416,27 +1492,16 @@ class FastSentenceTransformer(FastModel): if k in kwargs: st_kwargs[k] = kwargs[k] + # ST takes cache_folder, not cache_dir: map cache_dir onto it so this load hits the warm + # (None lets ST honor SENTENCE_TRANSFORMERS_HOME, matching the prefetch). + _st_cache = kwargs.get("cache_dir") or kwargs.get("cache_folder") + if _st_cache is not None: + st_kwargs["cache_folder"] = _st_cache + st_model = SentenceTransformer(model_name, **st_kwargs) return st_model - # sanity check, thanks Etherl: - if full_finetuning and (load_in_4bit or load_in_8bit): - print( - "Unsloth: You selected full finetuning support, but 4bit / 8bit is enabled - disabling LoRA / QLoRA." - ) - load_in_4bit = False - load_in_8bit = False - load_in_fp8 = False - load_in_16bit = False - - if int(load_in_4bit) + int(load_in_8bit) + int(load_in_16bit) >= 2: - raise RuntimeError( - "Unsloth: Can only load in 4bit or 8bit or 16bit, not a combination!\n" - "Also, we by default set `load_in_16bit = True`.\n" - "If you want 4bit LoRA finetuning, set `load_in_16bit = False` and `load_in_4bit = True`\n" - "If you want 8bit finetuning, set both `load_in_16bit = False` and `load_in_8bit = True`" - ) - + # Load-mode validation already ran before the prefetch above. if "auto_model" not in kwargs: kwargs["auto_model"] = AutoModel @@ -1533,7 +1598,8 @@ class FastSentenceTransformer(FastModel): elif is_mpnet: FastSentenceTransformer._patch_mpnet_v5() - # Load via native SentenceTransformer (bypasses Unsloth patching) + # ST takes cache_folder, not cache_dir: map cache_dir onto it so this load hits the warm + # (None lets ST honor SENTENCE_TRANSFORMERS_HOME, matching the prefetch). st_model = SentenceTransformer( model_name, device = st_device, @@ -1541,6 +1607,7 @@ class FastSentenceTransformer(FastModel): token = token, revision = revision, model_kwargs = model_kwargs, + cache_folder = kwargs.get("cache_dir") or kwargs.get("cache_folder"), ) # Store metadata for get_peft_model @@ -1646,7 +1713,18 @@ class FastSentenceTransformer(FastModel): # No modules.json -> force 16-bit: saving is custom for these models and # 4-bit would need dequant in save_pretrained_merged, not worth it. - has_modules_json = FastSentenceTransformer._module_path(model_name, token) is not None + # Resolve the warmed cache: hf_hub_download ignores SENTENCE_TRANSFORMERS_HOME, so pass it as cache_dir. + has_modules_json = ( + FastSentenceTransformer._module_path( + model_name, + token, + cache_dir = kwargs.get("cache_dir") + or kwargs.get("cache_folder") + or os.environ.get("SENTENCE_TRANSFORMERS_HOME"), + revision = revision, + ) + is not None + ) if not has_modules_json and load_in_4bit: print( @@ -1656,6 +1734,12 @@ class FastSentenceTransformer(FastModel): load_in_4bit = False load_in_16bit = True + # The fallback FastModel load reads HF cache_dir, not ST's cache_folder/SENTENCE_TRANSFORMERS_HOME. + # Point it at the warmed cache, but only when no explicit cache_dir was passed (which wins). + _st_cache_dir = kwargs.get("cache_folder") or os.environ.get("SENTENCE_TRANSFORMERS_HOME") + if _st_cache_dir is not None and "cache_dir" not in kwargs: + kwargs["cache_dir"] = _st_cache_dir + try: model, tokenizer = FastModel.from_pretrained( model_name = model_name, @@ -1697,6 +1781,12 @@ class FastSentenceTransformer(FastModel): max_seq_length, pooling_mode, trust_remote_code = trust_remote_code, + # Same resolved cache as above so the fallback module loads hit the warm, not Xet. + cache_dir = kwargs.get("cache_dir") + or kwargs.get("cache_folder") + or os.environ.get("SENTENCE_TRANSFORMERS_HOME"), + # Same revision as the weight load so modules hit the warm (None = default branch). + revision = revision, ) st_device = device_map diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index 870713a05c..f39c4a57b0 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -37,6 +37,7 @@ from ._utils import ( _get_text_only_config, _is_family_text_decoder, _apply_text_only_key_mapping, + _select_moe_detection_targets, set_task_config_attr, ) from ._utils import * @@ -559,6 +560,7 @@ def _construct_vlm_processor_fallback( model_type, token, trust_remote_code, + cache_dir = None, local_files_only = False, ): """Build a VLM processor manually when AutoProcessor.from_pretrained fails (some VLMs @@ -576,6 +578,7 @@ def _construct_vlm_processor_fallback( tokenizer_name, token = token, trust_remote_code = trust_remote_code, + cache_dir = cache_dir, local_files_only = local_files_only, ) # Load tokenizer via PreTrainedTokenizerFast (bypasses tokenizer_class check) @@ -584,6 +587,7 @@ def _construct_vlm_processor_fallback( padding_side = "left", token = token, trust_remote_code = trust_remote_code, + cache_dir = cache_dir, local_files_only = local_files_only, ) # Read tokenizer_config.json for special tokens: prefer the local file (offline @@ -609,6 +613,7 @@ def _construct_vlm_processor_fallback( tokenizer_name, "tokenizer_config.json", token = token, + cache_dir = cache_dir, local_files_only = local_files_only, ) with open(config_path, "r", encoding = "utf-8") as f: @@ -640,6 +645,7 @@ def _construct_vlm_processor_fallback( tokenizer_name, token = token, trust_remote_code = trust_remote_code, + cache_dir = cache_dir, local_files_only = local_files_only, ) proc_class_name = PROCESSOR_MAPPING_NAMES.get(config.model_type) @@ -880,6 +886,9 @@ class FastBaseModel: # For debugging - we use a download counter to see if environments are not breaking or if HF is down get_statistics(kwargs.get("local_files_only", False)) + # The base + tokenizer prefetch runs AFTER the load-mode validation below, so an invalid + # load_in_* combination fails without first downloading a snapshot. + if dtype is None: dtype = torch.float16 if not SUPPORTS_BFLOAT16 else torch.bfloat16 elif os.environ.get("UNSLOTH_FORCE_FLOAT32", "0") == "1": @@ -976,6 +985,53 @@ class FastBaseModel: raise RuntimeError( "Unsloth: Can only load in 4bit or 8bit or 16bit, not a combination!" ) + + # Prefetch the repo (killable child) so the in-process load below is a cache hit. vLLM owns the + # weight download only when actually available; if fast_inference was requested but vLLM is + # missing, the load falls through in-process, so weights must still be warmed here. + _vllm_owns_weights = fast_inference and is_vLLM_available() + _prefetched = maybe_prefetch_hf_snapshot( + model_name, + token = token, + revision = kwargs.get("revision"), + cache_dir = kwargs.get("cache_dir"), + local_files_only = kwargs.get("local_files_only", False), + fast_inference = _vllm_owns_weights, + subfolder = kwargs.get("subfolder"), + force_download = kwargs.get("force_download", False), + use_safetensors = kwargs.get("use_safetensors"), + from_tf = kwargs.get("from_tf", False), + from_flax = kwargs.get("from_flax", False), + # Bare load reads only ROOT weights; skip subdir weights. Ignored when a subfolder is set. + weights_at_root = True, + variant = kwargs.get("variant"), # forward so the warm keeps the variant .bin + gguf_file = kwargs.get( + "gguf_file" + ), # forward so the warm fetches the GGUF (else ignored) + ) + # Child did the forced download; clear the flag so the load reuses the warm cache. + if _prefetched and kwargs.get("force_download", False): + kwargs["force_download"] = False + + # Warm a SEPARATE tokenizer repo only (model_name is covered above). Not model_name here: this + # runs before fast_inference_setup may remap the repo, so it would warm the wrong one. + _tokenizer_repo = ( + tokenizer_name if (isinstance(tokenizer_name, str) and tokenizer_name) else model_name + ) + _warm_tokenizer_repo = ( + isinstance(_tokenizer_repo, str) + and bool(_tokenizer_repo) + and _tokenizer_repo != model_name + ) + if _warm_tokenizer_repo: + maybe_prefetch_hf_snapshot( + _tokenizer_repo, + token = token, + cache_dir = kwargs.get("cache_dir"), + local_files_only = kwargs.get("local_files_only", False), + tokenizer_only = True, + ) + _skip_modules = SKIP_QUANTIZATION_MODULES.copy() # Nemotron-H uses 'mixer' (not 'mamba') for Mamba layers. # Mamba fused kernels pass out_proj.weight directly to F.linear, @@ -1286,6 +1342,18 @@ class FastBaseModel: # Counteract saved tokenizers tokenizer_name = model_name if tokenizer_name is None else tokenizer_name + # On the vLLM path the tokenizer warm was deferred (fast_inference_setup may remap model_name). + # Warm the now-final tokenizer repo so the load below hits the cache (a cached/local repo is a no-op). + if _vllm_owns_weights and isinstance(tokenizer_name, str) and tokenizer_name: + maybe_prefetch_hf_snapshot( + tokenizer_name, + token = token, + revision = kwargs.get("revision"), + cache_dir = kwargs.get("cache_dir"), + local_files_only = kwargs.get("local_files_only", False), + tokenizer_only = True, + ) + # Fix _Unsloth_Patched_ prefix in local config files from old saves (issue #4085) if os.path.isdir(tokenizer_name): import json as _json @@ -1323,6 +1391,7 @@ class FastBaseModel: language = whisper_language, task = whisper_task, trust_remote_code = trust_remote_code, + cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, ) except Exception as _e: @@ -1335,6 +1404,7 @@ class FastBaseModel: padding_side = "left", token = token, trust_remote_code = trust_remote_code, + cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, ) except Exception as _e: @@ -1345,6 +1415,7 @@ class FastBaseModel: padding_side = "left", token = token, trust_remote_code = trust_remote_code, + cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, ) except Exception: @@ -1363,6 +1434,7 @@ class FastBaseModel: model_type_arch, token, trust_remote_code, + cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, ) except Exception as _fe: @@ -1448,6 +1520,7 @@ class FastBaseModel: padding_side = "left", token = token, trust_remote_code = trust_remote_code, + cache_dir = kwargs.get("cache_dir"), local_files_only = local_files_only, ) model, _fallback_tok = patch_tokenizer(model, _fallback_tok) @@ -1477,6 +1550,7 @@ class FastBaseModel: padding_side = "left", token = token, trust_remote_code = trust_remote_code, + cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, ) except Exception: @@ -1486,6 +1560,7 @@ class FastBaseModel: padding_side = "left", token = token, trust_remote_code = trust_remote_code, + cache_dir = kwargs.get("cache_dir"), local_files_only = lfo, ) @@ -1637,6 +1712,16 @@ class FastBaseModel: ) else: _audio_kwargs = {} + # Remember the caller's ORIGINAL explicit leaf list for MoE expert + # detection. When an explicit list is routed through get_peft_regex for + # family scoping below, the generated regex carries get_peft_regex's full + # "mlp|feed_forward|ffn|dense" component block even when the caller named + # only attention leaves (q/k/v/o_proj). Keying expert detection on that + # regex would train the experts for an attention-only request. The + # original list carries the true leaf intent, so use it for MoE detection; + # only the auto (None / "all-linear") path relies on the regex, whose mlp + # block is the sole remaining MLP-intent signal on fused-expert models. + _moe_detect_target = target_modules if type(target_modules) in (list, tuple) else None if target_modules is None or target_modules == "all-linear": target_modules = get_peft_regex( model, @@ -1714,9 +1799,21 @@ class FastBaseModel: loftq_config, lora_dropout, bias, init_lora_weights, model ) - # Auto-detect MoE models and populate target_parameters for expert layers + # Auto-detect MoE models and populate target_parameters for expert layers. + # Prefer the caller's ORIGINAL explicit leaf list over the scoped regex so an + # attention-only request does not train experts via get_peft_regex's mlp block, + # but only when MLP and language families are both still in scope. If the caller + # scoped MLP or language OFF (finetune_mlp_modules / finetune_language_layers + # False), the scoped regex already dropped the experts, so honor it instead of + # re-introducing the original list's gate/up/down leaves. if target_parameters is None: - target_parameters = get_moe_target_parameters(model, target_modules) + _moe_targets = _select_moe_detection_targets( + _moe_detect_target, + target_modules, + finetune_mlp_modules = finetune_mlp_modules, + finetune_language_layers = finetune_language_layers, + ) + target_parameters = get_moe_target_parameters(model, _moe_targets) if finetune_last_n_layers is not None and layers_to_transform is None: _total_layers = _get_total_transformer_layers(model) @@ -2215,3 +2312,16 @@ def check_dataset_for_missing_videos( warnings.warn(error_msg, stacklevel = 2) return missing + + +# Auto-enable grouped-GEMM MoE (transformers<5 ModuleList experts); see llama.py. +try: + from unsloth_zoo.temporary_patches.moe_grouped_modulelist import wrap_loader_for_grouped_moe + FastBaseModel.from_pretrained = staticmethod( + wrap_loader_for_grouped_moe(FastBaseModel.from_pretrained) + ) + FastBaseModel.get_peft_model = staticmethod( + wrap_loader_for_grouped_moe(FastBaseModel.get_peft_model) + ) +except Exception: + pass diff --git a/unsloth/save.py b/unsloth/save.py index a6697e98a1..020c63a9e2 100644 --- a/unsloth/save.py +++ b/unsloth/save.py @@ -2926,15 +2926,16 @@ def unsloth_save_pretrained_gguf( "Unsloth: quantization_method can only be a string or a list of strings" ) for i, quant_method in enumerate(quantization_method): - quant_method = quant_method.lower() + if quant_method is None: + quant_method = "q8_0" + else: + quant_method = quant_method.lower() if quant_method == "not_quantized": quant_method = "f16" elif quant_method == "fast_quantized": quant_method = "q8_0" elif quant_method == "quantized": quant_method = "q4_k_m" - elif quant_method is None: - quant_method = "q8_0" quantization_methods.append(quant_method.lower()) try: @@ -3727,15 +3728,16 @@ def save_to_gguf_generic( "Unsloth: quantization_method can only be a string or a list of strings" ) for i, quant_method in enumerate(quantization_method): - quant_method = quant_method.lower() + if quant_method is None: + quant_method = "q8_0" + else: + quant_method = quant_method.lower() if quant_method == "not_quantized": quant_method = "f16" elif quant_method == "fast_quantized": quant_method = "q8_0" elif quant_method == "quantized": quant_method = "q4_k_m" - elif quant_method is None: - quant_method = "q8_0" new_quantization_methods.append(quant_method.lower()) else: new_quantization_methods.append(quantization_type.lower()) diff --git a/unsloth/tokenizer_utils.py b/unsloth/tokenizer_utils.py index 93dfa9b2ad..3a91ef188d 100644 --- a/unsloth/tokenizer_utils.py +++ b/unsloth/tokenizer_utils.py @@ -563,8 +563,11 @@ def _load_correct_tokenizer( # /tmp of Kaggle seems has a 80GB limit! # Let's utilize them cache_dir = os.path.join(KAGGLE_TMP, cache_dir) - else: + elif cache_dir == "huggingface_tokenizers_cache": + # This default name is Colab/Kaggle-only; elsewhere use the HF default cache. cache_dir = None + # else: keep a caller-supplied cache_dir so the tokenizer loads from the prefetch-warmed dir instead + # of risking an in-process Hub/Xet transfer. # Try loading the slow tokenizer. If it fails, then try Fast only # Mainly to solve Deepseek models with no tokenizer.model file @@ -1323,6 +1326,7 @@ def check_tokenizer( padding_side = "right", token = None, _reload = True, + cache_dir = None, ): # Checks tokenizer for out of bounds ids. # Mainly a fix for https://huggingface.co/berkeley-nest/Starling-LM-7B-alpha @@ -1413,10 +1417,11 @@ def check_tokenizer( f"Fix your tokenizer since it'll perform out of bounds memory accesses." ) - if IS_COLAB_ENVIRONMENT or IS_KAGGLE_ENVIRONMENT: - cache_dir = "huggingface_tokenizers_cache" - else: - cache_dir = None + # Reuse a caller-supplied cache_dir (warmed cache) for the repair reload; else the + # Colab/Kaggle sentinel (HF default elsewhere), as load_correct_tokenizer does. + reload_cache_dir = cache_dir + if reload_cache_dir is None and (IS_COLAB_ENVIRONMENT or IS_KAGGLE_ENVIRONMENT): + reload_cache_dir = "huggingface_tokenizers_cache" # Sometimes slow tokenizer does not work like Deepseek try: @@ -1430,7 +1435,7 @@ def check_tokenizer( use_fast = False, legacy = False, from_slow = True, - cache_dir = cache_dir, + cache_dir = reload_cache_dir, ) return check_tokenizer( model = model, @@ -1440,6 +1445,7 @@ def check_tokenizer( padding_side = padding_side, token = token, _reload = False, + cache_dir = cache_dir, ) break except: diff --git a/unsloth/utils/attention_dispatch.py b/unsloth/utils/attention_dispatch.py index 2e984bad0a..68fb33dad9 100644 --- a/unsloth/utils/attention_dispatch.py +++ b/unsloth/utils/attention_dispatch.py @@ -17,6 +17,7 @@ from __future__ import annotations +import os from dataclasses import dataclass from typing import Any, Optional, Tuple @@ -42,6 +43,17 @@ if HAS_XFORMERS and torch.cuda.is_available(): HAS_XFORMERS = False SDPA_HAS_GQA = "enable_gqa" in (scaled_dot_product_attention.__doc__ or "") +# PrefixGrouper kernel, resolved once when the env gate is on so PG-off users never load +# torch flex_attention. +_flex_shared_prefix_attention = None +if os.environ.get("UNSLOTH_GRPO_PREFIX_GROUPER", "1").lower() not in ("0", "false", "no", "off"): + try: + from .prefix_grouper_kernel import ( + flex_shared_prefix_attention as _flex_shared_prefix_attention, + ) + except Exception: + _flex_shared_prefix_attention = None + FLASH_VARLEN = "flash_varlen" FLASH_DENSE = "flash_dense" XFORMERS = "xformers" @@ -84,6 +96,9 @@ class AttentionContext: attention_mask: Optional[Tensor] causal_mask: Optional[Any] sliding_window: Optional[int] = None + # PrefixGrouper: non-None routes Q/K/V through the FlexAttention shared-prefix kernel; + # None leaves every existing construction/behavior unchanged. + prefix_seg_info: Optional[Any] = None def select_attention_backend(use_varlen: bool = False) -> str: @@ -99,6 +114,33 @@ def select_attention_backend(use_varlen: bool = False) -> str: return SDPA +def resolve_prefix_seg_info(kwargs, past_key_value, attention_mask): + """PrefixGrouper shared-prefix segment table resolver for the arch attention forwards. + + The GRPO PrefixGrouper packed path rides a ``PrefixSegInfo`` in through ``**kwargs`` + (same route as ``packed_seq_lengths``). When present, the forward must route Q/K/V + through the FlexAttention shared-prefix kernel via ``AttentionContext.prefix_seg_info``. + + Returns the seg table (or ``None`` when PrefixGrouper did not group this batch -- the + unchanged path). Hardened: the shared-prefix stream is NOT a plain causal sequence, so running + it under a KV cache or an explicit padding mask would silently produce wrong logprobs. + That combination can only arise from misuse (PrefixGrouper only rides in via the GRPO + logprob forward, which is mask-free prefill), so we RAISE loudly instead of degrading + to a wrong result. + + Factored here so every arch (llama/mistral/qwen3/gemma2/cohere/granite/falcon_h1) + shares one implementation and cannot drift. + """ + seg = kwargs.get("prefix_seg_info", None) + if seg is not None and (past_key_value is not None or attention_mask is not None): + raise RuntimeError( + "PrefixGrouper: prefix_seg_info requires prefill with no KV cache and no " + f"attention_mask (got past_key_value={past_key_value is not None}, " + f"attention_mask={attention_mask is not None})." + ) + return seg + + def run_attention( *, config: AttentionConfig, context: AttentionContext, Q: Tensor, K: Tensor, V: Tensor ) -> Tensor: @@ -111,6 +153,28 @@ def run_attention( and SDPA handle packing via a block-diagonal mask. """ + # PrefixGrouper shared-prefix attention (GRPO dedup). Q/K/V here are [bsz, H, T, D]; + # the kernel takes/returns [1, T, H, D], matching the other backends. The field is + # only set when the env gate is on and grouping succeeded; None keeps every backend + # byte-identical. + if context.prefix_seg_info is not None: + flex_shared_prefix_attention = _flex_shared_prefix_attention + if flex_shared_prefix_attention is None: + # gate flipped on after import (or one-time load failed): resolve lazily. + from ..utils.prefix_grouper_kernel import flex_shared_prefix_attention + + scale = None + if config.flash_varlen_kwargs: + scale = config.flash_varlen_kwargs.get("softmax_scale") + A = flex_shared_prefix_attention( + Q.transpose(1, 2), + K.transpose(1, 2), + V.transpose(1, 2), + context.prefix_seg_info, + scale = scale, + ) + return A # [1, T, n_heads, head_dim] + backend = config.backend if backend == FLASH_VARLEN and context.seq_info is None: backend = FLASH_DENSE if HAS_FLASH_ATTENTION else SDPA @@ -337,5 +401,6 @@ __all__ = [ "AttentionConfig", "AttentionContext", "select_attention_backend", + "resolve_prefix_seg_info", "run_attention", ] diff --git a/unsloth/utils/prefix_grouper.py b/unsloth/utils/prefix_grouper.py new file mode 100644 index 0000000000..4e6ff9672c --- /dev/null +++ b/unsloth/utils/prefix_grouper.py @@ -0,0 +1,351 @@ +# 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 Affero 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 Affero General Public License for more details. +# +# You should have received a copy of the GNU Affero General Public License +# along with this program. If not, see . + +"""PrefixGrouper layout builder + completion-logprob extraction for the Unsloth GRPO +packed path (all archs that route through the varlen attention dispatch). + +Given the de-padded, LEFT-PACKED input_ids the packed GRPO path already works with, this +module: + + 1. Detects consecutive ``num_generations`` rows that share a prompt prefix (byte- + identical prompt precondition; falls back / returns None otherwise). + 2. Builds ONE flat shared-prefix stream across all groups + ``[ prefix_g0, suf_g0_0 .. suf_g0_{G-1}, prefix_g1, ... ]`` with position_ids that + continue each prefix positionally, plus a ``PrefixSegInfo`` segment table for the + FlexAttention shared-prefix kernel. + 3. Extracts completion logprobs via the index map (completion pos ``j==0`` predicted + from the shared prefix's last token; ``j>=1`` from the preceding suffix token) and + scatters them back into ``[total_rows, W]`` EXACTLY where the full-row packed path + puts them (dest = ``orig_row*L + orig_col``), so grpo_compute_loss / completion_mask + / TIS / metrics are byte-untouched. + +The flat stream is built by GATHERING original (row, col) coordinates out of input_ids, +so the grad path's autograd flows to the same embedding rows as today (the shared prefix +now contributes grad once = the sum of the G repeats, which is mathematically identical). + +``chunked_hidden_states_selective_log_softmax`` (from unsloth_zoo, passed in) is reused +verbatim over the gathered predicting-position hidden states, so fp32 accumulation, +logit_scale/softcapping/temperature are all preserved. + +Env: + UNSLOTH_GRPO_PREFIX_GROUPER=1 engage (default ON; set 0 to disable). Auto-off under vLLM. + UNSLOTH_GRPO_PREFIX_GROUPER_TOKR=1.3 tok_r auto-gate threshold (env-overridable) + UNSLOTH_GRPO_PREFIX_GROUPER_VERIFY=1 first-step self-verify (default ON) + UNSLOTH_GRPO_PREFIX_GROUPER_TOL=0.7 self-verify PASS band (nats) +""" + +from __future__ import annotations + +import os +from dataclasses import dataclass +from typing import List, Optional, Tuple + +import torch + +from .prefix_grouper_kernel import build_seg_info_multigroup, PrefixSegInfo + + +# --------------------------------------------------------------------------- +# Env helpers +# --------------------------------------------------------------------------- +def env_on(name: str, default: str = "0") -> bool: + return os.environ.get(name, default).lower() not in ("0", "false", "no", "off") + + +# One-time env reads; the helpers stay callable since unsloth_zoo imports and calls them. +_ENABLED = env_on("UNSLOTH_GRPO_SEQ_PACKING", "1") and env_on("UNSLOTH_GRPO_PREFIX_GROUPER", "1") +_VERIFY_ON = env_on("UNSLOTH_GRPO_PREFIX_GROUPER_VERIFY", "1") +_TOKR_THRESHOLD = float(os.environ.get("UNSLOTH_GRPO_PREFIX_GROUPER_TOKR", "1.3")) +_TOL_OK = float(os.environ.get("UNSLOTH_GRPO_PREFIX_GROUPER_TOL", "0.7")) + + +def prefix_grouper_enabled() -> bool: + """PrefixGrouper requires seq-packing on (it reuses its de-pad + scatter machinery).""" + return _ENABLED + + +def verify_on() -> bool: + return _VERIFY_ON + + +def tokr_threshold() -> float: + return _TOKR_THRESHOLD + + +def tol_ok() -> float: + return _TOL_OK + + +# diff >= TOL_KILL = broken mask/isolation -> structure permanently unsafe; between +# tol_ok and TOL_KILL -> fall back for this shape but keep trying others. +TOL_KILL = 1.5 + + +@dataclass +class GroupLayout: + """Everything the GRPO forward needs to run + extract the shared-prefix path.""" + + flat_ids: torch.Tensor # [1, T] (T == seg.T) + position_ids: torch.Tensor # [1, T] + prefix_seg_info: PrefixSegInfo + # per completion target token, aligned 1:1: + tgt_rows: torch.Tensor # [N] original row index + tgt_cols: torch.Tensor # [N] original padded column in that row + tgt_pred: torch.Tensor # [N] flat predicting index (into the T stream) + tgt_flat: torch.Tensor # [N] flat index of the target token itself (into T) + total_rows: int + L: int # original padded seq length (input_ids.shape[1]) + W: int # logits_to_keep + max_left_pad (scatter width) + tok_r: float + signature: Tuple + + def extract_logps( + self, + hidden, + lm_head, + chunked_fn, + chunks, + logit_scale_multiply, + logit_scale_divide, + logit_softcapping, + temperature, + ) -> torch.Tensor: + """hidden: [1, T, Hdim] (pre-lm_head hidden states, UNSLOTH_RETURN_HIDDEN_STATES=1). + Returns [total_rows, W] float32, byte-compatible with the packed path result.""" + # In a sharded model hidden may live on the lm-head device; move the small index + # maps to hidden.device before indexing. + device = hidden.device + pred_h = hidden[0, self.tgt_pred.to(device), :].unsqueeze(0) # [1, N, Hdim] + tgt_ids = self.flat_ids[0, self.tgt_flat].to(device).unsqueeze(0) # [1, N] + sel = chunked_fn( + pred_h, + lm_head, + tgt_ids, + chunks, + logit_scale_multiply, + logit_scale_divide, + logit_softcapping, + temperature, + )[0] # [N] logprobs + dest = self.tgt_rows.to(device) * self.L + self.tgt_cols.to(device) + result = ( + torch.zeros(self.total_rows * self.L, dtype = torch.float32, device = device) + .index_put((dest,), sel.to(torch.float32)) + .view(self.total_rows, self.L)[:, -self.W :] + ) + return result + + +def _build_groups(ids_cpu, real_cols_cpu, cstart_cpu, num_generations, total_rows): + """CPU-side grouping. Returns group dicts or None. Mirrors the packed _pk_* partition. + + A row's REAL tokens are the columns where input != pad. Its completion region (what + the packed path scatters, then completion_mask masks) is the real columns with + original col >= cstart_r, where cstart_r = (L - logits_to_keep) - left_pad_r. The + prompt is the real columns < cstart_r. Within a GRPO group all G rows share the same + prompt => same left_pad => same cstart => the prompt real columns are BYTE-IDENTICAL + across the group (the shared prefix). We require that byte-identity (falls back + otherwise). No prompt-tail special-casing: every suffix token is scattered exactly + like the packed path; completion_mask masks the leading prompt-tail positions. + """ + G = num_generations + if G is None or G < 2 or total_rows % G != 0: + return None + groups = [] + for g0 in range(0, total_rows, G): + rows = list(range(g0, g0 + G)) + prompt_cols_per_row = [] # real cols < cstart + prompt_toks_per_row = [] + comp_cols_per_row = [] # real cols >= cstart (the completion region packed scatters) + for r in rows: + cs = cstart_cpu[r] + rc = real_cols_cpu[r] + p_cols = [c for c in rc if c < cs] + c_cols = [c for c in rc if c >= cs] + prompt_cols_per_row.append(p_cols) + prompt_toks_per_row.append([ids_cpu[r][c] for c in p_cols]) + comp_cols_per_row.append(c_cols) + if any(len(p) == 0 for p in prompt_toks_per_row): + return None + # require BYTE-IDENTICAL prompts across the group (shared-prefix precondition). + P = len(prompt_toks_per_row[0]) + if any(len(prompt_toks_per_row[k]) != P for k in range(1, G)): + return None + p0 = prompt_toks_per_row[0] + if any(prompt_toks_per_row[k] != p0 for k in range(1, G)): + return None + if P == 0: + return None + R_list = [len(c) for c in comp_cols_per_row] + if sum(R_list) == 0: + return None + groups.append( + dict( + rows = rows, + P = P, + prefix_cols = prompt_cols_per_row[0], # shared prompt real columns (row0) + prefix_row = rows[0], + R_list = R_list, + suf_cols = comp_cols_per_row, # per-row completion-region real columns + ) + ) + return groups + + +def _tok_r(groups) -> float: + tok_full = 0 + tok_sp = 0 + for gm in groups: + P = gm["P"] + Rs = gm["R_list"] + tok_full += sum(P + r for r in Rs) # G*P + sumR + tok_sp += P + sum(Rs) # P + sumR + return (tok_full / tok_sp) if tok_sp else 1.0 + + +def build_group_layout( + input_ids, + logits_to_keep, + pad_id, + num_generations, + left_pad_tokens_per_prompt, + *, + apply_tokr_gate = True, + max_segment_cap = None, +): + """Build the shared-prefix GroupLayout, or return None to fall back to the packed path. + + input_ids : [B, L]. GRPO's layout is left-padded in the prompt and right-padded in + the completion. Real tokens of a row are a contiguous run not necessarily + starting at column 0. + logits_to_keep : int + left_pad_tokens_per_prompt : [B] long tensor (per-row left-pad count in the prompt). + """ + device = input_ids.device + total_rows, L = input_ids.shape + keep = input_ids != pad_id + # completion start column per row (matches create_completion_attention_mask / _pk_cstart). + cstart = ((L - logits_to_keep) - left_pad_tokens_per_prompt).to(torch.long) + cstart_cpu = cstart.tolist() + ids_cpu = input_ids.tolist() + # per-row real (non-pad) columns. GRPO rows are one contiguous real run, so derive + # [first, first+n) on GPU; the O(B*L) scan is only a non-contiguous fallback. + n_real = keep.sum(dim = 1) + first = torch.argmax(keep.to(torch.int8), dim = 1) + ar = torch.arange(L, device = device) + contiguous = bool( + (keep == ((ar >= first.unsqueeze(1)) & (ar < (first + n_real).unsqueeze(1)))).all() + ) + if contiguous: + real_cols_cpu = [list(range(f, f + n)) for f, n in zip(first.tolist(), n_real.tolist())] + else: + keep_cpu = keep.tolist() + real_cols_cpu = [[c for c in range(L) if keep_cpu[r][c]] for r in range(total_rows)] + + groups = _build_groups(ids_cpu, real_cols_cpu, cstart_cpu, num_generations, total_rows) + if groups is None: + return None + + # sliding-window guard: a group's PG span is P + max(R); fall back if it exceeds the window. + if max_segment_cap is not None: + for gm in groups: + if gm["P"] + max(gm["R_list"]) > max_segment_cap: + return None + + tok_r = _tok_r(groups) + if apply_tokr_gate and tok_r < tokr_threshold(): + return None # low reuse -> not worth it; use the full-row packed path + + # Build flat stream by gathering original (row, col) coordinates. + group_specs = [(gm["P"], gm["R_list"]) for gm in groups] + seg, group_meta = build_seg_info_multigroup(group_specs, device) + + flat_src_rows: List[int] = [] + flat_src_cols: List[int] = [] + pos_list: List[int] = [] + tgt_rows: List[int] = [] + tgt_cols: List[int] = [] + tgt_pred: List[int] = [] + tgt_flat: List[int] = [] + + for gm, meta in zip(groups, group_meta): + rows = gm["rows"] + P = gm["P"] + r0 = gm["prefix_row"] + prefix_cols = gm["prefix_cols"] # ORIGINAL real prompt columns (len P) of row0 + plast = meta["prefix_last_index"] # base + P - 1 + # gather the shared prefix once, from row0. + flat_src_rows.extend([r0] * P) + flat_src_cols.extend(prefix_cols) + pos_list.extend(range(P)) + # suffixes: every suffix token is a completion-region target (scattered like the + # packed path; completion_mask hides prompt-tail positions). + for i, r in enumerate(rows): + cols = gm["suf_cols"][i] + r_i = len(cols) + s, e = meta["suffix_slices"][i] # flat offsets [s, e) + flat_src_rows.extend([r] * r_i) + flat_src_cols.extend(cols) + pos_list.extend(range(P, P + r_i)) + for j in range(r_i): + # pos 0 is predicted from the prefix's last token; j>=1 from the previous suffix token. + pred = plast if j == 0 else (s + j - 1) + tgt_rows.append(r) + tgt_cols.append(cols[j]) # ORIGINAL padded column in row r + tgt_pred.append(pred) + tgt_flat.append(s + j) # flat index of the target token itself + + T = len(flat_src_rows) + assert T == seg.T, f"flat stream len {T} != seg.T {seg.T}" + fr = torch.tensor(flat_src_rows, device = device, dtype = torch.long) + fc = torch.tensor(flat_src_cols, device = device, dtype = torch.long) + flat_ids = input_ids[fr, fc].unsqueeze(0) # [1, T] (grad-safe gather) + position_ids = torch.tensor(pos_list, device = device, dtype = torch.long).unsqueeze(0) + + max_left_pad = int(left_pad_tokens_per_prompt.max().item()) if total_rows else 0 + W = logits_to_keep + max_left_pad + + # self-verify cache key: the mask/index-map/scatter logic is structural, so key on + # (num_groups, group_sizes), not exact lengths -- GRPO lengths change every step and + # keying on T would re-verify forever ("verify once, then trust", like the packed path). + grp_sizes = tuple(sorted(len(gm["R_list"]) for gm in groups)) + sig = (len(groups), grp_sizes) + + return GroupLayout( + flat_ids = flat_ids, + position_ids = position_ids, + prefix_seg_info = seg, + tgt_rows = torch.tensor(tgt_rows, device = device, dtype = torch.long), + tgt_cols = torch.tensor(tgt_cols, device = device, dtype = torch.long), + tgt_pred = torch.tensor(tgt_pred, device = device, dtype = torch.long), + tgt_flat = torch.tensor(tgt_flat, device = device, dtype = torch.long), + total_rows = total_rows, + L = L, + W = W, + tok_r = tok_r, + signature = sig, + ) + + +__all__ = [ + "GroupLayout", + "build_group_layout", + "prefix_grouper_enabled", + "verify_on", + "tokr_threshold", + "tol_ok", + "TOL_KILL", + "env_on", +] diff --git a/unsloth/utils/prefix_grouper_kernel.py b/unsloth/utils/prefix_grouper_kernel.py new file mode 100644 index 0000000000..9a9719b015 --- /dev/null +++ b/unsloth/utils/prefix_grouper_kernel.py @@ -0,0 +1,436 @@ +# 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 Affero 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 Affero General Public License for more details. +# +# You should have received a copy of the GNU Affero General Public License +# along with this program. If not, see . + +"""FlexAttention shared-prefix kernel for PrefixGrouper (GRPO shared-prompt dedup). + +In GRPO every prompt spawns ``G = num_generations`` completions that share the same +prompt prefix. The full-row packed path forwards the identical prefix ``G`` times. +PrefixGrouper stores the prefix ONCE and concatenates only the ``G`` suffixes, with an +attention layout where each suffix token attends to ``[the single shared prefix] + +[causal within its own suffix]``. This kernel expresses that one-prefix -> many-suffix +fan-out via a ``torch.nn.attention.flex_attention`` block mask, so the masked-out +cross-suffix / cross-group blocks are never computed and the ``P + G*R`` FLOP saving is +realised (not merely a masked dense ``O(T^2)``). + +Mask semantics (identical to the certified SDPA oracle): + + keep(q_idx, kv_idx) = same_group(q, kv) AND + ( is_prefix[kv_idx] # full prefix visibility + OR ( suffix_of_kv[kv_idx] == suffix_of_kv[q_idx] # same suffix ... + AND kv_idx <= q_idx ) ) # ... causal within it + +This module is self-contained (no dependency on any temp/ scratch dir) so PrefixGrouper +works from the installed source after a fresh compile. It is only imported lazily from +``attention_dispatch.run_attention`` when ``prefix_seg_info`` is present, which itself is +only ever set when ``UNSLOTH_GRPO_PREFIX_GROUPER`` is on and grouping succeeded, so the +default (off) path never touches this file. + +Provided entry points: + * ``PrefixSegInfo`` : per-flat-token segment metadata + cache signature. + * ``build_seg_info_multigroup``: build PrefixSegInfo for many groups packed flat. + * ``build_seg_info_from_layout``: build PrefixSegInfo for ONE group (test helper). + * ``get_block_mask`` : cached create_block_mask keyed on the signature. + * ``flex_shared_prefix_attention(Q, K, V, prefix_seg_info)`` + Q/K/V of shape [1, T, n_heads, head_dim]; returns [1, T, n_heads, head_dim], + IDENTICAL semantics to the SDPA oracle. +""" + +from __future__ import annotations + +import os +from dataclasses import dataclass +from typing import Dict, List, Optional, Tuple + +import torch +from torch.nn.attention.flex_attention import ( + BlockMask, + create_block_mask, + flex_attention, +) + +# GRPO feeds many distinct segment lengths; at dynamo's default recompile_limit (8) the +# compiled kernel silently reuses a mismatched specialisation (wrong results). Raise it. +torch._dynamo.config.recompile_limit = max(getattr(torch._dynamo.config, "recompile_limit", 8), 256) +torch._dynamo.config.accumulated_recompile_limit = max( + getattr(torch._dynamo.config, "accumulated_recompile_limit", 256), 2048 +) + + +# Compiled kernels: torch.compile fuses the sparse mask into one kernel. dynamic=True is +# required: T changes almost every GRPO batch and dynamic=False recompiles per T (~14s +# each). T is still padded to a multiple of 128 (_pad_len) for the backward kernel. +_flex_attention_compiled = torch.compile(flex_attention, dynamic = True) +_create_block_mask_compiled = torch.compile(create_block_mask, dynamic = True) + +# Flash block sizes by Q dtype (env-overridable). The two disjoint key runs (prefix + +# own-suffix) stress online-softmax accumulation: fp32 needs 32/32 for a ~1e-6 floor; +# bf16 passes parity at 128/64 and is ~5x faster (128/128 OOMs Triton on B200). +_FP32_BLOCK_M = int(os.environ.get("PG_FLEX_BLOCK_M", "32")) +_FP32_BLOCK_N = int(os.environ.get("PG_FLEX_BLOCK_N", "32")) +_BF16_BLOCK_M = int(os.environ.get("PG_FLEX_BF16_BLOCK_M", "128")) +_BF16_BLOCK_N = int(os.environ.get("PG_FLEX_BF16_BLOCK_N", "64")) + + +def _kernel_options_for_dtype(dtype): + """Pick the numerically-safe flash block sizes for the Q dtype.""" + if dtype == torch.bfloat16 or dtype == torch.float16: + return {"BLOCK_M": _BF16_BLOCK_M, "BLOCK_N": _BF16_BLOCK_N} + return {"BLOCK_M": _FP32_BLOCK_M, "BLOCK_N": _FP32_BLOCK_N} + + +# Backward-compat constant (fp32 default). +_FLEX_KERNEL_OPTIONS = {"BLOCK_M": _FP32_BLOCK_M, "BLOCK_N": _FP32_BLOCK_N} + +# The compiled backward trips an Inductor assertion when T is not a multiple of 128, so +# pad the flat sequence. Pad tokens form a group that attends to / is attended by nothing +# (all-masked rows return 0, not NaN) and are sliced off the output. +_PAD_MULTIPLE = 128 +_PAD_GROUP = -99 # sentinel group id / suffix id for pad tokens + + +def _pad_len(T: int) -> int: + return ((T + _PAD_MULTIPLE - 1) // _PAD_MULTIPLE) * _PAD_MULTIPLE + + +# --------------------------------------------------------------------------- +# Segment metadata +# --------------------------------------------------------------------------- + + +@dataclass +class PrefixSegInfo: + """Per-flat-token segment metadata driving the shared-prefix block mask. + + The label tensors are 1-D of length ``T_pad`` (>= real ``T``, padded up to a multiple + of 128 so the backward kernel compiles). Positions ``[T:T_pad)`` are pad tokens + (group/suffix == _PAD_GROUP) that attend to nothing. + + Attributes + ---------- + group_of_kv : LongTensor [T_pad] + Group id per flat token (0..num_groups-1); _PAD_GROUP for pad tokens. + is_prefix : BoolTensor [T_pad] + True iff the token is a prefix token of its group (False for pad). + suffix_of_kv : LongTensor [T_pad] + Suffix id per flat token; -1 for prefix, _PAD_GROUP for pad. Suffix ids are + globally unique across groups. + signature : hashable + Cache key for the block mask (depends only on the labels + T_pad). + T : int + Real flat sequence length (Q/K/V of this length are padded internally). + T_pad : int + Padded length (multiple of 128) at which the block mask is built. + """ + + group_of_kv: torch.Tensor + is_prefix: torch.Tensor + suffix_of_kv: torch.Tensor + signature: Tuple + T: int + T_pad: int + + +def _pad_labels(group_of_kv, is_prefix, suffix_of_kv, device): + """Pad the label tensors up to a multiple of 128 with pad-token sentinels.""" + T = int(group_of_kv.numel()) + T_pad = _pad_len(T) + if T_pad == T: + return group_of_kv, is_prefix, suffix_of_kv, T, T_pad + pad = T_pad - T + group_of_kv = torch.cat( + [group_of_kv, torch.full((pad,), _PAD_GROUP, dtype = torch.long, device = device)] + ) + is_prefix = torch.cat([is_prefix, torch.zeros(pad, dtype = torch.bool, device = device)]) + suffix_of_kv = torch.cat( + [suffix_of_kv, torch.full((pad,), _PAD_GROUP, dtype = torch.long, device = device)] + ) + return group_of_kv, is_prefix, suffix_of_kv, T, T_pad + + +def build_seg_info_from_layout(layout, device: Optional[torch.device] = None) -> PrefixSegInfo: + """Build PrefixSegInfo for ONE group from an object with ``.flat_ids``, ``.P`` and + ``.suffix_slices`` (used by the parity test / oracle helpers).""" + if device is None: + device = layout.flat_ids.device + T = int(layout.flat_ids.shape[1]) + P = int(layout.P) + + group_of_kv = torch.zeros(T, dtype = torch.long, device = device) # single group -> 0 + is_prefix = torch.zeros(T, dtype = torch.bool, device = device) + is_prefix[:P] = True + suffix_of_kv = torch.full((T,), -1, dtype = torch.long, device = device) + for i, (s, e) in enumerate(layout.suffix_slices): + suffix_of_kv[s:e] = i + + group_of_kv, is_prefix, suffix_of_kv, T, T_pad = _pad_labels( + group_of_kv, is_prefix, suffix_of_kv, device + ) + sig = ("single", T_pad, P, tuple((s, e) for (s, e) in layout.suffix_slices)) + return PrefixSegInfo( + group_of_kv = group_of_kv, + is_prefix = is_prefix, + suffix_of_kv = suffix_of_kv, + signature = sig, + T = T, + T_pad = T_pad, + ) + + +def build_seg_info_multigroup( + group_specs: List[Tuple[int, List[int]]], device: torch.device +) -> Tuple[PrefixSegInfo, List[dict]]: + """Build PrefixSegInfo for several shared-prefix groups packed block-diagonally. + + Parameters + ---------- + group_specs : list of (P_g, [R_{g,0}, R_{g,1}, ...]) + For each group: prefix length and the list of suffix lengths. + + Returns + ------- + seg : PrefixSegInfo + group_meta : list of dicts with 'base', 'P', 'prefix_last_index', 'suffix_slices' + (flat offsets), enough to build the completion index map. + """ + group_of_list = [] + is_prefix_list = [] + suffix_of_list = [] + group_meta = [] + + base = 0 + suffix_counter = 0 + sig_parts = [] + for gid, (P, R_list) in enumerate(group_specs): + # prefix + group_of_list.append(torch.full((P,), gid, dtype = torch.long, device = device)) + is_prefix_list.append(torch.ones(P, dtype = torch.bool, device = device)) + suffix_of_list.append(torch.full((P,), -1, dtype = torch.long, device = device)) + prefix_last_index = base + P - 1 + suffix_slices = [] + cursor = base + P + for r in R_list: + group_of_list.append(torch.full((r,), gid, dtype = torch.long, device = device)) + is_prefix_list.append(torch.zeros(r, dtype = torch.bool, device = device)) + suffix_of_list.append(torch.full((r,), suffix_counter, dtype = torch.long, device = device)) + suffix_slices.append((cursor, cursor + r)) + cursor += r + suffix_counter += 1 + group_meta.append( + { + "base": base, + "P": P, + "prefix_last_index": prefix_last_index, + "suffix_slices": suffix_slices, + } + ) + sig_parts.append((P, tuple(R_list))) + base = cursor + + group_of_kv = torch.cat(group_of_list) + is_prefix = torch.cat(is_prefix_list) + suffix_of_kv = torch.cat(suffix_of_list) + group_of_kv, is_prefix, suffix_of_kv, T, T_pad = _pad_labels( + group_of_kv, is_prefix, suffix_of_kv, device + ) + sig = ("multi", T_pad, tuple(sig_parts)) + seg = PrefixSegInfo( + group_of_kv = group_of_kv, + is_prefix = is_prefix, + suffix_of_kv = suffix_of_kv, + signature = sig, + T = T, + T_pad = T_pad, + ) + return seg, group_meta + + +# --------------------------------------------------------------------------- +# Block-mask builder + cache, keyed on (signature, device): the mask depends only on the +# per-token labels and T, so it is reused across layers and steps. + +_BLOCK_MASK_CACHE: Dict[Tuple, BlockMask] = {} + + +def _make_mask_mod(group_of_kv, is_prefix, suffix_of_kv): + """Return a mask_mod closure over the (device) label tensors. + + keep(q, kv) = same_group AND + ( is_prefix[kv] AND kv <= q # causal within/ into prefix + OR ( suffix_of_kv[kv] == suffix_of_kv[q] # same suffix ... + AND (not is_prefix[q]) # q is a suffix token ... + AND kv <= q ) ) # ... causal within it + + The single ``kv <= q`` guard on the is_prefix branch gives BOTH prefix-causal + behaviour (a prefix q sees only earlier prefix tokens) AND full-prefix-visibility for + suffixes (every prefix index < every suffix index in a group, so kv <= q always holds + for a suffix q vs a prefix kv of its group), matching the SDPA oracle exactly. + """ + + def mask_mod(b, h, q_idx, kv_idx): + same_group = group_of_kv[q_idx] == group_of_kv[kv_idx] + kv_is_prefix = is_prefix[kv_idx] + causal = kv_idx <= q_idx + same_suffix = (suffix_of_kv[kv_idx] == suffix_of_kv[q_idx]) & (~is_prefix[q_idx]) + keep = same_group & ((kv_is_prefix & causal) | (same_suffix & causal)) + return keep + + return mask_mod + + +def get_block_mask( + seg: PrefixSegInfo, + device: torch.device, + compile_mask: bool = True, +) -> BlockMask: + """Return a cached BlockMask for the segment signature (built once, reused). + + CRITICAL: the block mask is cached and shared across BOTH the no-grad old/ref logprob + forward (which runs under torch.inference_mode) and the grad training forward. If the + mask were first built under inference_mode, its tensors would be INFERENCE tensors that + "cannot be saved for backward" when reused in the grad forward. We therefore build the + mask with inference mode explicitly DISABLED, so the same cached BlockMask is a normal + tensor usable by autograd. (The mask depends only on integer labels; it needs no grad.) + """ + key = (seg.signature, str(device)) + bm = _BLOCK_MASK_CACHE.get(key) + if bm is not None: + return bm + + # Move labels to the consumer (Q) device: with a sharded model the seg tensors live on + # input_ids.device and would index cross-device. Copies once per (signature, device). + # These copies must also run with inference mode DISABLED (same reason as the mask build): + # when this entry is first built under the no-grad old/ref forward's inference_mode and + # device != seg.device, a .to(device) copy would be an inference tensor that mask_mod + # captures, which then cannot be saved for backward when the grad training forward reuses + # the cached mask. + builder = _create_block_mask_compiled if compile_mask else create_block_mask + with torch.inference_mode(False): + mask_mod = _make_mask_mod( + seg.group_of_kv.to(device), seg.is_prefix.to(device), seg.suffix_of_kv.to(device) + ) + bm = builder( + mask_mod, + B = 1, + H = None, + Q_LEN = seg.T_pad, + KV_LEN = seg.T_pad, + device = device, + ) + # FIFO bound: GRPO lengths change nearly every step, so evict the oldest to cap GPU pins. + if len(_BLOCK_MASK_CACHE) >= 8: + _BLOCK_MASK_CACHE.pop(next(iter(_BLOCK_MASK_CACHE))) + _BLOCK_MASK_CACHE[key] = bm + return bm + + +def clear_block_mask_cache(): + _BLOCK_MASK_CACHE.clear() + + +def _pad_qkv_seq(x: torch.Tensor, T_pad: int) -> torch.Tensor: + """Zero-pad a [B, H, T, D] tensor along the sequence dim up to T_pad.""" + T = x.shape[2] + if T_pad == T: + return x + pad = torch.zeros(x.shape[0], x.shape[1], T_pad - T, x.shape[3], device = x.device, dtype = x.dtype) + return torch.cat([x, pad], dim = 2) + + +def _run_flex(q, k, v, block_mask, enable_gqa, scale, compiled, T, T_pad): + """Pad q/k/v to T_pad, run flex, slice the output back to T. q/k/v: [B,H,T,D].""" + qp = _pad_qkv_seq(q, T_pad) + kp = _pad_qkv_seq(k, T_pad) + vp = _pad_qkv_seq(v, T_pad) + if compiled: + out = _flex_attention_compiled( + qp, + kp, + vp, + block_mask = block_mask, + enable_gqa = enable_gqa, + scale = scale, + kernel_options = _kernel_options_for_dtype(qp.dtype), + ) + else: + # eager path (fp64 parity): dense scores, no kernel_options. + out = flex_attention( + qp, + kp, + vp, + block_mask = block_mask, + enable_gqa = enable_gqa, + scale = scale, + ) + return out[:, :, :T, :] + + +# --------------------------------------------------------------------------- +# The kernel entry point +# --------------------------------------------------------------------------- + + +def flex_shared_prefix_attention( + Q: torch.Tensor, + K: torch.Tensor, + V: torch.Tensor, + prefix_seg_info: PrefixSegInfo, + scale: Optional[float] = None, + block_mask: Optional[BlockMask] = None, + compiled: bool = True, +) -> torch.Tensor: + """Shared-prefix attention via FlexAttention. + + Parameters + ---------- + Q, K, V : Tensor [1, T, n_heads, head_dim] + (Q has n_heads, K/V have n_kv_heads for GQA). + prefix_seg_info : PrefixSegInfo + scale : optional float, softmax scale (defaults to 1/sqrt(head_dim)). + block_mask : optional precomputed BlockMask (else built/cached from seg info). + + Returns + ------- + Tensor [1, T, n_heads, head_dim], identical semantics to the SDPA oracle branch. + """ + assert Q.dim() == 4 and Q.shape[0] == 1, f"expected [1,T,H,D], got {tuple(Q.shape)}" + device = Q.device + # FlexAttention wants [B, H, T, D]. + q = Q.transpose(1, 2) # [1, n_heads, T, D] + k = K.transpose(1, 2) # [1, n_kv_heads, T, D] + v = V.transpose(1, 2) + + n_heads = q.shape[1] + n_kv = k.shape[1] + enable_gqa = n_heads != n_kv + T = q.shape[2] + T_pad = prefix_seg_info.T_pad + assert T == prefix_seg_info.T, f"Q length {T} != seg.T {prefix_seg_info.T}" + + if block_mask is None: + block_mask = get_block_mask(prefix_seg_info, device, compile_mask = compiled) + + out = _run_flex(q, k, v, block_mask, enable_gqa, scale, compiled, T, T_pad) + # back to [1, T, n_heads, D] + return out.transpose(1, 2).contiguous() + + +__all__ = [ + "PrefixSegInfo", + "build_seg_info_multigroup", + "build_seg_info_from_layout", + "get_block_mask", + "clear_block_mask_cache", + "flex_shared_prefix_attention", +]