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",
+]