Merge remote-tracking branch 'origin/main' into docker-blackwell-build

Sync the docker image branch with main (15 commits) so the branch's Python
files carry main's current formatting and pre-commit.ci runs cleanly on a
non-stale checkout.
This commit is contained in:
Daniel Han 2026-07-06 14:20:09 +00:00
commit 782a7d5335
33 changed files with 4495 additions and 806 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

123
tests/test_fp8_tiny_e8m0.py Normal file
View file

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

View file

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

View file

@ -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 <https://www.gnu.org/licenses/>.
"""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 <subfolder>/<gguf_file>, 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}"
)

View file

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

View file

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

View file

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

View file

@ -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="<name>" fetches additional_chat_templates/<name>.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.<variant>.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 <subfolder>/ 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"(?<![qkvo]_)(?<!out_)(?<!in_)proj", target_modules):
return set(_MOE_BROAD_MLP_TARGETS)
# The auto regex on fused-expert models lists only attention Linears as
# leaves; its mlp tag block is the remaining MLP-intent signal. A regex
# like "(mlp|self_attn).(q_proj|o_proj)" has neither and stays attention-only.
if "mlp|feed_forward|ffn|dense" in target_modules:
return set(_MOE_BROAD_MLP_TARGETS)
return set()
@ -3598,6 +4080,31 @@ def get_moe_target_parameters(model, target_modules = None) -> 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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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 <https://www.gnu.org/licenses/>.
"""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",
]

View file

@ -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 <https://www.gnu.org/licenses/>.
"""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",
]