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:
commit
782a7d5335
33 changed files with 4495 additions and 806 deletions
1
.github/workflows/consolidated-tests-ci.yml
vendored
1
.github/workflows/consolidated-tests-ci.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
4
.github/workflows/lockfile-audit.yml
vendored
4
.github/workflows/lockfile-audit.yml
vendored
|
|
@ -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'
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ──
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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) => {
|
||||
|
|
|
|||
|
|
@ -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({
|
||||
|
|
|
|||
81
tests/saving/test_quant_method_none_normalization.py
Normal file
81
tests/saving/test_quant_method_none_normalization.py
Normal 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"]
|
||||
190
tests/test_attn_impl_honor_explicit.py
Normal file
190
tests/test_attn_impl_honor_explicit.py
Normal 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
123
tests/test_fp8_tiny_e8m0.py
Normal 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"]))
|
||||
|
|
@ -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
|
||||
|
|
|
|||
916
tests/test_prefetch_snapshot_scope.py
Normal file
916
tests/test_prefetch_snapshot_scope.py
Normal 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}"
|
||||
)
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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!")
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
351
unsloth/utils/prefix_grouper.py
Normal file
351
unsloth/utils/prefix_grouper.py
Normal 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",
|
||||
]
|
||||
436
unsloth/utils/prefix_grouper_kernel.py
Normal file
436
unsloth/utils/prefix_grouper_kernel.py
Normal 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",
|
||||
]
|
||||
Loading…
Add table
Add a link
Reference in a new issue