Merge remote-tracking branch 'origin/main' into cuda-torch-index-override
This commit is contained in:
commit
5d69193e1e
13 changed files with 2115 additions and 532 deletions
|
|
@ -84,7 +84,7 @@
|
|||
"id": "277e431e"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": "import sys\nsys.path.insert(0, \"/content/unsloth/studio/backend\")\nfrom colab import start\nstart()"
|
||||
"source": "import sys\nsys.path.insert(0, \"/content/unsloth/studio/backend\")\nfrom colab import start\n\n# Default: in-tab iframe only. start() blocks to keep the kernel alive.\nstart()\n\n# For a shareable Cloudflare link, replace start() above with:\n# start(cloudflare=True)"
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
|
|
|
|||
|
|
@ -103,24 +103,132 @@ def show_link(port: int = 8888, *, _url: "str | None" = None):
|
|||
display(HTML(html))
|
||||
|
||||
|
||||
def _is_studio_healthy(port: int, timeout: float = 2.0) -> bool:
|
||||
"""Return True if a Studio backend is already answering health checks on *port*."""
|
||||
import urllib.request
|
||||
def _bootstrap_password_pending() -> bool:
|
||||
"""True while the default admin still owes a bootstrap-password change.
|
||||
|
||||
While pending, main.py injects that password into same-origin GETs, and a public
|
||||
tunnel GET (no Origin) reads as same-origin, so sharing the link would leak admin
|
||||
access. Fails safe to pending if the state cannot be read.
|
||||
"""
|
||||
try:
|
||||
with urllib.request.urlopen(f"http://localhost:{port}/api/health", timeout = timeout):
|
||||
return True
|
||||
from auth.storage import requires_password_change, DEFAULT_ADMIN_USERNAME
|
||||
return bool(requires_password_change(DEFAULT_ADMIN_USERNAME))
|
||||
except Exception as e:
|
||||
logger.info(f"Could not check admin password state ({e}); refusing tunnel to be safe.")
|
||||
return True
|
||||
|
||||
|
||||
def start_cloudflare_tunnel(port: int) -> "str | None":
|
||||
"""Open a shareable Cloudflare quick tunnel to localhost:*port*, or None.
|
||||
|
||||
run_server suppresses the tunnel on Colab by design, so we start it directly.
|
||||
Refused while the bootstrap password is pending; any failure collapses to None
|
||||
and the Colab proxy still works.
|
||||
"""
|
||||
if _bootstrap_password_pending():
|
||||
logger.warning(
|
||||
"Cloudflare link not started: the admin account still has its temporary "
|
||||
"bootstrap password, which is exposed to anyone who can load the page. "
|
||||
"Open Studio in this tab, log in and change the admin password, then re-run "
|
||||
"start(cloudflare=True) to get the shareable link."
|
||||
)
|
||||
return None
|
||||
try:
|
||||
from cloudflare_tunnel import start_studio_tunnel
|
||||
except Exception as e:
|
||||
logger.info(f"Cloudflare tunnel unavailable ({e}); using Colab proxy only.")
|
||||
return None
|
||||
try:
|
||||
url = start_studio_tunnel(port)
|
||||
except Exception as e:
|
||||
logger.info(f"Cloudflare tunnel failed to start ({e}); using Colab proxy only.")
|
||||
return None
|
||||
# Success is logged by _show_and_embed; note only misses here.
|
||||
if not url:
|
||||
logger.info("Cloudflare tunnel did not produce a URL; using Colab proxy only.")
|
||||
return url
|
||||
|
||||
|
||||
def _publish_cloudflare_url(cloudflare_url: "str | None") -> None:
|
||||
"""Publish a directly-started tunnel URL onto app.state so /api/health advertises it.
|
||||
|
||||
run_server only sets this when it opens the tunnel itself, which it skips on Colab,
|
||||
so we set it here. Otherwise the frontend's API examples fall back to an
|
||||
unreachable server_url. Best-effort.
|
||||
"""
|
||||
if not cloudflare_url:
|
||||
return
|
||||
try:
|
||||
from main import app as _studio_app
|
||||
_studio_app.state.cloudflare_url = cloudflare_url
|
||||
except Exception as e:
|
||||
logger.info(f"Could not publish Cloudflare URL to /api/health ({e}).")
|
||||
|
||||
|
||||
def _stop_cloudflare_tunnel() -> None:
|
||||
"""Best-effort teardown of the Cloudflare tunnel started by start_cloudflare_tunnel."""
|
||||
try:
|
||||
from cloudflare_tunnel import stop_studio_tunnel
|
||||
stop_studio_tunnel()
|
||||
except Exception:
|
||||
pass
|
||||
# Stop /api/health advertising a dead tunnel.
|
||||
try:
|
||||
from main import app as _studio_app
|
||||
_studio_app.state.cloudflare_url = None
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _is_studio_healthy(port: int, timeout: float = 2.0) -> bool:
|
||||
"""True only if Unsloth Studio (not some other app) answers /api/health on *port*.
|
||||
|
||||
The service-marker check stops the reuse path reusing or tunneling a foreign
|
||||
process that merely serves /api/health.
|
||||
"""
|
||||
import json, urllib.request
|
||||
try:
|
||||
with urllib.request.urlopen(f"http://localhost:{port}/api/health", timeout = timeout) as r:
|
||||
return json.loads(r.read()).get("service") == "Unsloth UI Backend"
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _show_and_embed(port: int):
|
||||
"""Embed the Studio inline for *port* with a branded header bar.
|
||||
|
||||
Fetches the proxy URL once (registering the port), then renders header bar +
|
||||
iframe. Falls back to serve_kernel_port_as_iframe if IPython HTML is unavailable.
|
||||
def _shareable_link_html(cloudflare_url: str) -> str:
|
||||
"""Branded card for the shareable Cloudflare link, styled like the show_link banner."""
|
||||
return f"""
|
||||
<div style="display: inline-block; padding: 20px; background: #ffffff; border: 2px solid #000000;
|
||||
border-radius: 12px; margin: 10px 0; font-family: system-ui, -apple-system, sans-serif;">
|
||||
<h2 style="color: #000000; margin: 0 0 12px 0; font-size: 26px; font-weight: 800;
|
||||
display: flex; align-items: center; gap: 12px;">
|
||||
<img src="https://github.com/unslothai/unsloth/raw/main/studio/frontend/public/unsloth-gem.png"
|
||||
height="48" style="display:block;">
|
||||
Shareable Studio Link is Ready!
|
||||
</h2>
|
||||
<a href="{cloudflare_url}" onclick="var w=window.open(this.href,'_blank');if(!w){{return true;}}return false;"
|
||||
style="display: inline-flex; align-items: center; gap: 10px; padding: 14px 28px;
|
||||
background: #000000; color: white; text-decoration: none; border-radius: 8px;
|
||||
font-weight: 800; font-size: 16px; cursor: pointer;">
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="18" height="18" viewBox="0 0 24 24" fill="white"><polygon points="5,3 19,12 5,21"/></svg>
|
||||
Open Unsloth Studio
|
||||
</a>
|
||||
<p style="color: #333333; margin: 12px 0 0 0; font-size: 14px; font-weight: bold;">
|
||||
This Cloudflare HTTPS link works from any device — share it with anyone. The Colab view below only works in this tab.
|
||||
</p>
|
||||
<p style="color: #333333; margin: 16px 0 0 0; font-size: 13px; font-family: monospace; font-weight: bold;">
|
||||
🔗 {cloudflare_url}
|
||||
</p>
|
||||
</div>
|
||||
"""
|
||||
|
||||
|
||||
def _show_and_embed(port: int, *, cloudflare_url: "str | None" = None):
|
||||
"""Render the Studio header + iframe for *port*, with a shareable-link card above
|
||||
when *cloudflare_url* is set. Falls back to serve_kernel_port_as_iframe."""
|
||||
url = get_colab_url(port)
|
||||
logger.info(f"🌐 Unsloth Studio URL: {url}")
|
||||
if cloudflare_url:
|
||||
logger.info(f"🔗 Shareable Cloudflare link: {cloudflare_url}")
|
||||
|
||||
try:
|
||||
from IPython.display import HTML, display
|
||||
|
|
@ -136,6 +244,9 @@ def _show_and_embed(port: int):
|
|||
except (ValueError, IndexError):
|
||||
short_url = url
|
||||
|
||||
if cloudflare_url:
|
||||
display(HTML(_shareable_link_html(cloudflare_url)))
|
||||
|
||||
display(
|
||||
HTML(f"""
|
||||
<div style="font-family:system-ui,-apple-system,sans-serif;margin:8px 0;
|
||||
|
|
@ -164,13 +275,18 @@ def _show_and_embed(port: int):
|
|||
pass
|
||||
|
||||
|
||||
def start(port: int = 8888):
|
||||
"""
|
||||
Start Unsloth Studio server in Colab and display the URL.
|
||||
def start(port: int = 8888, *, cloudflare: bool = False):
|
||||
"""Start Unsloth Studio in Colab and display the URL.
|
||||
|
||||
Args:
|
||||
port: Port to bind/serve on.
|
||||
cloudflare: Opt in to a shareable Cloudflare HTTPS link reachable from any
|
||||
device (default OFF). It exposes Studio's login page beyond Colab, so it
|
||||
stays an explicit opt-in; the default shows only the in-tab proxy iframe.
|
||||
|
||||
Usage:
|
||||
from colab import start
|
||||
start()
|
||||
start() # Colab-proxy iframe only (default)
|
||||
start(cloudflare=True) # also open a shareable Cloudflare link
|
||||
"""
|
||||
import time
|
||||
|
||||
|
|
@ -180,13 +296,18 @@ def start(port: int = 8888):
|
|||
# the port, so just re-show the link and iframe.
|
||||
if _is_studio_healthy(port):
|
||||
logger.info(f" Studio is already running on port {port} — reusing existing server.")
|
||||
_show_and_embed(port)
|
||||
# try/finally: tear the tunnel down even if interrupted mid-start/render.
|
||||
try:
|
||||
cf_url = start_cloudflare_tunnel(port) if cloudflare else None
|
||||
_publish_cloudflare_url(cf_url)
|
||||
_show_and_embed(port, cloudflare_url = cf_url)
|
||||
for _ in range(10000):
|
||||
time.sleep(300)
|
||||
print("=", end = "", flush = True)
|
||||
except KeyboardInterrupt:
|
||||
logger.info("\nUnsloth Studio keepalive stopped.")
|
||||
finally:
|
||||
_stop_cloudflare_tunnel()
|
||||
return
|
||||
|
||||
logger.info(" Loading backend...")
|
||||
|
|
@ -202,7 +323,15 @@ def start(port: int = 8888):
|
|||
|
||||
logger.info(" Starting server...")
|
||||
try:
|
||||
app = run_server(host = "0.0.0.0", port = port, frontend_path = frontend_path, silent = True)
|
||||
# cloudflare=False: this helper owns the tunnel. run_server's default True
|
||||
# would tunnel this 0.0.0.0 bind if Colab detection fails, breaking the opt-out.
|
||||
app = run_server(
|
||||
host = "0.0.0.0",
|
||||
port = port,
|
||||
frontend_path = frontend_path,
|
||||
silent = True,
|
||||
cloudflare = False,
|
||||
)
|
||||
except SystemExit as exc:
|
||||
logger.error(f"❌ Unsloth Studio failed to start: {exc}")
|
||||
return
|
||||
|
|
@ -236,16 +365,21 @@ def start(port: int = 8888):
|
|||
)
|
||||
return
|
||||
|
||||
_show_and_embed(actual_port)
|
||||
|
||||
# Keep kernel alive so the daemon server thread runs; handle KeyboardInterrupt
|
||||
# cleanly so interrupting the cell gives a readable message.
|
||||
# Open the tunnel now the server is healthy, publish its URL for /api/health, and
|
||||
# tear it down on interrupt (try/finally) rather than orphan the process.
|
||||
try:
|
||||
cf_url = start_cloudflare_tunnel(actual_port) if cloudflare else None
|
||||
_publish_cloudflare_url(cf_url)
|
||||
_show_and_embed(actual_port, cloudflare_url = cf_url)
|
||||
|
||||
# Keep kernel alive so the daemon server thread runs.
|
||||
for _ in range(10000):
|
||||
time.sleep(300)
|
||||
print("=", end = "", flush = True)
|
||||
except KeyboardInterrupt:
|
||||
logger.info("\nUnsloth Studio keepalive stopped.")
|
||||
finally:
|
||||
_stop_cloudflare_tunnel()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import tempfile
|
|||
from loggers import get_logger
|
||||
import os
|
||||
import shutil
|
||||
import contextlib
|
||||
from pathlib import Path
|
||||
from typing import Optional, Tuple, List
|
||||
from unsloth import FastLanguageModel, FastVisionModel, _IS_MLX
|
||||
|
|
@ -37,6 +38,45 @@ logger = get_logger(__name__)
|
|||
_LLAMA_CPP_SCRIPTS_WARNING_EMITTED = False
|
||||
|
||||
|
||||
def _hf_offline(timeout = 3):
|
||||
"""True if export should avoid the Hub: honors the HF offline env vars, else does one
|
||||
cheap TCP reachability probe so a network-down load uses local files / the HF cache
|
||||
instead of hanging on connection timeouts. Proxy-aware (probes the proxy egress when
|
||||
one is configured); disable the probe with UNSLOTH_OFFLINE_PROBE=0."""
|
||||
_offline = {"1", "true", "yes", "on"}
|
||||
if (
|
||||
os.environ.get("HF_HUB_OFFLINE", "").strip().lower() in _offline
|
||||
or os.environ.get("TRANSFORMERS_OFFLINE", "").strip().lower() in _offline
|
||||
):
|
||||
return True
|
||||
if os.environ.get("UNSLOTH_OFFLINE_PROBE", "1").strip().lower() in {"0", "false", "no", "off"}:
|
||||
return False # probe disabled -> assume online; loads still pass local_files_only on env
|
||||
|
||||
# Shared bounded, proxy-aware probe (also used by the export worker before version activation).
|
||||
from utils.transformers_version import hf_endpoint_unreachable
|
||||
|
||||
if hf_endpoint_unreachable(timeout):
|
||||
logger.warning("Hugging Face endpoint unreachable; loading checkpoint in offline mode")
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
# Reuse Unsloth's lock-guarded forced-offline context; no-op fallback if it moves.
|
||||
try:
|
||||
from unsloth.models.loader_utils import _force_hf_offline
|
||||
except Exception:
|
||||
import contextlib as _contextlib
|
||||
|
||||
@_contextlib.contextmanager
|
||||
def _force_hf_offline():
|
||||
yield
|
||||
|
||||
|
||||
def _offline_window_if(local_files_only):
|
||||
"""Forced-offline window when offline was detected, else a no-op context."""
|
||||
return _force_hf_offline() if local_files_only else contextlib.nullcontext()
|
||||
|
||||
|
||||
def _is_wsl():
|
||||
"""Detect if running under Windows Subsystem for Linux."""
|
||||
try:
|
||||
|
|
@ -175,10 +215,19 @@ class ExportBackend:
|
|||
|
||||
model_id = base_model or checkpoint_path
|
||||
|
||||
# Token the type-detection probes too, else a gated multimodal base
|
||||
# 404s here and falls through to the text loader.
|
||||
self._audio_type = detect_audio_type(model_id, hf_token = token)
|
||||
self.is_vision = not self._audio_type and is_vision_model(model_id, hf_token = token)
|
||||
# Skip the Hub when offline so a no-internet export uses the local cache.
|
||||
local_files_only = _hf_offline()
|
||||
|
||||
# Run the type-detection probes in the forced-offline window (else a gated
|
||||
# base 404s); it covers is_vision_model's Hub reads + the transformers-5
|
||||
# subprocess, and local_files_only makes detect_audio_type's requests.get skip.
|
||||
with _offline_window_if(local_files_only):
|
||||
self._audio_type = detect_audio_type(
|
||||
model_id, hf_token = token, local_files_only = local_files_only
|
||||
)
|
||||
self.is_vision = not self._audio_type and is_vision_model(
|
||||
model_id, hf_token = token, local_files_only = local_files_only
|
||||
)
|
||||
|
||||
if self._audio_type == "csm":
|
||||
from unsloth import FastModel
|
||||
|
|
@ -193,6 +242,7 @@ class ExportBackend:
|
|||
load_in_4bit = False,
|
||||
trust_remote_code = trust_remote_code,
|
||||
token = token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
|
||||
elif self._audio_type == "whisper":
|
||||
|
|
@ -207,6 +257,7 @@ class ExportBackend:
|
|||
auto_model = WhisperForConditionalGeneration,
|
||||
trust_remote_code = trust_remote_code,
|
||||
token = token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
|
||||
elif self._audio_type == "snac":
|
||||
|
|
@ -218,6 +269,7 @@ class ExportBackend:
|
|||
load_in_4bit = load_in_4bit,
|
||||
trust_remote_code = trust_remote_code,
|
||||
token = token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
|
||||
elif self._audio_type == "bicodec":
|
||||
|
|
@ -230,6 +282,7 @@ class ExportBackend:
|
|||
load_in_4bit = False,
|
||||
trust_remote_code = trust_remote_code,
|
||||
token = token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
|
||||
elif self._audio_type == "dac":
|
||||
|
|
@ -241,6 +294,7 @@ class ExportBackend:
|
|||
load_in_4bit = False,
|
||||
trust_remote_code = trust_remote_code,
|
||||
token = token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
|
||||
elif self.is_vision:
|
||||
|
|
@ -252,6 +306,7 @@ class ExportBackend:
|
|||
load_in_4bit = load_in_4bit,
|
||||
trust_remote_code = trust_remote_code,
|
||||
token = token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
tokenizer = processor # vision: processor acts as tokenizer
|
||||
|
||||
|
|
@ -264,6 +319,7 @@ class ExportBackend:
|
|||
load_in_4bit = load_in_4bit,
|
||||
trust_remote_code = trust_remote_code,
|
||||
token = token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
|
||||
if _IS_MLX:
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ Pattern follows core/inference/worker.py and core/training/worker.py.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import errno
|
||||
import structlog
|
||||
from loggers import get_logger
|
||||
|
|
@ -171,6 +172,57 @@ def _activate_transformers_version(model_name: str, hf_token: str | None = None)
|
|||
activate_transformers_for_subprocess(model_name, hf_token)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _offline_window_if_unreachable(step = "loading"):
|
||||
"""Force HF offline for a network-touching step (transformers version activation, or the
|
||||
load preflights that hit the Hub) when the endpoint is unreachable, then restore the prior
|
||||
env. Keeps a no-network export from hanging on Hub calls that run before load_checkpoint's
|
||||
own probe, while letting this persistent worker re-decide per operation once back online.
|
||||
|
||||
Post-ML-import (the load preflights), huggingface_hub has already read its in-process
|
||||
offline constant and cached sessions, so env alone is too late: defer to the loader's
|
||||
_force_hf_offline (env + in-process flags + session reset). Pre-import (activation),
|
||||
huggingface_hub is not loaded yet, so setting the env vars suffices for its urllib probes."""
|
||||
saved: dict[str, str | None] = {}
|
||||
force_ctx = None
|
||||
try:
|
||||
from utils.transformers_version import _env_offline, hf_endpoint_unreachable
|
||||
probe_enabled = os.environ.get("UNSLOTH_OFFLINE_PROBE", "1").strip().lower() not in (
|
||||
"0",
|
||||
"false",
|
||||
"no",
|
||||
"off",
|
||||
)
|
||||
if not _env_offline() and probe_enabled and hf_endpoint_unreachable():
|
||||
logger.warning("Hugging Face endpoint unreachable; %s offline", step)
|
||||
if "huggingface_hub" in sys.modules:
|
||||
try:
|
||||
from unsloth.models.loader_utils import _force_hf_offline
|
||||
force_ctx = _force_hf_offline()
|
||||
force_ctx.__enter__() # sets env + in-process flags + resets sessions
|
||||
except Exception:
|
||||
force_ctx = None
|
||||
if force_ctx is None:
|
||||
for k in ("HF_HUB_OFFLINE", "TRANSFORMERS_OFFLINE"):
|
||||
saved[k] = os.environ.get(k)
|
||||
os.environ[k] = "1"
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
if force_ctx is not None:
|
||||
try:
|
||||
force_ctx.__exit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
for k, v in saved.items():
|
||||
if v is None:
|
||||
os.environ.pop(k, None)
|
||||
else:
|
||||
os.environ[k] = v
|
||||
|
||||
|
||||
def _send_response(resp_queue: Any, response: dict) -> None:
|
||||
"""Send a response to the parent process."""
|
||||
try:
|
||||
|
|
@ -459,19 +511,20 @@ def run_export_process(*, cmd_queue: Any, resp_queue: Any, config: dict) -> None
|
|||
checkpoint_path = config["checkpoint_path"]
|
||||
|
||||
# ── 1. Activate correct transformers version BEFORE any ML imports ──
|
||||
try:
|
||||
_activate_transformers_version(checkpoint_path, config.get("hf_token") or None)
|
||||
except Exception as exc:
|
||||
_send_response(
|
||||
resp_queue,
|
||||
{
|
||||
"type": "error",
|
||||
"error": f"Failed to activate transformers version: {exc}",
|
||||
"stack": traceback.format_exc(limit = 20),
|
||||
"ts": time.time(),
|
||||
},
|
||||
)
|
||||
return
|
||||
with _offline_window_if_unreachable(step = "activating transformers"):
|
||||
try:
|
||||
_activate_transformers_version(checkpoint_path, config.get("hf_token") or None)
|
||||
except Exception as exc:
|
||||
_send_response(
|
||||
resp_queue,
|
||||
{
|
||||
"type": "error",
|
||||
"error": f"Failed to activate transformers version: {exc}",
|
||||
"stack": traceback.format_exc(limit = 20),
|
||||
"ts": time.time(),
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
# ── 1b. Check Triton on Windows (must precede import torch) ──
|
||||
if sys.platform == "win32":
|
||||
|
|
@ -534,7 +587,10 @@ def run_export_process(*, cmd_queue: Any, resp_queue: Any, config: dict) -> None
|
|||
try:
|
||||
backend = ExportBackend()
|
||||
|
||||
_handle_load(backend, config, resp_queue)
|
||||
# Offline window covers the load preflights (malware/consent scans hit the Hub)
|
||||
# before load_checkpoint runs its own probe; restored after so later loads re-decide.
|
||||
with _offline_window_if_unreachable():
|
||||
_handle_load(backend, config, resp_queue)
|
||||
|
||||
except Exception as exc:
|
||||
_send_response(
|
||||
|
|
@ -570,7 +626,9 @@ def run_export_process(*, cmd_queue: Any, resp_queue: Any, config: dict) -> None
|
|||
if cmd_type == "load":
|
||||
# Load a new checkpoint, reusing this subprocess.
|
||||
backend.cleanup_memory()
|
||||
_handle_load(backend, cmd, resp_queue)
|
||||
# Offline window also covers this load's Hub preflights (re-probed per load).
|
||||
with _offline_window_if_unreachable():
|
||||
_handle_load(backend, cmd, resp_queue)
|
||||
|
||||
elif cmd_type == "export":
|
||||
_handle_export(backend, cmd, resp_queue)
|
||||
|
|
|
|||
|
|
@ -60,6 +60,7 @@ from utils.transformers_version import (
|
|||
activate_transformers_for_subprocess,
|
||||
_venv_dir_is_valid,
|
||||
_ensure_venv_dir,
|
||||
hf_endpoint_unreachable,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -2428,3 +2429,124 @@ class TestMalformedInputRobustness:
|
|||
|
||||
def test_empty_name_returns_default(self):
|
||||
assert get_transformers_tier("") == "default"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Offline negatives must not poison the version caches (persistent worker)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOfflineCacheNotPoisoned:
|
||||
"""An offline first load must not leave a stale negative for a later online read."""
|
||||
|
||||
def setup_method(self):
|
||||
_tokenizer_class_cache.clear()
|
||||
_config_json_cache.clear()
|
||||
|
||||
def test_offline_tokenizer_assumption_not_cached(self, monkeypatch):
|
||||
import utils.transformers_version as tv
|
||||
|
||||
monkeypatch.setattr(tv, "_env_offline", lambda: True)
|
||||
# No local file, not a local dir -> offline branch returns False without caching.
|
||||
assert _check_tokenizer_config_needs_v5("org/uncached") is False
|
||||
assert ("org/uncached", None) not in _tokenizer_class_cache
|
||||
|
||||
def test_offline_then_online_refetches(self, monkeypatch):
|
||||
import utils.transformers_version as tv
|
||||
|
||||
# 1) Offline: returns False, nothing cached.
|
||||
monkeypatch.setattr(tv, "_env_offline", lambda: True)
|
||||
assert _check_tokenizer_config_needs_v5("org/needs5") is False
|
||||
assert ("org/needs5", None) not in _tokenizer_class_cache
|
||||
|
||||
# 2) Back online: the real fetch runs (cache was not poisoned) and is honored.
|
||||
monkeypatch.setattr(tv, "_env_offline", lambda: False)
|
||||
|
||||
class _Resp:
|
||||
def read(self):
|
||||
return json.dumps({"tokenizer_class": "TokenizersBackend"}).encode()
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
monkeypatch.setattr("urllib.request.urlopen", lambda req, timeout = 10: _Resp())
|
||||
assert _check_tokenizer_config_needs_v5("org/needs5") is True
|
||||
|
||||
def test_offline_config_miss_not_cached(self, monkeypatch):
|
||||
import utils.transformers_version as tv
|
||||
|
||||
monkeypatch.setattr(tv, "_env_offline", lambda: True)
|
||||
monkeypatch.setattr(tv, "_config_json_from_hf_cache", lambda name: None)
|
||||
assert _load_config_json("org/uncached-config", None) is None
|
||||
assert ("org/uncached-config", None) not in _config_json_cache
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# hf_endpoint_unreachable — bounded, proxy/egress-aware reachability probe
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestHfEndpointUnreachable:
|
||||
def test_reachable_returns_false(self, monkeypatch):
|
||||
class _Resp:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
monkeypatch.setattr("urllib.request.urlopen", lambda *a, **k: _Resp())
|
||||
assert hf_endpoint_unreachable(timeout = 2) is False
|
||||
|
||||
def test_gateway_error_is_unreachable(self, monkeypatch):
|
||||
import urllib.error
|
||||
|
||||
def _gw(*a, **k):
|
||||
raise urllib.error.HTTPError("http://x", 504, "Gateway Timeout", {}, None)
|
||||
|
||||
monkeypatch.setattr("urllib.request.urlopen", _gw)
|
||||
assert hf_endpoint_unreachable(timeout = 2) is True
|
||||
|
||||
def test_other_http_status_is_reachable(self, monkeypatch):
|
||||
import urllib.error
|
||||
|
||||
def _405(*a, **k):
|
||||
raise urllib.error.HTTPError("http://x", 405, "Method Not Allowed", {}, None)
|
||||
|
||||
monkeypatch.setattr("urllib.request.urlopen", _405)
|
||||
assert hf_endpoint_unreachable(timeout = 2) is False
|
||||
|
||||
def test_tls_failure_is_reachable(self, monkeypatch):
|
||||
import ssl
|
||||
import urllib.error
|
||||
|
||||
def _tls(*a, **k):
|
||||
raise urllib.error.URLError(ssl.SSLCertVerificationError("self-signed"))
|
||||
|
||||
monkeypatch.setattr("urllib.request.urlopen", _tls)
|
||||
# TLS reached the server: treat as reachable so the load surfaces the cert error.
|
||||
assert hf_endpoint_unreachable(timeout = 2) is False
|
||||
|
||||
def test_dns_failure_is_unreachable(self, monkeypatch):
|
||||
import socket
|
||||
import urllib.error
|
||||
|
||||
def _dns(*a, **k):
|
||||
raise urllib.error.URLError(socket.gaierror(-2, "Name or service not known"))
|
||||
|
||||
monkeypatch.setattr("urllib.request.urlopen", _dns)
|
||||
assert hf_endpoint_unreachable(timeout = 2) is True
|
||||
|
||||
def test_hung_probe_is_bounded(self, monkeypatch):
|
||||
import time
|
||||
|
||||
def _hang(*a, **k):
|
||||
time.sleep(30)
|
||||
|
||||
monkeypatch.setattr("urllib.request.urlopen", _hang)
|
||||
t0 = time.time()
|
||||
result = hf_endpoint_unreachable(timeout = 2)
|
||||
assert result is True and (time.time() - t0) < 6.0
|
||||
|
|
|
|||
|
|
@ -71,7 +71,7 @@ class TestVisionCacheHitMiss:
|
|||
"""Two calls for the same model invoke the uncached fn once."""
|
||||
assert is_vision_model("org/my-vlm") is True
|
||||
assert is_vision_model("org/my-vlm") is True
|
||||
mock_uncached.assert_called_once_with("org/my-vlm", None)
|
||||
mock_uncached.assert_called_once_with("org/my-vlm", None, local_files_only = False)
|
||||
|
||||
@patch("utils.models.model_config._is_vision_model_uncached", return_value = False)
|
||||
def test_different_models_each_detected(self, mock_uncached):
|
||||
|
|
@ -97,7 +97,7 @@ class TestVisionCacheStoresFalse:
|
|||
assert is_vision_model("org/text-only") is False
|
||||
assert is_vision_model("org/text-only") is False
|
||||
mock_uncached.assert_called_once()
|
||||
assert _vision_detection_cache[("org/text-only", None)] is False
|
||||
assert _vision_detection_cache[("org/text-only", None, False)] is False
|
||||
|
||||
|
||||
# Subprocess path (transformers 5.x) caching
|
||||
|
|
@ -120,7 +120,7 @@ class TestVisionCacheSubprocessPath:
|
|||
assert is_vision_model("unsloth/Qwen3.5-2B") is True
|
||||
|
||||
mock_subprocess.assert_called_once()
|
||||
assert _vision_detection_cache[("unsloth/Qwen3.5-2B", None)] is True
|
||||
assert _vision_detection_cache[("unsloth/Qwen3.5-2B", None, False)] is True
|
||||
|
||||
@patch("utils.models.model_config._raw_config_has_vision_config", return_value = True)
|
||||
@patch("utils.models.model_config._is_vision_model_subprocess", return_value = None)
|
||||
|
|
@ -133,7 +133,9 @@ class TestVisionCacheSubprocessPath:
|
|||
assert is_vision_model("unsloth/gemma-4-E4B-it") is True
|
||||
assert is_vision_model("unsloth/gemma-4-E4B-it") is True
|
||||
|
||||
mock_raw_config.assert_called_once_with("unsloth/gemma-4-E4B-it", hf_token = None)
|
||||
mock_raw_config.assert_called_once_with(
|
||||
"unsloth/gemma-4-E4B-it", hf_token = None, local_files_only = False
|
||||
)
|
||||
mock_subprocess.assert_not_called()
|
||||
|
||||
|
||||
|
|
@ -405,6 +407,43 @@ class TestVisionCacheTokenHandling:
|
|||
mock_uncached.assert_called_once()
|
||||
|
||||
|
||||
class TestVisionCacheLocalOnly:
|
||||
"""local_files_only is in the cache key: an offline negative must not be reused by a
|
||||
later online probe (else a VLM is routed through the text loader until restart)."""
|
||||
|
||||
def test_local_only_negative_does_not_poison_online(self, monkeypatch):
|
||||
import utils.models.model_config as mc
|
||||
|
||||
mc._vision_detection_cache.clear()
|
||||
monkeypatch.setattr(mc, "is_local_path", lambda *_a, **_k: False)
|
||||
monkeypatch.setattr(mc, "resolve_cached_repo_id_case", lambda n, *_a, **_k: n)
|
||||
# Pin env-offline off so the key tracks the kwarg.
|
||||
monkeypatch.setattr(mc, "_env_offline", lambda: False)
|
||||
|
||||
seen = []
|
||||
|
||||
def _probe(
|
||||
name,
|
||||
hf_token = None,
|
||||
local_files_only = False,
|
||||
):
|
||||
seen.append(local_files_only)
|
||||
# Offline can't fetch -> not a VLM; online reveals the VLM.
|
||||
return False if local_files_only else True
|
||||
|
||||
monkeypatch.setattr(mc, "_is_vision_model_uncached", _probe)
|
||||
|
||||
# Offline probe caches False under a local-only key.
|
||||
assert mc.is_vision_model("some/vlm", local_files_only = True) is False
|
||||
# A later online probe must re-run (different key) and detect the VLM.
|
||||
assert mc.is_vision_model("some/vlm", local_files_only = False) is True
|
||||
assert seen == [True, False]
|
||||
# The online positive is then cached for subsequent online callers.
|
||||
assert mc.is_vision_model("some/vlm", local_files_only = False) is True
|
||||
assert seen == [True, False]
|
||||
mc._vision_detection_cache.clear()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Direct unit tests for _raw_config_has_vision_config
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -570,7 +609,11 @@ class TestAudioDetectionCacheTokenAware:
|
|||
mc._audio_detection_cache.clear()
|
||||
calls = []
|
||||
|
||||
def _fake(name, hf_token = None):
|
||||
def _fake(
|
||||
name,
|
||||
hf_token = None,
|
||||
local_files_only = False,
|
||||
):
|
||||
calls.append(hf_token)
|
||||
# Gated repo: only an authenticated probe can read the tokenizer.
|
||||
return ("bicodec", True) if hf_token else (None, True)
|
||||
|
|
@ -601,7 +644,11 @@ class TestAudioDetectionCacheTokenAware:
|
|||
|
||||
transient_calls = []
|
||||
|
||||
def _transient(name, hf_token = None):
|
||||
def _transient(
|
||||
name,
|
||||
hf_token = None,
|
||||
local_files_only = False,
|
||||
):
|
||||
transient_calls.append(hf_token)
|
||||
return (None, False) # network/5xx -- not cacheable
|
||||
|
||||
|
|
@ -613,7 +660,11 @@ class TestAudioDetectionCacheTokenAware:
|
|||
|
||||
definitive_calls = []
|
||||
|
||||
def _definitive(name, hf_token = None):
|
||||
def _definitive(
|
||||
name,
|
||||
hf_token = None,
|
||||
local_files_only = False,
|
||||
):
|
||||
definitive_calls.append(hf_token)
|
||||
return (None, True) # read the config, no audio tokens
|
||||
|
||||
|
|
@ -623,3 +674,94 @@ class TestAudioDetectionCacheTokenAware:
|
|||
# Probed once: the definitive None was cached.
|
||||
assert definitive_calls == [None]
|
||||
mc._audio_detection_cache.clear()
|
||||
|
||||
def test_local_only_negative_does_not_poison_online(self, monkeypatch):
|
||||
"""An offline negative must not be reused by a later online probe (else an audio
|
||||
model is routed through the text loader until restart)."""
|
||||
import utils.models.model_config as mc
|
||||
|
||||
mc._audio_detection_cache.clear()
|
||||
monkeypatch.setattr(mc, "is_local_path", lambda *_a, **_k: False)
|
||||
monkeypatch.setattr(mc, "resolve_cached_repo_id_case", lambda n, *_a, **_k: n)
|
||||
# Pin env-offline off so the key tracks the kwarg.
|
||||
monkeypatch.setattr(mc, "_env_offline", lambda: False)
|
||||
|
||||
seen = []
|
||||
|
||||
def _probe(
|
||||
name,
|
||||
hf_token = None,
|
||||
local_files_only = False,
|
||||
):
|
||||
seen.append(local_files_only)
|
||||
# Offline: nothing on disk -> not audio; online reveals the audio model.
|
||||
return (None, True) if local_files_only else ("snac", True)
|
||||
|
||||
monkeypatch.setattr(mc, "_detect_audio_from_tokenizer", _probe)
|
||||
|
||||
# Offline probe caches None under a local-only key.
|
||||
assert mc.detect_audio_type("some/audio-model", local_files_only = True) is None
|
||||
# A later online probe must re-run (different key) and detect the audio model.
|
||||
assert mc.detect_audio_type("some/audio-model", local_files_only = False) == "snac"
|
||||
assert seen == [True, False]
|
||||
# The online positive is then cached for subsequent online callers.
|
||||
assert mc.detect_audio_type("some/audio-model", local_files_only = False) == "snac"
|
||||
assert seen == [True, False]
|
||||
mc._audio_detection_cache.clear()
|
||||
|
||||
def test_env_offline_negative_does_not_poison_online(self, monkeypatch):
|
||||
"""An env-offline probe (default local_files_only=False) must cache under the
|
||||
effective-offline key, so clearing the env var later doesn't leak a stale negative."""
|
||||
import utils.models.model_config as mc
|
||||
|
||||
mc._audio_detection_cache.clear()
|
||||
monkeypatch.setattr(mc, "is_local_path", lambda *_a, **_k: False)
|
||||
monkeypatch.setattr(mc, "resolve_cached_repo_id_case", lambda n, *_a, **_k: n)
|
||||
|
||||
env_offline = {"v": True}
|
||||
monkeypatch.setattr(mc, "_env_offline", lambda: env_offline["v"])
|
||||
|
||||
seen = []
|
||||
|
||||
def _probe(
|
||||
name,
|
||||
hf_token = None,
|
||||
local_files_only = False,
|
||||
):
|
||||
seen.append(local_files_only)
|
||||
return (None, True) if local_files_only else ("snac", True)
|
||||
|
||||
monkeypatch.setattr(mc, "_detect_audio_from_tokenizer", _probe)
|
||||
|
||||
# Env offline + default kwarg -> probe runs offline; None cached under the offline key.
|
||||
assert mc.detect_audio_type("some/audio-model") is None
|
||||
assert seen == [True]
|
||||
# Env var cleared: a fresh online probe must re-run (different key) and detect.
|
||||
env_offline["v"] = False
|
||||
assert mc.detect_audio_type("some/audio-model") == "snac"
|
||||
assert seen == [True, False]
|
||||
mc._audio_detection_cache.clear()
|
||||
|
||||
|
||||
class TestEnvOfflineParsing:
|
||||
"""_env_offline accepts the canonical truthy set (strip+lower, on/true/yes/1); it gates
|
||||
the requests.get fallback and the cache keys, so 'on' or ' 1 ' must still count as offline."""
|
||||
|
||||
def test_truthy_values_recognized(self, monkeypatch):
|
||||
import utils.models.model_config as mc
|
||||
for var in ("HF_HUB_OFFLINE", "TRANSFORMERS_OFFLINE"):
|
||||
for val in ("1", "true", "TRUE", "yes", "Yes", "on", "ON", " 1 ", " on ", "\ttrue\n"):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
monkeypatch.setenv(var, val)
|
||||
assert mc._env_offline() is True, f"{var}={val!r} should be offline"
|
||||
|
||||
def test_falsy_values_not_offline(self, monkeypatch):
|
||||
import utils.models.model_config as mc
|
||||
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
assert mc._env_offline() is False
|
||||
for val in ("", "0", "false", "no", "off", "2", "onn"):
|
||||
monkeypatch.setenv("HF_HUB_OFFLINE", val)
|
||||
assert mc._env_offline() is False, f"HF_HUB_OFFLINE={val!r} should not be offline"
|
||||
|
|
|
|||
|
|
@ -44,13 +44,15 @@ from utils.subprocess_compat import (
|
|||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
_OFFLINE_TRUE_VALUES = {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
def _env_offline() -> bool:
|
||||
"""True if HF_HUB_OFFLINE or TRANSFORMERS_OFFLINE is set to a truthy value."""
|
||||
return os.environ.get("HF_HUB_OFFLINE", "").lower() in (
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
) or os.environ.get("TRANSFORMERS_OFFLINE", "").lower() in ("1", "true", "yes")
|
||||
"""True if an HF offline env var is truthy (canonical strip+lower parse, on/true/yes/1)."""
|
||||
return (
|
||||
os.environ.get("HF_HUB_OFFLINE", "").strip().lower() in _OFFLINE_TRUE_VALUES
|
||||
or os.environ.get("TRANSFORMERS_OFFLINE", "").strip().lower() in _OFFLINE_TRUE_VALUES
|
||||
)
|
||||
|
||||
|
||||
# ── Model size extraction ────────────────────────────────────
|
||||
|
|
@ -471,6 +473,7 @@ def load_model_config(
|
|||
use_auth: bool = False,
|
||||
token: Optional[str] = None,
|
||||
trust_remote_code: bool = False,
|
||||
local_files_only: bool = False,
|
||||
):
|
||||
"""Load model config with optional authentication control.
|
||||
|
||||
|
|
@ -478,12 +481,18 @@ def load_model_config(
|
|||
metadata lookups must never execute a model repo's ``auto_map`` Python.
|
||||
Deliberate remote-code loads pass the flag explicitly through
|
||||
``FastLanguageModel.from_pretrained`` with the user's own consent.
|
||||
|
||||
``local_files_only`` keeps the config read on the local HF cache (offline
|
||||
export), so an offline probe never blocks on the network.
|
||||
"""
|
||||
from transformers import AutoConfig
|
||||
|
||||
if token:
|
||||
return AutoConfig.from_pretrained(
|
||||
model_name, trust_remote_code = trust_remote_code, token = token
|
||||
model_name,
|
||||
trust_remote_code = trust_remote_code,
|
||||
token = token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
|
||||
if not use_auth:
|
||||
|
|
@ -493,12 +502,14 @@ def load_model_config(
|
|||
model_name,
|
||||
trust_remote_code = trust_remote_code,
|
||||
token = None,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
|
||||
# Default auth (cached tokens)
|
||||
return AutoConfig.from_pretrained(
|
||||
model_name,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -598,7 +609,9 @@ def _is_vlm(config) -> bool:
|
|||
|
||||
|
||||
def _raw_config_has_vision_config(
|
||||
model_name: str, hf_token: Optional[str] = None
|
||||
model_name: str,
|
||||
hf_token: Optional[str] = None,
|
||||
local_files_only: bool = False,
|
||||
) -> Optional[bool]:
|
||||
try:
|
||||
if is_local_path(model_name):
|
||||
|
|
@ -610,6 +623,7 @@ def _raw_config_has_vision_config(
|
|||
repo_id = model_name,
|
||||
filename = "config.json",
|
||||
token = hf_token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
)
|
||||
config = json.loads(config_path.read_text())
|
||||
|
|
@ -776,27 +790,20 @@ def _token_fingerprint(token: Optional[str]) -> Optional[str]:
|
|||
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
# Cache vision detection per session to avoid repeated subprocess spawns.
|
||||
# Keyed by (normalized_model_name, token_fingerprint) to handle gated models.
|
||||
# Only definitive results are cached; transient failures (network, timeouts)
|
||||
# are NOT cached so they can be retried.
|
||||
_vision_detection_cache: Dict[Tuple[str, Optional[str]], bool] = {}
|
||||
# Vision detection cache keyed by (name, token, local_files_only); only definitive results cached.
|
||||
_vision_detection_cache: Dict[Tuple[str, Optional[str], bool], bool] = {}
|
||||
_vision_cache_lock = threading.Lock()
|
||||
|
||||
|
||||
def is_vision_model(model_name: str, hf_token: Optional[str] = None) -> bool:
|
||||
"""
|
||||
Detect vision-language models (VLMs) via architecture in config. Works for
|
||||
fine-tuned models since they inherit the base architecture.
|
||||
|
||||
Models needing transformers 5.x are checked in a .venv_t5/ subprocess.
|
||||
Results are cached per (model_name, token_fingerprint) for the process
|
||||
lifetime; transient failures are not cached so they can be retried.
|
||||
|
||||
Args:
|
||||
model_name: Model identifier (HF repo or local path)
|
||||
hf_token: Optional HF token for gated/private models
|
||||
"""
|
||||
def is_vision_model(
|
||||
model_name: str,
|
||||
hf_token: Optional[str] = None,
|
||||
local_files_only: bool = False,
|
||||
) -> bool:
|
||||
"""Detect VLMs via the config architecture (works for fine-tunes); transformers-5.x
|
||||
models are checked in a .venv_t5/ subprocess. Cached per (model_name, token,
|
||||
local_files_only) minus transient failures; local_files_only is in the key so an
|
||||
offline probe never shares an online entry."""
|
||||
# Local GGUF models are served by llama-server. Their multimodal
|
||||
# capability comes from a companion mmproj, not a Transformers config.
|
||||
# Do not cache this lookup: a projector may be added beside an existing
|
||||
|
|
@ -829,7 +836,10 @@ def is_vision_model(model_name: str, hf_token: Optional[str] = None) -> bool:
|
|||
exc,
|
||||
)
|
||||
resolved_name = model_name
|
||||
cache_key = (resolved_name, _token_fingerprint(hf_token))
|
||||
# Key on effective offline (kwarg OR env) so an offline probe can't poison a later
|
||||
# online lookup once the env var is cleared.
|
||||
effective_offline = bool(local_files_only or _env_offline())
|
||||
cache_key = (resolved_name, _token_fingerprint(hf_token), effective_offline)
|
||||
|
||||
# Lock-free fast path for cache hits. Sentinel distinguishes "key not found"
|
||||
# from "value is False" in a single atomic dict.get() call.
|
||||
|
|
@ -840,7 +850,7 @@ def is_vision_model(model_name: str, hf_token: Optional[str] = None) -> bool:
|
|||
|
||||
# Compute outside the lock so long-running detection isn't serialized across
|
||||
# models. Two concurrent calls may both run, but produce the same result.
|
||||
result = _is_vision_model_uncached(resolved_name, hf_token)
|
||||
result = _is_vision_model_uncached(resolved_name, hf_token, local_files_only = effective_offline)
|
||||
# Only cache definitive results; None is a transient failure, retry later.
|
||||
if result is not None:
|
||||
with _vision_cache_lock:
|
||||
|
|
@ -849,7 +859,11 @@ def is_vision_model(model_name: str, hf_token: Optional[str] = None) -> bool:
|
|||
return False
|
||||
|
||||
|
||||
def _is_vision_model_uncached(model_name: str, hf_token: Optional[str] = None) -> Optional[bool]:
|
||||
def _is_vision_model_uncached(
|
||||
model_name: str,
|
||||
hf_token: Optional[str] = None,
|
||||
local_files_only: bool = False,
|
||||
) -> Optional[bool]:
|
||||
"""Uncached vision detection; use is_vision_model() instead.
|
||||
|
||||
Returns True/False for definitive results, or None on transient errors
|
||||
|
|
@ -858,15 +872,17 @@ def _is_vision_model_uncached(model_name: str, hf_token: Optional[str] = None) -
|
|||
# Try the raw-config reader FIRST (code-free, version-independent): it classifies
|
||||
# repo-code VLMs like DeepSeek-OCR via declarative vision_config with no remote-code
|
||||
# execution or transformers-5.x subprocess.
|
||||
raw = _raw_config_has_vision_config(model_name, hf_token = hf_token)
|
||||
raw = _raw_config_has_vision_config(
|
||||
model_name, hf_token = hf_token, local_files_only = local_files_only
|
||||
)
|
||||
if raw is not None:
|
||||
return raw
|
||||
|
||||
# Raw read failed transiently: fall back to AutoConfig with remote code DISABLED
|
||||
# (in a transformers-5.x subprocess when the main process can't parse the arch).
|
||||
# Raw read failed transiently: fall back to AutoConfig (remote code DISABLED), via a
|
||||
# transformers-5.x subprocess if needed. Skip that subprocess offline (it probes the network).
|
||||
from utils.transformers_version import needs_transformers_5
|
||||
|
||||
if needs_transformers_5(model_name):
|
||||
if not local_files_only and needs_transformers_5(model_name):
|
||||
logger.info(
|
||||
"Model '%s' needs transformers 5.x -- checking vision via subprocess",
|
||||
model_name,
|
||||
|
|
@ -874,7 +890,12 @@ def _is_vision_model_uncached(model_name: str, hf_token: Optional[str] = None) -
|
|||
return _is_vision_model_subprocess(model_name, hf_token = hf_token)
|
||||
|
||||
try:
|
||||
config = load_model_config(model_name, use_auth = True, token = hf_token)
|
||||
config = load_model_config(
|
||||
model_name,
|
||||
use_auth = True,
|
||||
token = hf_token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
|
||||
if _is_vlm(config):
|
||||
model_type = getattr(config, "model_type", None)
|
||||
|
|
@ -914,9 +935,9 @@ def _is_vision_model_uncached(model_name: str, hf_token: Optional[str] = None) -
|
|||
|
||||
VALID_AUDIO_TYPES = ("snac", "csm", "bicodec", "dac", "whisper", "audio_vlm")
|
||||
|
||||
# Keyed by (normalized_name, token_fingerprint) like the vision cache, so an
|
||||
# unauthenticated miss (None) cannot poison a later authenticated lookup.
|
||||
_audio_detection_cache: Dict[Tuple[str, Optional[str]], Optional[str]] = {}
|
||||
# Keyed like the vision cache by (name, token, local_files_only) so an unauthenticated
|
||||
# or offline miss cannot poison a later authenticated / online lookup.
|
||||
_audio_detection_cache: Dict[Tuple[str, Optional[str], bool], Optional[str]] = {}
|
||||
|
||||
# Tokenizer token patterns → audio_type (all 6 types from tokenizer_config.json)
|
||||
_AUDIO_TOKEN_PATTERNS = {
|
||||
|
|
@ -935,12 +956,20 @@ _AUDIO_TOKEN_PATTERNS = {
|
|||
}
|
||||
|
||||
|
||||
def detect_audio_type(model_name: str, hf_token: Optional[str] = None) -> Optional[str]:
|
||||
def detect_audio_type(
|
||||
model_name: str,
|
||||
hf_token: Optional[str] = None,
|
||||
local_files_only: bool = False,
|
||||
) -> Optional[str]:
|
||||
"""Detect if a model is an audio model and return its type.
|
||||
|
||||
Works for any model via tokenizer_config.json special tokens.
|
||||
Returns an audio_type string ('snac', 'csm', 'bicodec', 'dac', 'whisper',
|
||||
'audio_vlm') or None.
|
||||
|
||||
When local_files_only is True (offline export) the remote HuggingFace fetch
|
||||
is skipped so detection never blocks on a network read; only the local HF
|
||||
cache is consulted.
|
||||
"""
|
||||
# Normalize casing + include the token fingerprint (mirrors is_vision_model).
|
||||
try:
|
||||
|
|
@ -950,11 +979,16 @@ def detect_audio_type(model_name: str, hf_token: Optional[str] = None) -> Option
|
|||
resolved_name = resolve_cached_repo_id_case(model_name)
|
||||
except Exception:
|
||||
resolved_name = model_name
|
||||
cache_key = (resolved_name, _token_fingerprint(hf_token))
|
||||
# Key on effective offline (kwarg OR env), matching where the remote fetch is skipped,
|
||||
# so an offline negative can't poison a later online probe.
|
||||
effective_offline = bool(local_files_only or _env_offline())
|
||||
cache_key = (resolved_name, _token_fingerprint(hf_token), effective_offline)
|
||||
if cache_key in _audio_detection_cache:
|
||||
return _audio_detection_cache[cache_key]
|
||||
|
||||
result, definitive = _detect_audio_from_tokenizer(model_name, hf_token)
|
||||
result, definitive = _detect_audio_from_tokenizer(
|
||||
model_name, hf_token, local_files_only = effective_offline
|
||||
)
|
||||
# Cache only definitive results; a transient read failure stays None and retries.
|
||||
if definitive:
|
||||
_audio_detection_cache[cache_key] = result
|
||||
|
|
@ -964,12 +998,15 @@ def detect_audio_type(model_name: str, hf_token: Optional[str] = None) -> Option
|
|||
|
||||
|
||||
def _detect_audio_from_tokenizer(
|
||||
model_name: str, hf_token: Optional[str] = None
|
||||
model_name: str,
|
||||
hf_token: Optional[str] = None,
|
||||
local_files_only: bool = False,
|
||||
) -> Tuple[Optional[str], bool]:
|
||||
"""Detect audio type from tokenizer special tokens.
|
||||
|
||||
Checks local HF cache first, then fetches tokenizer_config.json from HF;
|
||||
examines added_tokens_decoder for distinctive patterns.
|
||||
Checks local HF cache first, then (unless local_files_only) fetches
|
||||
tokenizer_config.json from HF; examines added_tokens_decoder for distinctive
|
||||
patterns.
|
||||
|
||||
Returns (audio_type_or_None, definitive). definitive is False only on a
|
||||
transient read failure (network/timeout/5xx) so the caller skips caching and
|
||||
|
|
@ -1009,7 +1046,11 @@ def _detect_audio_from_tokenizer(
|
|||
except Exception as e:
|
||||
logger.debug(f"Could not check local cache for {model_name}: {e}")
|
||||
|
||||
# 2) Fall back to HuggingFace API
|
||||
# 2) Fall back to the HuggingFace API. This raw requests.get ignores the HF offline
|
||||
# flag, so gate it on local_files_only OR the env vars to skip the network offline.
|
||||
if local_files_only or _env_offline():
|
||||
return None, read_any
|
||||
|
||||
try:
|
||||
import requests
|
||||
import os
|
||||
|
|
|
|||
|
|
@ -46,13 +46,55 @@ from utils.subprocess_compat import (
|
|||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
_OFFLINE_TRUE_VALUES = {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
def _env_offline() -> bool:
|
||||
"""True if HF_HUB_OFFLINE or TRANSFORMERS_OFFLINE is set to a truthy value."""
|
||||
return os.environ.get("HF_HUB_OFFLINE", "").lower() in (
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
) or os.environ.get("TRANSFORMERS_OFFLINE", "").lower() in ("1", "true", "yes")
|
||||
"""True if an HF offline env var is truthy (canonical strip+lower parse); gates the urllib fetches below."""
|
||||
return (
|
||||
os.environ.get("HF_HUB_OFFLINE", "").strip().lower() in _OFFLINE_TRUE_VALUES
|
||||
or os.environ.get("TRANSFORMERS_OFFLINE", "").strip().lower() in _OFFLINE_TRUE_VALUES
|
||||
)
|
||||
|
||||
|
||||
def hf_endpoint_unreachable(timeout: int = 3) -> bool:
|
||||
"""Bounded reachability probe to the HF endpoint. A HEAD request runs in a daemon thread
|
||||
joined with a deadline, so a resolver blackhole cannot block past ~timeout+1s. True if
|
||||
unreachable. urllib natively honors *_PROXY / NO_PROXY, so this verifies real egress
|
||||
(the proxy can reach HF), not just that the proxy is up. No ML imports, so it is safe to
|
||||
call before transformers version activation. Mirrors the probe in export._hf_offline."""
|
||||
import ssl
|
||||
import threading
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
|
||||
endpoint = os.environ.get("HF_ENDPOINT", "https://huggingface.co")
|
||||
if "://" not in endpoint:
|
||||
endpoint = "https://" + endpoint
|
||||
|
||||
result = {"online": False}
|
||||
|
||||
def _probe():
|
||||
try:
|
||||
req = urllib.request.Request(endpoint, method = "HEAD")
|
||||
with urllib.request.urlopen(req, timeout = timeout):
|
||||
result["online"] = True
|
||||
except urllib.error.HTTPError as exc:
|
||||
# The server/proxy answered: reachable unless it is a gateway error.
|
||||
result["online"] = exc.code not in (502, 503, 504)
|
||||
except urllib.error.URLError as exc:
|
||||
# A TLS/cert failure means we DID reach the server; treat as reachable so the real
|
||||
# load surfaces it (consistent with _is_offline_related_error not retrying TLS).
|
||||
result["online"] = isinstance(exc.reason, ssl.SSLError)
|
||||
except ssl.SSLError:
|
||||
result["online"] = True
|
||||
except Exception:
|
||||
result["online"] = False
|
||||
|
||||
t = threading.Thread(target = _probe, daemon = True)
|
||||
t.start()
|
||||
t.join(timeout + 1)
|
||||
return t.is_alive() or not result["online"]
|
||||
|
||||
|
||||
def _safe_is_file(p: Path) -> bool:
|
||||
|
|
@ -151,6 +193,8 @@ _TRANSFORMERS_5_TOKENIZER_CLASSES: set[str] = {
|
|||
|
||||
# Caches keyed on (model_name, token-hash) so authed/unauthed reads stay separate (a
|
||||
# gated/private repo's unauthenticated miss must not poison a later authenticated lookup).
|
||||
# Offline negatives are NOT written (see the _env_offline branches) so they cannot poison a
|
||||
# later online read in this persistent worker.
|
||||
_tokenizer_class_cache: dict[tuple[str, str | None], bool] = {}
|
||||
_config_json_cache: dict[tuple[str, str | None], dict | None] = {}
|
||||
_config_needs_510_cache: dict[tuple[str, str | None], bool] = {}
|
||||
|
|
@ -525,9 +569,9 @@ def _check_tokenizer_config_needs_v5(model_name: str, hf_token: str | None = Non
|
|||
if _safe_is_dir(local_path):
|
||||
return False
|
||||
|
||||
# Offline: skip the 10s urllib fetch (fail-open to lower tier).
|
||||
# Offline: skip the 10s urllib fetch (fail-open to lower tier). Do NOT cache this
|
||||
# assumed negative, so a later online read of the same id re-fetches the real value.
|
||||
if _env_offline():
|
||||
_tokenizer_class_cache[cache_key] = False
|
||||
return False
|
||||
|
||||
# --- Fall back to fetching from HuggingFace ----------------------------
|
||||
|
|
@ -633,9 +677,11 @@ def _load_config_json(model_name: str, hf_token: str | None = None) -> dict | No
|
|||
return None
|
||||
|
||||
if _env_offline():
|
||||
# No network: a previously downloaded repo can still tier from the hub cache.
|
||||
# No network: a previously downloaded repo can still tier from the hub cache. Cache a
|
||||
# real hit, but never the miss (None) so a later online read still fetches the config.
|
||||
cfg = _config_json_from_hf_cache(model_name)
|
||||
_config_json_cache[cache_key] = cfg
|
||||
if cfg is not None:
|
||||
_config_json_cache[cache_key] = cfg
|
||||
return cfg
|
||||
|
||||
import urllib.error
|
||||
|
|
|
|||
539
tests/test_offline_loading_helpers.py
Normal file
539
tests/test_offline_loading_helpers.py
Normal file
|
|
@ -0,0 +1,539 @@
|
|||
"""Unit tests for the offline-loading helpers in unsloth/models/loader_utils.py:
|
||||
error classification, _force_hf_offline flip/restore, and the retry orchestrator.
|
||||
Pure CPU, no network, no GPU."""
|
||||
|
||||
import os
|
||||
import socket
|
||||
|
||||
import pytest
|
||||
|
||||
from unsloth.models import loader_utils as L
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _env_says_offline / _get_effective_local_files_only
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_OFFLINE_TRUE = ("1", "true", "yes", "on", "ON", " 1 ", "\tyes\n")
|
||||
_OFFLINE_FALSE = ("0", "no", "false", "off", "", " ", "maybe")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", _OFFLINE_TRUE)
|
||||
def test_env_says_offline_truthy(monkeypatch, value):
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
monkeypatch.setenv("HF_HUB_OFFLINE", value)
|
||||
assert L._env_says_offline() is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", _OFFLINE_FALSE)
|
||||
def test_env_says_offline_falsy(monkeypatch, value):
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
monkeypatch.setenv("HF_HUB_OFFLINE", value)
|
||||
assert L._env_says_offline() is False
|
||||
|
||||
|
||||
def test_env_says_offline_absent(monkeypatch):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
assert L._env_says_offline() is False
|
||||
|
||||
|
||||
def test_env_says_offline_transformers_var(monkeypatch):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.setenv("TRANSFORMERS_OFFLINE", "1")
|
||||
assert L._env_says_offline() is True
|
||||
|
||||
|
||||
def test_effective_lfo_kwarg_wins(monkeypatch):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
assert L._get_effective_local_files_only({"local_files_only": True}) is True
|
||||
|
||||
|
||||
def test_effective_lfo_env_only(monkeypatch):
|
||||
monkeypatch.setenv("HF_HUB_OFFLINE", "1")
|
||||
assert L._get_effective_local_files_only({}) is True
|
||||
|
||||
|
||||
def test_effective_lfo_neither(monkeypatch):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
assert L._get_effective_local_files_only({"local_files_only": False}) is False
|
||||
|
||||
|
||||
def test_effective_lfo_is_read_only():
|
||||
# Must not pop local_files_only: the weight load reuses the same kwarg.
|
||||
kwargs = {"local_files_only": True}
|
||||
L._get_effective_local_files_only(kwargs)
|
||||
assert kwargs == {"local_files_only": True}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _is_offline_related_error
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _http_error(status):
|
||||
import requests
|
||||
|
||||
resp = requests.Response()
|
||||
resp.status_code = status
|
||||
return requests.exceptions.HTTPError("http %s" % status, response = resp)
|
||||
|
||||
|
||||
def test_none_is_not_offline():
|
||||
assert L._is_offline_related_error(None) is False
|
||||
|
||||
|
||||
def test_plain_connection_error_is_offline():
|
||||
assert L._is_offline_related_error(ConnectionError("down")) is True
|
||||
|
||||
|
||||
def test_timeout_error_is_offline():
|
||||
assert L._is_offline_related_error(TimeoutError("slow")) is True
|
||||
|
||||
|
||||
def test_plain_file_not_found_propagates():
|
||||
assert L._is_offline_related_error(FileNotFoundError("config.json")) is False
|
||||
|
||||
|
||||
def test_unrelated_error_is_not_offline():
|
||||
assert L._is_offline_related_error(ValueError("bad arg")) is False
|
||||
|
||||
|
||||
def test_requests_connection_error_is_offline():
|
||||
import requests
|
||||
assert L._is_offline_related_error(requests.exceptions.ConnectionError("x")) is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", (500, 502, 503, 504))
|
||||
def test_http_5xx_is_offline(status):
|
||||
assert L._is_offline_related_error(_http_error(status)) is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", (400, 401, 403, 404))
|
||||
def test_http_4xx_propagates(status):
|
||||
assert L._is_offline_related_error(_http_error(status)) is False
|
||||
|
||||
|
||||
def test_status_less_http_with_network_wording_is_offline():
|
||||
import requests
|
||||
err = requests.exceptions.HTTPError("Couldn't connect to the server")
|
||||
assert L._is_offline_related_error(err) is True
|
||||
|
||||
|
||||
def test_status_less_http_without_network_wording_propagates():
|
||||
import requests
|
||||
err = requests.exceptions.HTTPError("I'm a teapot")
|
||||
assert L._is_offline_related_error(err) is False
|
||||
|
||||
|
||||
def test_gaierror_dns_failure_is_offline():
|
||||
assert L._is_offline_related_error(socket.gaierror(-2, "Name or service not known")) is True
|
||||
|
||||
|
||||
def test_gaierror_without_wording_is_offline_by_type():
|
||||
# Matched by type, so a locale-specific / empty message still classifies offline.
|
||||
assert L._is_offline_related_error(socket.gaierror(-2, "")) is True
|
||||
|
||||
|
||||
def test_urllib_urlerror_is_offline():
|
||||
import urllib.error
|
||||
assert L._is_offline_related_error(urllib.error.URLError("connection failed")) is True
|
||||
|
||||
|
||||
def test_urllib_httperror_404_propagates():
|
||||
import urllib.error
|
||||
err = urllib.error.HTTPError("http://x", 404, "Not Found", {}, None)
|
||||
assert L._is_offline_related_error(err) is False
|
||||
|
||||
|
||||
def test_urllib_httperror_503_is_offline():
|
||||
import urllib.error
|
||||
err = urllib.error.HTTPError("http://x", 503, "Service Unavailable", {}, None)
|
||||
assert L._is_offline_related_error(err) is True
|
||||
|
||||
|
||||
def test_ssl_error_is_not_offline():
|
||||
# TLS/cert failure must surface, not silently fall back to cached files.
|
||||
import ssl
|
||||
assert L._is_offline_related_error(ssl.SSLError("certificate verify failed")) is False
|
||||
|
||||
|
||||
def test_requests_ssl_error_is_not_offline():
|
||||
# requests.SSLError subclasses ConnectionError, but is still a TLS failure -> not offline.
|
||||
requests = pytest.importorskip("requests")
|
||||
assert L._is_offline_related_error(requests.exceptions.SSLError("bad cert")) is False
|
||||
|
||||
|
||||
def test_urlerror_wrapping_ssl_is_not_offline():
|
||||
import ssl
|
||||
import urllib.error
|
||||
|
||||
err = urllib.error.URLError(ssl.SSLCertVerificationError("self-signed certificate"))
|
||||
assert L._is_offline_related_error(err) is False
|
||||
|
||||
|
||||
def test_ssl_node_does_not_hide_deeper_connection_cause():
|
||||
# Skipping a TLS node must not abort the walk: a genuine outage deeper still counts.
|
||||
import ssl
|
||||
|
||||
outer = RuntimeError("load failed")
|
||||
mid = ssl.SSLError("cert")
|
||||
mid.__context__ = ConnectionError("down")
|
||||
outer.__cause__ = mid
|
||||
assert L._is_offline_related_error(outer) is True
|
||||
|
||||
|
||||
def test_oserror_network_unreachable_is_offline():
|
||||
assert L._is_offline_related_error(OSError("Network is unreachable")) is True
|
||||
|
||||
|
||||
def test_offline_mode_is_enabled_is_offline():
|
||||
errors = pytest.importorskip("huggingface_hub.errors")
|
||||
assert L._is_offline_related_error(errors.OfflineModeIsEnabled("offline")) is True
|
||||
|
||||
|
||||
def test_local_entry_not_found_is_offline():
|
||||
# Both a FileNotFoundError and an HfHubHTTPError, but means "not cached + Hub down" -> offline.
|
||||
errors = pytest.importorskip("huggingface_hub.errors")
|
||||
assert L._is_offline_related_error(errors.LocalEntryNotFoundError("missing")) is True
|
||||
|
||||
|
||||
def test_chained_cause_connection_error_is_offline():
|
||||
err = RuntimeError("combined load failure")
|
||||
err.__cause__ = ConnectionError("down")
|
||||
assert L._is_offline_related_error(err) is True
|
||||
|
||||
|
||||
def test_chained_context_connection_error_is_offline():
|
||||
try:
|
||||
try:
|
||||
raise ConnectionError("down")
|
||||
except ConnectionError:
|
||||
raise RuntimeError("wrap")
|
||||
except RuntimeError as e:
|
||||
err = e
|
||||
assert L._is_offline_related_error(err) is True
|
||||
|
||||
|
||||
def test_chained_cause_404_still_propagates():
|
||||
err = RuntimeError("combined load failure")
|
||||
err.__cause__ = _http_error(404)
|
||||
assert L._is_offline_related_error(err) is False
|
||||
|
||||
|
||||
def test_cause_context_cycle_terminates():
|
||||
a = RuntimeError("a")
|
||||
b = RuntimeError("b")
|
||||
a.__context__ = b
|
||||
b.__context__ = a
|
||||
# Must not hang; neither is network-related.
|
||||
assert L._is_offline_related_error(a) is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _force_hf_offline
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _inprocess_offline_flags():
|
||||
flags = []
|
||||
try:
|
||||
import huggingface_hub.constants as hfc
|
||||
if hasattr(hfc, "HF_HUB_OFFLINE"):
|
||||
flags.append(hfc.HF_HUB_OFFLINE)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
import transformers.utils.hub as tuh
|
||||
for attr in ("_is_offline_mode", "OFFLINE"):
|
||||
if hasattr(tuh, attr):
|
||||
flags.append(getattr(tuh, attr))
|
||||
except Exception:
|
||||
pass
|
||||
return flags
|
||||
|
||||
|
||||
def test_force_offline_sets_and_restores_absent_env(monkeypatch):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
with L._force_hf_offline():
|
||||
assert os.environ.get("HF_HUB_OFFLINE") == "1"
|
||||
assert os.environ.get("TRANSFORMERS_OFFLINE") == "1"
|
||||
# Absent before -> absent after (not left as "1").
|
||||
assert os.environ.get("HF_HUB_OFFLINE") is None
|
||||
assert os.environ.get("TRANSFORMERS_OFFLINE") is None
|
||||
|
||||
|
||||
def test_force_offline_preserves_prior_env_value(monkeypatch):
|
||||
monkeypatch.setenv("HF_HUB_OFFLINE", "0")
|
||||
with L._force_hf_offline():
|
||||
assert os.environ.get("HF_HUB_OFFLINE") == "1"
|
||||
assert os.environ.get("HF_HUB_OFFLINE") == "0"
|
||||
|
||||
|
||||
def test_force_offline_flips_inprocess_constants():
|
||||
before = _inprocess_offline_flags()
|
||||
with L._force_hf_offline():
|
||||
during = _inprocess_offline_flags()
|
||||
assert during, "expected at least one in-process offline flag to inspect"
|
||||
assert all(flag is True for flag in during)
|
||||
assert _inprocess_offline_flags() == before
|
||||
|
||||
|
||||
def test_force_offline_nesting_shares_one_flip(monkeypatch):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
with L._force_hf_offline():
|
||||
with L._force_hf_offline():
|
||||
assert os.environ.get("HF_HUB_OFFLINE") == "1"
|
||||
# Inner exit must NOT restore while the outer window is still open.
|
||||
assert os.environ.get("HF_HUB_OFFLINE") == "1"
|
||||
assert os.environ.get("HF_HUB_OFFLINE") is None
|
||||
|
||||
|
||||
def test_force_offline_restores_on_exception(monkeypatch):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
with pytest.raises(RuntimeError):
|
||||
with L._force_hf_offline():
|
||||
assert os.environ.get("HF_HUB_OFFLINE") == "1"
|
||||
raise RuntimeError("boom")
|
||||
assert os.environ.get("HF_HUB_OFFLINE") is None
|
||||
assert os.environ.get("TRANSFORMERS_OFFLINE") is None
|
||||
|
||||
|
||||
def test_force_offline_depth_returns_to_zero():
|
||||
assert L._force_offline_depth == 0
|
||||
with L._force_hf_offline():
|
||||
assert L._force_offline_depth == 1
|
||||
assert L._force_offline_depth == 0
|
||||
|
||||
|
||||
def test_reset_hf_sessions_is_safe():
|
||||
# Best-effort no-op when the hub helper is missing; must never raise.
|
||||
L._reset_hf_sessions()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _has_local_tokenizer_files / _resolve_checkpoint_tokenizer_name
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _touch(path, name):
|
||||
open(os.path.join(path, name), "w").close()
|
||||
|
||||
|
||||
def test_has_local_tokenizer_json(tmp_path):
|
||||
_touch(tmp_path, "tokenizer.json")
|
||||
assert L._has_local_tokenizer_files(str(tmp_path)) is True
|
||||
|
||||
|
||||
def test_has_local_tokenizer_model(tmp_path):
|
||||
_touch(tmp_path, "tokenizer.model")
|
||||
assert L._has_local_tokenizer_files(str(tmp_path)) is True
|
||||
|
||||
|
||||
def test_has_local_tokenizer_bpe_needs_merges(tmp_path):
|
||||
# vocab.json alone is not loadable BPE; it needs merges.txt.
|
||||
_touch(tmp_path, "vocab.json")
|
||||
assert L._has_local_tokenizer_files(str(tmp_path)) is False
|
||||
_touch(tmp_path, "merges.txt")
|
||||
assert L._has_local_tokenizer_files(str(tmp_path)) is True
|
||||
|
||||
|
||||
def test_has_local_tokenizer_empty_dir(tmp_path):
|
||||
assert L._has_local_tokenizer_files(str(tmp_path)) is False
|
||||
|
||||
|
||||
def test_resolve_tokenizer_explicit_override_wins(tmp_path):
|
||||
kwargs = {"tokenizer_name": "base/repo"}
|
||||
assert L._resolve_checkpoint_tokenizer_name(str(tmp_path), kwargs) == "base/repo"
|
||||
# tokenizer_name is always popped (it is passed explicitly downstream too).
|
||||
assert "tokenizer_name" not in kwargs
|
||||
|
||||
|
||||
def test_resolve_tokenizer_self_sufficient_dir(tmp_path):
|
||||
_touch(tmp_path, "tokenizer_config.json")
|
||||
_touch(tmp_path, "tokenizer.json")
|
||||
kwargs = {}
|
||||
assert L._resolve_checkpoint_tokenizer_name(str(tmp_path), kwargs) == str(tmp_path)
|
||||
|
||||
|
||||
def test_resolve_tokenizer_config_without_files_falls_back(tmp_path):
|
||||
# Has tokenizer_config.json but no loadable tokenizer file -> base repo.
|
||||
_touch(tmp_path, "tokenizer_config.json")
|
||||
assert L._resolve_checkpoint_tokenizer_name(str(tmp_path), {}) is None
|
||||
|
||||
|
||||
def test_resolve_tokenizer_nonexistent_dir_falls_back():
|
||||
assert L._resolve_checkpoint_tokenizer_name("/no/such/dir", {}) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _offline_aware_load (the retry orchestrator)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_retry_once_on_offline_error_then_succeed(monkeypatch):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
calls = []
|
||||
|
||||
@L._offline_aware_load
|
||||
def fake(*args, **kwargs):
|
||||
calls.append(dict(kwargs))
|
||||
if len(calls) == 1:
|
||||
raise ConnectionError("network down")
|
||||
return "ok"
|
||||
|
||||
assert fake("model") == "ok"
|
||||
assert len(calls) == 2
|
||||
assert not calls[0].get("local_files_only")
|
||||
assert calls[1].get("local_files_only") is True
|
||||
assert L._force_offline_depth == 0
|
||||
|
||||
|
||||
def test_no_retry_on_non_offline_error(monkeypatch):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
calls = []
|
||||
|
||||
@L._offline_aware_load
|
||||
def fake(*args, **kwargs):
|
||||
calls.append(1)
|
||||
raise ValueError("genuine bug, not a network issue")
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
fake("model")
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
def test_no_retry_when_already_offline_via_kwarg(monkeypatch):
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
calls = []
|
||||
|
||||
@L._offline_aware_load
|
||||
def fake(*args, **kwargs):
|
||||
calls.append(dict(kwargs))
|
||||
# Offline window is active for the single attempt.
|
||||
assert os.environ.get("HF_HUB_OFFLINE") == "1"
|
||||
return "ok"
|
||||
|
||||
assert fake("model", local_files_only = True) == "ok"
|
||||
assert len(calls) == 1
|
||||
assert L._force_offline_depth == 0
|
||||
|
||||
|
||||
def test_offline_error_when_already_offline_propagates(monkeypatch):
|
||||
# Already offline -> no online attempt to retry, so the error propagates once.
|
||||
monkeypatch.setenv("HF_HUB_OFFLINE", "1")
|
||||
calls = []
|
||||
|
||||
@L._offline_aware_load
|
||||
def fake(*args, **kwargs):
|
||||
calls.append(1)
|
||||
raise ConnectionError("still down")
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
fake("model")
|
||||
assert len(calls) == 1
|
||||
assert L._force_offline_depth == 0
|
||||
|
||||
|
||||
def test_kwargs_preserved_across_retry(monkeypatch):
|
||||
# Callee popping config/tokenizer_name must not change what the retry sees:
|
||||
# fn(*args, **kwargs) re-packs a fresh **kwargs per call.
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
seen = []
|
||||
|
||||
@L._offline_aware_load
|
||||
def fake(model_name, **kwargs):
|
||||
cfg = kwargs.pop("config", None)
|
||||
tok = kwargs.pop("tokenizer_name", None)
|
||||
seen.append((cfg, tok))
|
||||
if len(seen) == 1:
|
||||
raise ConnectionError("down")
|
||||
return cfg, tok
|
||||
|
||||
assert fake("m", config = "CFG", tokenizer_name = "TOK") == ("CFG", "TOK")
|
||||
assert seen == [("CFG", "TOK"), ("CFG", "TOK")]
|
||||
|
||||
|
||||
def test_retry_runs_gc_collect_between_attempts(monkeypatch):
|
||||
# The retry lives OUTSIDE the except so the failed attempt's traceback (a
|
||||
# partial model) is freed by gc.collect() before the second load reallocates.
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
gc_calls = []
|
||||
monkeypatch.setattr(L.gc, "collect", lambda *a, **k: gc_calls.append(1))
|
||||
calls = []
|
||||
|
||||
@L._offline_aware_load
|
||||
def fake(*args, **kwargs):
|
||||
calls.append(1)
|
||||
if len(calls) == 1:
|
||||
raise ConnectionError("down")
|
||||
# By the retry attempt, gc.collect() must already have fired.
|
||||
assert gc_calls, "gc.collect must run before the offline retry"
|
||||
return "ok"
|
||||
|
||||
gc_calls.clear()
|
||||
assert fake("model") == "ok"
|
||||
assert len(calls) == 2
|
||||
assert len(gc_calls) == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _force_hf_offline — constant restore (no stale offline pin)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_force_offline_restores_freshly_imported_constant(monkeypatch):
|
||||
# If huggingface_hub.constants is first imported inside the window, the saved value must
|
||||
# be the pre-window state, not the just-forced "1"; otherwise the process pins offline.
|
||||
import sys
|
||||
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
|
||||
saved_mod = sys.modules.get("huggingface_hub.constants")
|
||||
saved_val = getattr(saved_mod, "HF_HUB_OFFLINE", None) if saved_mod else None
|
||||
try:
|
||||
sys.modules.pop("huggingface_hub.constants", None) # simulate "not imported yet"
|
||||
with L._force_hf_offline():
|
||||
import huggingface_hub.constants as hfc_in
|
||||
assert hfc_in.HF_HUB_OFFLINE is True # forced offline inside the window
|
||||
import huggingface_hub.constants as hfc_after
|
||||
|
||||
assert hfc_after.HF_HUB_OFFLINE is False # restored, not pinned True
|
||||
assert os.environ.get("HF_HUB_OFFLINE") is None
|
||||
finally:
|
||||
if saved_mod is not None:
|
||||
sys.modules["huggingface_hub.constants"] = saved_mod
|
||||
if saved_val is not None:
|
||||
saved_mod.HF_HUB_OFFLINE = saved_val
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _resolve_checkpoint_tokenizer_name — VLM needs local processor files
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resolve_tokenizer_vlm_without_processor_falls_back(tmp_path):
|
||||
# VLM checkpoint with tokenizer files but no processor config -> base repo (None), so its
|
||||
# cached processor still loads instead of AutoProcessor failing on the local dir.
|
||||
_touch(tmp_path, "tokenizer_config.json")
|
||||
_touch(tmp_path, "tokenizer.json")
|
||||
assert L._resolve_checkpoint_tokenizer_name(str(tmp_path), {}, require_processor = True) is None
|
||||
|
||||
|
||||
def test_resolve_tokenizer_vlm_with_processor_uses_local_dir(tmp_path):
|
||||
_touch(tmp_path, "tokenizer_config.json")
|
||||
_touch(tmp_path, "tokenizer.json")
|
||||
_touch(tmp_path, "preprocessor_config.json")
|
||||
assert L._resolve_checkpoint_tokenizer_name(str(tmp_path), {}, require_processor = True) == str(
|
||||
tmp_path
|
||||
)
|
||||
|
|
@ -2527,144 +2527,148 @@ class FastLlamaModel:
|
|||
kwargs = add_dtype_kwargs(dtype, kwargs)
|
||||
|
||||
raise_handler = RaiseUninitialized()
|
||||
if num_labels is not None:
|
||||
# Transformers 5.x @strict config classes reject unexpected kwargs
|
||||
# like num_labels and max_position_embeddings. Set on the config
|
||||
# object directly and pass config= instead.
|
||||
set_task_config_attr(model_config, "num_labels", num_labels)
|
||||
if max_position_embeddings is not None:
|
||||
model_config.max_position_embeddings = max_position_embeddings
|
||||
# Pop config-level attrs that would be rejected by @strict model init
|
||||
for _cfg_key in ("id2label", "label2id", "rope_scaling"):
|
||||
_cfg_val = kwargs.pop(_cfg_key, None)
|
||||
if _cfg_val is not None:
|
||||
if _cfg_key in ("id2label", "label2id"):
|
||||
set_task_config_attr(model_config, _cfg_key, _cfg_val)
|
||||
else:
|
||||
setattr(model_config, _cfg_key, _cfg_val)
|
||||
model = AutoModelForSequenceClassification.from_pretrained(
|
||||
model_name,
|
||||
config = model_config,
|
||||
device_map = device_map,
|
||||
# torch_dtype = dtype, # transformers changed torch_dtype to dtype
|
||||
# quantization_config = bnb_config,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
attn_implementation = preferred_attn_impl,
|
||||
**kwargs,
|
||||
)
|
||||
# Defensive: ensure the task head is in a floating dtype, guarding
|
||||
# against any path leaving it as integer storage. See unslothai/unsloth#5027.
|
||||
for _head_name in ("score", "classifier", "qa_outputs"):
|
||||
_head = getattr(model, _head_name, None)
|
||||
if (
|
||||
_head is not None
|
||||
and hasattr(_head, "weight")
|
||||
and not _head.weight.is_floating_point()
|
||||
):
|
||||
_head.to(dtype)
|
||||
# Attach dispatch hooks for bnb multi-device loads.
|
||||
from unsloth.models.vision import _attach_bnb_multidevice_hooks
|
||||
|
||||
_attach_bnb_multidevice_hooks(
|
||||
model,
|
||||
load_in_4bit = load_in_4bit,
|
||||
load_in_8bit = kwargs.get("load_in_8bit", False),
|
||||
offload_embedding = False,
|
||||
fast_inference = fast_inference,
|
||||
)
|
||||
elif not fast_inference:
|
||||
if user_config is not None:
|
||||
# Transformers 5.x @strict model init rejects extra kwargs next
|
||||
# to config=; set the override on the config and pass the single
|
||||
# config object through so user overrides reach the actual load.
|
||||
try:
|
||||
if num_labels is not None:
|
||||
# Transformers 5.x @strict config classes reject unexpected kwargs
|
||||
# like num_labels and max_position_embeddings. Set on the config
|
||||
# object directly and pass config= instead.
|
||||
set_task_config_attr(model_config, "num_labels", num_labels)
|
||||
if max_position_embeddings is not None:
|
||||
model_config.max_position_embeddings = max_position_embeddings
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
# Pop config-level attrs that would be rejected by @strict model init
|
||||
for _cfg_key in ("id2label", "label2id", "rope_scaling"):
|
||||
_cfg_val = kwargs.pop(_cfg_key, None)
|
||||
if _cfg_val is not None:
|
||||
if _cfg_key in ("id2label", "label2id"):
|
||||
set_task_config_attr(model_config, _cfg_key, _cfg_val)
|
||||
else:
|
||||
setattr(model_config, _cfg_key, _cfg_val)
|
||||
model = AutoModelForSequenceClassification.from_pretrained(
|
||||
model_name,
|
||||
config = model_config,
|
||||
device_map = device_map,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
attn_implementation = preferred_attn_impl,
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
model_name,
|
||||
device_map = device_map,
|
||||
# torch_dtype = dtype, # transformers changed torch_dtype to dtype
|
||||
# quantization_config = bnb_config,
|
||||
token = token,
|
||||
max_position_embeddings = max_position_embeddings,
|
||||
trust_remote_code = trust_remote_code,
|
||||
attn_implementation = preferred_attn_impl,
|
||||
**kwargs,
|
||||
)
|
||||
# Attach dispatch hooks for bnb multi-device loads.
|
||||
from unsloth.models.vision import _attach_bnb_multidevice_hooks
|
||||
# Defensive: ensure the task head is in a floating dtype, guarding
|
||||
# against any path leaving it as integer storage. See unslothai/unsloth#5027.
|
||||
for _head_name in ("score", "classifier", "qa_outputs"):
|
||||
_head = getattr(model, _head_name, None)
|
||||
if (
|
||||
_head is not None
|
||||
and hasattr(_head, "weight")
|
||||
and not _head.weight.is_floating_point()
|
||||
):
|
||||
_head.to(dtype)
|
||||
# Attach dispatch hooks for bnb multi-device loads.
|
||||
from unsloth.models.vision import _attach_bnb_multidevice_hooks
|
||||
|
||||
_attach_bnb_multidevice_hooks(
|
||||
model,
|
||||
load_in_4bit = load_in_4bit,
|
||||
load_in_8bit = kwargs.get("load_in_8bit", False),
|
||||
offload_embedding = False,
|
||||
fast_inference = False,
|
||||
)
|
||||
model.fast_generate = make_fast_generate_wrapper(model.generate)
|
||||
model.fast_generate_batches = None
|
||||
else:
|
||||
from unsloth_zoo.vllm_utils import (
|
||||
load_vllm,
|
||||
get_vllm_state_dict,
|
||||
convert_vllm_to_huggingface,
|
||||
generate_batches,
|
||||
)
|
||||
_attach_bnb_multidevice_hooks(
|
||||
model,
|
||||
load_in_4bit = load_in_4bit,
|
||||
load_in_8bit = kwargs.get("load_in_8bit", False),
|
||||
offload_embedding = False,
|
||||
fast_inference = fast_inference,
|
||||
)
|
||||
elif not fast_inference:
|
||||
if user_config is not None:
|
||||
# Transformers 5.x @strict model init rejects extra kwargs next
|
||||
# to config=; set the override on the config and pass the single
|
||||
# config object through so user overrides reach the actual load.
|
||||
if max_position_embeddings is not None:
|
||||
model_config.max_position_embeddings = max_position_embeddings
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
model_name,
|
||||
config = model_config,
|
||||
device_map = device_map,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
attn_implementation = preferred_attn_impl,
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
model_name,
|
||||
device_map = device_map,
|
||||
# torch_dtype = dtype, # transformers changed torch_dtype to dtype
|
||||
# quantization_config = bnb_config,
|
||||
token = token,
|
||||
max_position_embeddings = max_position_embeddings,
|
||||
trust_remote_code = trust_remote_code,
|
||||
attn_implementation = preferred_attn_impl,
|
||||
**kwargs,
|
||||
)
|
||||
# Attach dispatch hooks for bnb multi-device loads.
|
||||
from unsloth.models.vision import _attach_bnb_multidevice_hooks
|
||||
|
||||
fp8_mode = None
|
||||
if load_in_fp8 != False:
|
||||
fp8_mode = _get_fp8_mode_and_check_settings(
|
||||
load_in_fp8,
|
||||
fast_inference,
|
||||
_attach_bnb_multidevice_hooks(
|
||||
model,
|
||||
load_in_4bit = load_in_4bit,
|
||||
load_in_8bit = kwargs.get("load_in_8bit", False),
|
||||
offload_embedding = False,
|
||||
fast_inference = False,
|
||||
)
|
||||
model.fast_generate = make_fast_generate_wrapper(model.generate)
|
||||
model.fast_generate_batches = None
|
||||
else:
|
||||
from unsloth_zoo.vllm_utils import (
|
||||
load_vllm,
|
||||
get_vllm_state_dict,
|
||||
convert_vllm_to_huggingface,
|
||||
generate_batches,
|
||||
)
|
||||
|
||||
allowed_args = inspect.getfullargspec(load_vllm).args
|
||||
load_vllm_kwargs = dict(
|
||||
model_name = model_name,
|
||||
config = model_config,
|
||||
gpu_memory_utilization = gpu_memory_utilization,
|
||||
max_seq_length = max_seq_length,
|
||||
dtype = dtype,
|
||||
float8_kv_cache = float8_kv_cache,
|
||||
enable_lora = True,
|
||||
max_lora_rank = max_lora_rank,
|
||||
disable_log_stats = disable_log_stats,
|
||||
use_bitsandbytes = load_in_4bit,
|
||||
unsloth_vllm_standby = unsloth_vllm_standby,
|
||||
fp8_mode = fp8_mode,
|
||||
)
|
||||
for allowed_arg in allowed_args:
|
||||
if allowed_arg not in load_vllm_kwargs and allowed_arg in kwargs:
|
||||
load_vllm_kwargs[allowed_arg] = kwargs[allowed_arg]
|
||||
pass
|
||||
fp8_mode = None
|
||||
if load_in_fp8 != False:
|
||||
fp8_mode = _get_fp8_mode_and_check_settings(
|
||||
load_in_fp8,
|
||||
fast_inference,
|
||||
)
|
||||
|
||||
# Load vLLM first
|
||||
llm = load_vllm(**load_vllm_kwargs)
|
||||
allowed_args = inspect.getfullargspec(load_vllm).args
|
||||
load_vllm_kwargs = dict(
|
||||
model_name = model_name,
|
||||
config = model_config,
|
||||
gpu_memory_utilization = gpu_memory_utilization,
|
||||
max_seq_length = max_seq_length,
|
||||
dtype = dtype,
|
||||
float8_kv_cache = float8_kv_cache,
|
||||
enable_lora = True,
|
||||
max_lora_rank = max_lora_rank,
|
||||
disable_log_stats = disable_log_stats,
|
||||
use_bitsandbytes = load_in_4bit,
|
||||
unsloth_vllm_standby = unsloth_vllm_standby,
|
||||
fp8_mode = fp8_mode,
|
||||
)
|
||||
for allowed_arg in allowed_args:
|
||||
if allowed_arg not in load_vllm_kwargs and allowed_arg in kwargs:
|
||||
load_vllm_kwargs[allowed_arg] = kwargs[allowed_arg]
|
||||
pass
|
||||
|
||||
# Convert to HF format
|
||||
_, quant_state_dict = get_vllm_state_dict(
|
||||
llm,
|
||||
config = model_config,
|
||||
load_in_fp8 = load_in_fp8,
|
||||
)
|
||||
model = convert_vllm_to_huggingface(quant_state_dict, model_config, dtype, bnb_config)
|
||||
model.vllm_engine = llm
|
||||
llm.shared_weights = True
|
||||
model.fast_generate = model.vllm_engine.generate
|
||||
model.fast_generate_batches = functools.partial(generate_batches, model.vllm_engine)
|
||||
raise_handler.remove()
|
||||
# Return old flag
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = old_hf_transfer
|
||||
# Load vLLM first
|
||||
llm = load_vllm(**load_vllm_kwargs)
|
||||
|
||||
# Convert to HF format
|
||||
_, quant_state_dict = get_vllm_state_dict(
|
||||
llm,
|
||||
config = model_config,
|
||||
load_in_fp8 = load_in_fp8,
|
||||
)
|
||||
model = convert_vllm_to_huggingface(
|
||||
quant_state_dict, model_config, dtype, bnb_config
|
||||
)
|
||||
model.vllm_engine = llm
|
||||
llm.shared_weights = True
|
||||
model.fast_generate = model.vllm_engine.generate
|
||||
model.fast_generate_batches = functools.partial(generate_batches, model.vllm_engine)
|
||||
finally:
|
||||
raise_handler.remove()
|
||||
# Return old flag
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = old_hf_transfer
|
||||
|
||||
# Counteract saved tokenizers
|
||||
tokenizer_name = model_name if tokenizer_name is None else tokenizer_name
|
||||
|
|
|
|||
|
|
@ -38,6 +38,9 @@ from .loader_utils import (
|
|||
_tag_model_with_fp8_torchao_config,
|
||||
get_model_name,
|
||||
prepare_device_map,
|
||||
_offline_aware_load,
|
||||
_resolve_checkpoint_tokenizer_name,
|
||||
_is_offline_related_error,
|
||||
)
|
||||
import os, contextlib, sys
|
||||
|
||||
|
|
@ -284,6 +287,7 @@ def _fix_rope_inv_freq(model):
|
|||
|
||||
class FastLanguageModel(FastLlamaModel):
|
||||
@staticmethod
|
||||
@_offline_aware_load
|
||||
def from_pretrained(
|
||||
model_name = "unsloth/Llama-3.2-1B-Instruct",
|
||||
max_seq_length = 2048,
|
||||
|
|
@ -357,16 +361,7 @@ class FastLanguageModel(FastLlamaModel):
|
|||
if is_dist:
|
||||
device_map = distributed_device_map
|
||||
|
||||
# Honour offline env vars BEFORE FastModel delegation so 8bit /
|
||||
# full-finetuning / qat paths also receive local_files_only.
|
||||
if not kwargs.get("local_files_only", False):
|
||||
_offline = {"1", "true", "yes", "on"}
|
||||
if (
|
||||
os.environ.get("TRANSFORMERS_OFFLINE", "").strip().lower() in _offline
|
||||
or os.environ.get("HF_HUB_OFFLINE", "").strip().lower() in _offline
|
||||
):
|
||||
kwargs["local_files_only"] = True
|
||||
|
||||
# @_offline_aware_load already forced offline when needed; delegations inherit it.
|
||||
if load_in_8bit or full_finetuning or qat_scheme is not None:
|
||||
return FastModel.from_pretrained(
|
||||
model_name = model_name,
|
||||
|
|
@ -496,6 +491,8 @@ class FastLanguageModel(FastLlamaModel):
|
|||
|
||||
autoconfig_error = None
|
||||
peft_error = None
|
||||
autoconfig_exc = None
|
||||
peft_exc = None
|
||||
model_config = None
|
||||
peft_config = None
|
||||
local_files_only = kwargs.get("local_files_only", False)
|
||||
|
|
@ -513,6 +510,7 @@ class FastLanguageModel(FastLlamaModel):
|
|||
raise
|
||||
except Exception as error:
|
||||
autoconfig_error = str(error)
|
||||
autoconfig_exc = error
|
||||
if "architecture" in autoconfig_error:
|
||||
if "qwen3_5" in autoconfig_error:
|
||||
raise ImportError(
|
||||
|
|
@ -539,6 +537,7 @@ class FastLanguageModel(FastLlamaModel):
|
|||
raise
|
||||
except Exception as error:
|
||||
peft_error = str(error)
|
||||
peft_exc = error
|
||||
if "architecture" in peft_error:
|
||||
raise ValueError(
|
||||
f"`{model_name}` is not supported yet in `transformers=={transformers_version}`.\n"
|
||||
|
|
@ -557,6 +556,34 @@ class FastLanguageModel(FastLlamaModel):
|
|||
"We must only allow one config file.\n"
|
||||
"Please separate the LoRA and base models to 2 repos."
|
||||
)
|
||||
if not is_model and not is_peft:
|
||||
error = autoconfig_error if autoconfig_error is not None else peft_error
|
||||
# Old transformers version
|
||||
if "rope_scaling" in error.lower() and not SUPPORTS_LLAMA31:
|
||||
raise ImportError(
|
||||
f"Unsloth: Your transformers version of {transformers_version} does not support new RoPE scaling methods.\n"
|
||||
f"This includes Llama 3.1. The minimum required version is 4.43.2\n"
|
||||
f'Try `pip install --upgrade "transformers>=4.43.2"`\n'
|
||||
f"to obtain the latest transformers build, then restart this session."
|
||||
)
|
||||
# Create a combined error message showing both failures
|
||||
combined_error = (
|
||||
"Unsloth: Failed to load model. Both AutoConfig and PeftConfig loading failed.\n\n"
|
||||
f"AutoConfig error: {autoconfig_error}\n\n"
|
||||
f"PeftConfig error: {peft_error}\n\n"
|
||||
)
|
||||
# Chain an offline-related cause if either probe had one, so @_offline_aware_load
|
||||
# still retries from cache (e.g. adapter repo: permanent AutoConfig 404 + transient PeftConfig).
|
||||
_cause = next(
|
||||
(
|
||||
e
|
||||
for e in (autoconfig_exc, peft_exc)
|
||||
if e is not None and _is_offline_related_error(e)
|
||||
),
|
||||
autoconfig_exc or peft_exc,
|
||||
)
|
||||
raise RuntimeError(combined_error) from _cause
|
||||
|
||||
model_types = get_transformers_model_type(
|
||||
peft_config if peft_config is not None else model_config,
|
||||
trust_remote_code = trust_remote_code,
|
||||
|
|
@ -582,24 +609,6 @@ class FastLanguageModel(FastLlamaModel):
|
|||
# definitely exist -- no need for an extra HfFileSystem network call.
|
||||
both_exist = True
|
||||
|
||||
if not is_model and not is_peft:
|
||||
error = autoconfig_error if autoconfig_error is not None else peft_error
|
||||
# Old transformers version
|
||||
if "rope_scaling" in error.lower() and not SUPPORTS_LLAMA31:
|
||||
raise ImportError(
|
||||
f"Unsloth: Your transformers version of {transformers_version} does not support new RoPE scaling methods.\n"
|
||||
f"This includes Llama 3.1. The minimum required version is 4.43.2\n"
|
||||
f'Try `pip install --upgrade "transformers>=4.43.2"`\n'
|
||||
f"to obtain the latest transformers build, then restart this session."
|
||||
)
|
||||
# Create a combined error message showing both failures
|
||||
combined_error = (
|
||||
"Unsloth: Failed to load model. Both AutoConfig and PeftConfig loading failed.\n\n"
|
||||
f"AutoConfig error: {autoconfig_error}\n\n"
|
||||
f"PeftConfig error: {peft_error}\n\n"
|
||||
)
|
||||
raise RuntimeError(combined_error)
|
||||
|
||||
# Get base model for PEFT:
|
||||
if is_peft:
|
||||
# Check base model again for PEFT
|
||||
|
|
@ -755,15 +764,8 @@ class FastLanguageModel(FastLlamaModel):
|
|||
use_gradient_checkpointing, max_seq_length, dtype
|
||||
)
|
||||
|
||||
# Check if this is local model since the tokenizer gets overwritten
|
||||
if (
|
||||
os.path.exists(os.path.join(old_model_name, "tokenizer_config.json"))
|
||||
and os.path.exists(os.path.join(old_model_name, "tokenizer.json"))
|
||||
and os.path.exists(os.path.join(old_model_name, "special_tokens_map.json"))
|
||||
):
|
||||
tokenizer_name = old_model_name
|
||||
else:
|
||||
tokenizer_name = kwargs.pop("tokenizer_name", None)
|
||||
# Keep the local checkpoint dir as tokenizer when self-sufficient (see _resolve_checkpoint_tokenizer_name).
|
||||
tokenizer_name = _resolve_checkpoint_tokenizer_name(old_model_name, kwargs)
|
||||
|
||||
if fast_inference:
|
||||
fast_inference, model_name = fast_inference_setup(model_name, model_config)
|
||||
|
|
@ -867,6 +869,7 @@ class FastLanguageModel(FastLlamaModel):
|
|||
old_model_name,
|
||||
token = token,
|
||||
revision = revision,
|
||||
local_files_only = local_files_only,
|
||||
is_trainable = True,
|
||||
trust_remote_code = trust_remote_code,
|
||||
)
|
||||
|
|
@ -928,6 +931,7 @@ class FastModel(FastBaseModel):
|
|||
return FastBaseModel.for_training(model, use_gradient_checkpointing)
|
||||
|
||||
@staticmethod
|
||||
@_offline_aware_load
|
||||
def from_pretrained(
|
||||
model_name = "unsloth/Llama-3.2-11B-Vision-Instruct-bnb-4bit",
|
||||
max_seq_length = 2048,
|
||||
|
|
@ -1134,18 +1138,12 @@ class FastModel(FastBaseModel):
|
|||
|
||||
autoconfig_error = None
|
||||
peft_error = None
|
||||
autoconfig_exc = None
|
||||
peft_exc = None
|
||||
model_config = None
|
||||
peft_config = None
|
||||
# @_offline_aware_load already forced offline when needed; nested calls inherit it.
|
||||
local_files_only = kwargs.get("local_files_only", False)
|
||||
# Mirror env-var fallback for direct callers (FastVisionModel / FastTextModel).
|
||||
if not local_files_only:
|
||||
_offline = {"1", "true", "yes", "on"}
|
||||
if (
|
||||
os.environ.get("TRANSFORMERS_OFFLINE", "").strip().lower() in _offline
|
||||
or os.environ.get("HF_HUB_OFFLINE", "").strip().lower() in _offline
|
||||
):
|
||||
local_files_only = True
|
||||
kwargs["local_files_only"] = True
|
||||
|
||||
# Text-diffusion slow-path dispatch, factored so both the normal route (below) and the
|
||||
# legacy-config fallback (in the AutoConfig except handler) share one call site.
|
||||
|
|
@ -1180,6 +1178,7 @@ class FastModel(FastBaseModel):
|
|||
raise
|
||||
except Exception as error:
|
||||
autoconfig_error = str(error)
|
||||
autoconfig_exc = error
|
||||
# Legacy text-diffusion configs use model_type "diffusion_gemma", which current
|
||||
# transformers does not register by name (it ships "diffusion_gemma4"). AutoConfig
|
||||
# raises before we can dispatch; route straight to the diffusion slow path, whose
|
||||
|
|
@ -1212,6 +1211,7 @@ class FastModel(FastBaseModel):
|
|||
raise
|
||||
except Exception as error:
|
||||
peft_error = str(error)
|
||||
peft_exc = error
|
||||
if "architecture" in peft_error:
|
||||
raise ValueError(
|
||||
f"`{model_name}` is not supported yet in `transformers=={transformers_version}`.\n"
|
||||
|
|
@ -1228,6 +1228,34 @@ class FastModel(FastBaseModel):
|
|||
"We must only allow one config file.\n"
|
||||
"Please separate the LoRA and base models to 2 repos."
|
||||
)
|
||||
if not is_model and not is_peft:
|
||||
error = autoconfig_error if autoconfig_error is not None else peft_error
|
||||
# Old transformers version
|
||||
if "rope_scaling" in error.lower() and not SUPPORTS_LLAMA31:
|
||||
raise ImportError(
|
||||
f"Unsloth: Your transformers version of {transformers_version} does not support new RoPE scaling methods.\n"
|
||||
f"This includes Llama 3.1. The minimum required version is 4.43.2\n"
|
||||
f'Try `pip install --upgrade "transformers>=4.43.2"`\n'
|
||||
f"to obtain the latest transformers build, then restart this session."
|
||||
)
|
||||
# Create a combined error message showing both failures
|
||||
combined_error = (
|
||||
"Unsloth: Failed to load model. Both AutoConfig and PeftConfig loading failed.\n\n"
|
||||
f"AutoConfig error: {autoconfig_error}\n\n"
|
||||
f"PeftConfig error: {peft_error}\n\n"
|
||||
)
|
||||
# Chain an offline-related cause if either probe had one, so @_offline_aware_load
|
||||
# still retries from cache (e.g. adapter repo: permanent AutoConfig 404 + transient PeftConfig).
|
||||
_cause = next(
|
||||
(
|
||||
e
|
||||
for e in (autoconfig_exc, peft_exc)
|
||||
if e is not None and _is_offline_related_error(e)
|
||||
),
|
||||
autoconfig_exc or peft_exc,
|
||||
)
|
||||
raise RuntimeError(combined_error) from _cause
|
||||
|
||||
model_types = get_transformers_model_type(
|
||||
peft_config if peft_config is not None else model_config,
|
||||
trust_remote_code = trust_remote_code,
|
||||
|
|
@ -1432,24 +1460,6 @@ class FastModel(FastBaseModel):
|
|||
# definitely exist -- no need for an extra HfFileSystem network call.
|
||||
both_exist = True
|
||||
|
||||
if not is_model and not is_peft:
|
||||
error = autoconfig_error if autoconfig_error is not None else peft_error
|
||||
# Old transformers version
|
||||
if "rope_scaling" in error.lower() and not SUPPORTS_LLAMA31:
|
||||
raise ImportError(
|
||||
f"Unsloth: Your transformers version of {transformers_version} does not support new RoPE scaling methods.\n"
|
||||
f"This includes Llama 3.1. The minimum required version is 4.43.2\n"
|
||||
f'Try `pip install --upgrade "transformers>=4.43.2"`\n'
|
||||
f"to obtain the latest transformers build, then restart this session."
|
||||
)
|
||||
# Create a combined error message showing both failures
|
||||
combined_error = (
|
||||
"Unsloth: Failed to load model. Both AutoConfig and PeftConfig loading failed.\n\n"
|
||||
f"AutoConfig error: {autoconfig_error}\n\n"
|
||||
f"PeftConfig error: {peft_error}\n\n"
|
||||
)
|
||||
raise RuntimeError(combined_error)
|
||||
|
||||
# Get base model for PEFT:
|
||||
if is_peft:
|
||||
# Check base model again for PEFT
|
||||
|
|
@ -1547,15 +1557,16 @@ class FastModel(FastBaseModel):
|
|||
if model_type in model_types_all:
|
||||
supports_sdpa = False
|
||||
|
||||
# Check if this is local model since the tokenizer gets overwritten
|
||||
if (
|
||||
os.path.exists(os.path.join(old_model_name, "tokenizer_config.json"))
|
||||
and os.path.exists(os.path.join(old_model_name, "tokenizer.json"))
|
||||
and os.path.exists(os.path.join(old_model_name, "special_tokens_map.json"))
|
||||
):
|
||||
tokenizer_name = old_model_name
|
||||
else:
|
||||
tokenizer_name = kwargs.pop("tokenizer_name", None)
|
||||
# Keep the local checkpoint dir as tokenizer when self-sufficient (see
|
||||
# _resolve_checkpoint_tokenizer_name). A VLM also needs local processor files, else
|
||||
# we fall back to the base repo so its cached processor loads.
|
||||
_ckpt_arch = getattr(model_config, "architectures", None) or []
|
||||
_ckpt_is_vlm = any(x.endswith("ForConditionalGeneration") for x in _ckpt_arch) or hasattr(
|
||||
model_config, "vision_config"
|
||||
)
|
||||
tokenizer_name = _resolve_checkpoint_tokenizer_name(
|
||||
old_model_name, kwargs, require_processor = _ckpt_is_vlm
|
||||
)
|
||||
|
||||
# Capture task intent before text_only can replace a parent VLM config
|
||||
# with its nested text config.
|
||||
|
|
@ -1783,6 +1794,7 @@ class FastModel(FastBaseModel):
|
|||
old_model_name,
|
||||
token = token,
|
||||
revision = revision,
|
||||
local_files_only = local_files_only,
|
||||
is_trainable = True,
|
||||
trust_remote_code = trust_remote_code,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -18,6 +18,9 @@ import os
|
|||
import torch
|
||||
import re
|
||||
import tempfile
|
||||
import contextlib
|
||||
import threading as _threading
|
||||
import functools
|
||||
from typing import Union
|
||||
from .mapper import (
|
||||
INT_TO_FLOAT_MAPPER,
|
||||
|
|
@ -237,7 +240,12 @@ def get_model_name(
|
|||
):
|
||||
new_model_name = BAD_MAPPINGS[new_model_name.lower()]
|
||||
|
||||
if new_model_name is None and model_name.count("/") == 1 and model_name[0].isalnum():
|
||||
if (
|
||||
new_model_name is None
|
||||
and model_name.count("/") == 1
|
||||
and model_name[0].isalnum()
|
||||
and not _env_says_offline() # offline: skip the remote (raw GitHub) mapper refresh
|
||||
):
|
||||
# Try checking if a new Unsloth version allows it!
|
||||
NEW_INT_TO_FLOAT_MAPPER, NEW_FLOAT_TO_INT_MAPPER, NEW_MAP_TO_UNSLOTH_16bit = (
|
||||
_get_new_mapper()
|
||||
|
|
@ -489,3 +497,330 @@ def _get_fp8_mode_and_check_settings(
|
|||
f"Using Triton kernels instead."
|
||||
)
|
||||
return fp8_mode
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Offline loading - single source of truth (shared by vision.py, loader.py and
|
||||
# the Studio exporter). Decide offline ONCE at the load boundary and force it
|
||||
# ONCE around the whole load, so every nested HF call inherits it.
|
||||
# =============================================================================
|
||||
|
||||
_OFFLINE_ENV_VALUES = {"1", "true", "yes", "on"}
|
||||
_OFFLINE_ENV_KEYS = ("HF_HUB_OFFLINE", "TRANSFORMERS_OFFLINE")
|
||||
|
||||
|
||||
def _env_says_offline():
|
||||
"""True if an HF offline env var is set to a truthy value."""
|
||||
return any(
|
||||
os.environ.get(_k, "").strip().lower() in _OFFLINE_ENV_VALUES for _k in _OFFLINE_ENV_KEYS
|
||||
)
|
||||
|
||||
|
||||
def _get_effective_local_files_only(kwargs):
|
||||
"""Offline if local_files_only is truthy or an HF offline env var is set. Read-only."""
|
||||
if kwargs.get("local_files_only", None):
|
||||
return True
|
||||
return _env_says_offline()
|
||||
|
||||
|
||||
def _is_offline_related_error(exc):
|
||||
"""True if exc (or its cause/context chain) is a lost-connection error, not a
|
||||
missing file. Plain FileNotFoundError propagates; LocalEntryNotFoundError is offline."""
|
||||
import socket
|
||||
import ssl
|
||||
import urllib.error
|
||||
|
||||
# Match network failures by type (locale independent), not just message wording.
|
||||
_net_types = [ConnectionError, TimeoutError, socket.gaierror, urllib.error.URLError]
|
||||
_offline_fnf_types = () # FileNotFoundError subclasses that count as offline
|
||||
# urllib HTTPError is a URLError subclass: judge by status (5xx offline, 4xx propagates).
|
||||
_http_types = (urllib.error.HTTPError,)
|
||||
# TLS/cert failures are security-sensitive (MITM, expired CA): never offline-retry them.
|
||||
_ssl_types = [ssl.SSLError]
|
||||
try:
|
||||
import requests
|
||||
|
||||
_net_types += [requests.exceptions.ConnectionError, requests.exceptions.Timeout]
|
||||
_http_types += (requests.exceptions.HTTPError,)
|
||||
_ssl_types.append(requests.exceptions.SSLError)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
from huggingface_hub.errors import (
|
||||
OfflineModeIsEnabled,
|
||||
HfHubHTTPError,
|
||||
LocalEntryNotFoundError,
|
||||
)
|
||||
|
||||
_net_types += [OfflineModeIsEnabled, LocalEntryNotFoundError]
|
||||
_offline_fnf_types = (LocalEntryNotFoundError,)
|
||||
_http_types += (HfHubHTTPError,)
|
||||
except Exception:
|
||||
pass
|
||||
_net_types = tuple(_net_types)
|
||||
_ssl_types = tuple(_ssl_types)
|
||||
|
||||
def _http_status(e):
|
||||
resp = getattr(e, "response", None)
|
||||
code = getattr(resp, "status_code", None)
|
||||
if code is None:
|
||||
code = getattr(e, "status_code", None)
|
||||
if code is None:
|
||||
code = getattr(e, "code", None) # urllib.error.HTTPError uses .code
|
||||
try:
|
||||
return int(code)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
_wording = (
|
||||
"couldn't connect",
|
||||
"could not connect",
|
||||
"connection error",
|
||||
"connectionerror",
|
||||
"max retries",
|
||||
"offline",
|
||||
"timed out",
|
||||
"timeout",
|
||||
"couldn't reach",
|
||||
"could not reach",
|
||||
"failed to resolve",
|
||||
"getaddrinfo",
|
||||
"name resolution",
|
||||
"no address associated",
|
||||
"network is unreachable",
|
||||
"connection refused",
|
||||
"we couldn't connect to",
|
||||
"proxyerror",
|
||||
# Raw socket.gaierror DNS wording (Linux / macOS)
|
||||
"name or service not known",
|
||||
"temporary failure in name resolution",
|
||||
"nodename nor servname provided",
|
||||
)
|
||||
seen = set()
|
||||
cur = exc
|
||||
while cur is not None and id(cur) not in seen:
|
||||
seen.add(id(cur))
|
||||
# TLS/cert failure (corporate MITM, expired CA): security-sensitive, never retry from
|
||||
# cache. Skip this node; a deeper cause in the chain may still be a genuine outage.
|
||||
if isinstance(cur, _ssl_types) or isinstance(getattr(cur, "reason", None), _ssl_types):
|
||||
cur = cur.__cause__ or cur.__context__
|
||||
continue
|
||||
is_fnf = isinstance(cur, FileNotFoundError) and not isinstance(cur, _offline_fnf_types)
|
||||
# urllib HTTPError is a URLError (net type) but must be judged by status code below,
|
||||
# unlike LocalEntryNotFoundError (an HfHubHTTPError that is always offline).
|
||||
if (
|
||||
isinstance(cur, _net_types)
|
||||
and not is_fnf
|
||||
and not isinstance(cur, urllib.error.HTTPError)
|
||||
):
|
||||
return True
|
||||
if isinstance(cur, _http_types):
|
||||
code = _http_status(cur)
|
||||
if code is not None and 500 <= code < 600:
|
||||
return True
|
||||
# No status -> wording fallback (coded 4xx already decided above).
|
||||
if code is None and not is_fnf and any(w in str(cur).lower() for w in _wording):
|
||||
return True
|
||||
# OSError wording fallback (HTTP status already decided above).
|
||||
elif isinstance(cur, OSError) and not is_fnf:
|
||||
if any(w in str(cur).lower() for w in _wording):
|
||||
return True
|
||||
cur = cur.__cause__ or cur.__context__
|
||||
return False
|
||||
|
||||
|
||||
# Process-wide HF offline state; the depth counter lets nested windows share one
|
||||
# flip (first entrant saves originals, last exit restores). Lock guards flip/restore.
|
||||
_force_offline_lock = _threading.RLock()
|
||||
_force_offline_depth = 0
|
||||
_force_offline_saved = [] # in-process module attributes
|
||||
_force_offline_saved_env = {} # HF offline env-var originals
|
||||
|
||||
|
||||
def _reset_hf_sessions():
|
||||
"""Clear hub's per-thread cached Sessions so the next rebuilds against the current
|
||||
offline flag. On hub 0.x the offline adapter is baked in at Session creation. Best-effort."""
|
||||
try:
|
||||
from huggingface_hub.utils._http import reset_sessions
|
||||
except Exception:
|
||||
try:
|
||||
from huggingface_hub.utils import reset_sessions
|
||||
except Exception:
|
||||
return
|
||||
try:
|
||||
reset_sessions()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _force_hf_offline():
|
||||
"""Force HF offline for the window. local_files_only alone is not enough
|
||||
(transformers < 5 still pings /api/models), so set BOTH the env vars (cover
|
||||
subprocesses + raw urllib/requests) AND the in-process hub/transformers constants.
|
||||
Process-global; the refcount keeps restore correct under nesting / overlap."""
|
||||
global _force_offline_depth, _force_offline_saved, _force_offline_saved_env
|
||||
with _force_offline_lock:
|
||||
if _force_offline_depth == 0:
|
||||
saved = []
|
||||
saved_env = {}
|
||||
# Snapshot in-process constants BEFORE forcing the env: a module first imported
|
||||
# here would otherwise initialize its constant from the just-set "1" and we would
|
||||
# save (then restore) True, pinning the process offline after the window.
|
||||
try:
|
||||
import huggingface_hub.constants as _hfc
|
||||
if hasattr(_hfc, "HF_HUB_OFFLINE"):
|
||||
saved.append((_hfc, "HF_HUB_OFFLINE", _hfc.HF_HUB_OFFLINE))
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
import transformers.utils.hub as _tuh
|
||||
for _attr in ("_is_offline_mode", "OFFLINE"):
|
||||
if hasattr(_tuh, _attr):
|
||||
saved.append((_tuh, _attr, getattr(_tuh, _attr)))
|
||||
except Exception:
|
||||
pass
|
||||
# Now force the env vars and flip the snapshotted constants to offline.
|
||||
for _k in _OFFLINE_ENV_KEYS:
|
||||
saved_env[_k] = os.environ.get(_k)
|
||||
os.environ[_k] = "1"
|
||||
for _obj, _attr, _ in saved:
|
||||
try:
|
||||
setattr(_obj, _attr, True)
|
||||
except Exception:
|
||||
pass
|
||||
_force_offline_saved = saved
|
||||
_force_offline_saved_env = saved_env
|
||||
# Rebuild cached sessions so they pick up the offline adapter.
|
||||
_reset_hf_sessions()
|
||||
_force_offline_depth += 1
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
with _force_offline_lock:
|
||||
_force_offline_depth -= 1
|
||||
if _force_offline_depth == 0:
|
||||
for obj, attr, val in _force_offline_saved:
|
||||
try:
|
||||
setattr(obj, attr, val)
|
||||
except Exception:
|
||||
pass
|
||||
_force_offline_saved = []
|
||||
for _k, _v in _force_offline_saved_env.items():
|
||||
if _v is None:
|
||||
os.environ.pop(_k, None)
|
||||
else:
|
||||
os.environ[_k] = _v
|
||||
_force_offline_saved_env = {}
|
||||
# Drop offline-mounted sessions so later online calls rebuild for the network.
|
||||
_reset_hf_sessions()
|
||||
|
||||
|
||||
def _progress_bars_were_disabled():
|
||||
"""Snapshot HF progress-bar state (None if unknown); pairs with _restore_progress_bars."""
|
||||
try:
|
||||
from huggingface_hub.utils import are_progress_bars_disabled
|
||||
return are_progress_bars_disabled()
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _restore_progress_bars(were_disabled):
|
||||
"""Re-enable HF progress bars only if a failed attempt left them disabled after they
|
||||
were enabled (a loader disables them around config probes and skips re-enabling on
|
||||
error). No-op if the user had them disabled or the state is unknown."""
|
||||
if were_disabled is False:
|
||||
try:
|
||||
from huggingface_hub.utils import enable_progress_bars
|
||||
enable_progress_bars()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _offline_aware_load(fn):
|
||||
"""Decide offline ONCE (local_files_only kwarg or env) and force it around the
|
||||
whole load. If we started online and hit a network error, retry once forced-offline.
|
||||
Network-up online path is unchanged: no window, no retry."""
|
||||
|
||||
@functools.wraps(fn)
|
||||
def _wrapper(*args, **kwargs):
|
||||
if _get_effective_local_files_only(kwargs):
|
||||
kwargs["local_files_only"] = True
|
||||
with _force_hf_offline():
|
||||
return fn(*args, **kwargs)
|
||||
_pb_were_disabled = _progress_bars_were_disabled() # restore before any retry
|
||||
try:
|
||||
return fn(*args, **kwargs)
|
||||
except Exception as e:
|
||||
# Skip if not network-related, or already retried by a nested decorator
|
||||
# (else outer layers reload the whole model again).
|
||||
if not _is_offline_related_error(e) or getattr(e, "_unsloth_offline_retried", False):
|
||||
raise
|
||||
# Retry OUTSIDE the except so the failed attempt's traceback (a partial model)
|
||||
# is freed before reallocating, else a large VLM can OOM on the second load.
|
||||
try:
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
if hasattr(torch, "xpu") and torch.xpu.is_available():
|
||||
torch.xpu.empty_cache()
|
||||
except Exception:
|
||||
pass
|
||||
# A failed attempt may have left HF progress bars disabled; restore before retry.
|
||||
_restore_progress_bars(_pb_were_disabled)
|
||||
kwargs["local_files_only"] = True
|
||||
try:
|
||||
with _force_hf_offline():
|
||||
return fn(*args, **kwargs)
|
||||
except Exception as e:
|
||||
# Tag so an enclosing _offline_aware_load skips its own redundant retry.
|
||||
try:
|
||||
e._unsloth_offline_retried = True
|
||||
except Exception:
|
||||
pass
|
||||
raise
|
||||
|
||||
return _wrapper
|
||||
|
||||
|
||||
def _has_local_tokenizer_files(path):
|
||||
"""True if a local dir has a loadable tokenizer (BPE vocab.json needs merges.txt;
|
||||
special_tokens_map.json is not required)."""
|
||||
return (
|
||||
os.path.exists(os.path.join(path, "tokenizer.json"))
|
||||
or os.path.exists(os.path.join(path, "tokenizer.model"))
|
||||
or (
|
||||
os.path.exists(os.path.join(path, "vocab.json"))
|
||||
and os.path.exists(os.path.join(path, "merges.txt"))
|
||||
)
|
||||
or os.path.exists(os.path.join(path, "vocab.txt"))
|
||||
or os.path.exists(os.path.join(path, "spiece.model"))
|
||||
)
|
||||
|
||||
|
||||
def _has_local_processor_files(path):
|
||||
"""True if a local dir ships a processor/image-processor config (a VLM needs this to
|
||||
build AutoProcessor; tokenizer files alone are not enough)."""
|
||||
return os.path.exists(os.path.join(path, "processor_config.json")) or os.path.exists(
|
||||
os.path.join(path, "preprocessor_config.json")
|
||||
)
|
||||
|
||||
|
||||
def _resolve_checkpoint_tokenizer_name(
|
||||
old_model_name,
|
||||
kwargs,
|
||||
require_processor = False,
|
||||
):
|
||||
"""tokenizer_name for a PEFT/checkpoint load: caller override, else the local checkpoint
|
||||
dir if self-sufficient, else None (base repo). Always popped from kwargs (also passed
|
||||
explicitly downstream). For a VLM (require_processor), the dir must also ship processor
|
||||
files; otherwise fall back to the base repo whose cached processor still loads."""
|
||||
explicit = kwargs.pop("tokenizer_name", None)
|
||||
if explicit is not None:
|
||||
return explicit
|
||||
has_config = os.path.exists(os.path.join(old_model_name, "tokenizer_config.json"))
|
||||
if not (has_config and _has_local_tokenizer_files(old_model_name)):
|
||||
return None
|
||||
if require_processor and not _has_local_processor_files(old_model_name):
|
||||
return None
|
||||
return old_model_name
|
||||
|
|
|
|||
|
|
@ -449,6 +449,14 @@ def unsloth_base_fast_generate(self, *args, **kwargs):
|
|||
return output
|
||||
|
||||
|
||||
# Offline helpers live in loader_utils.py (shared canonical source).
|
||||
from .loader_utils import (
|
||||
_get_effective_local_files_only,
|
||||
_is_offline_related_error,
|
||||
_offline_aware_load,
|
||||
)
|
||||
|
||||
|
||||
def _missing_torchvision_error(error = None):
|
||||
"""True if a VLM processor failed to load due to missing torchvision (#4202).
|
||||
|
||||
|
|
@ -466,13 +474,18 @@ def _missing_torchvision_error(error = None):
|
|||
return False
|
||||
|
||||
|
||||
def _construct_vlm_processor_fallback(tokenizer_name, model_type, token, trust_remote_code):
|
||||
"""Construct a VLM processor manually when AutoProcessor.from_pretrained fails.
|
||||
|
||||
Some VLMs (e.g., LFM2.5-VL) have tokenizer_class entries that AutoTokenizer
|
||||
cannot resolve. This function loads the image processor and tokenizer separately,
|
||||
sets required special token attributes, and constructs the processor.
|
||||
"""
|
||||
def _construct_vlm_processor_fallback(
|
||||
tokenizer_name,
|
||||
model_type,
|
||||
token,
|
||||
trust_remote_code,
|
||||
local_files_only = False,
|
||||
):
|
||||
"""Build a VLM processor manually when AutoProcessor.from_pretrained fails (some VLMs
|
||||
have unresolvable tokenizer_class entries): load the image processor + tokenizer
|
||||
separately and combine. Returns (processor_or_None, error_or_None) so the caller can
|
||||
tell an offline failure (retry from cache) from a genuine one."""
|
||||
_fb_err = None
|
||||
try:
|
||||
from transformers import AutoImageProcessor, PreTrainedTokenizerFast, AutoConfig
|
||||
from transformers.models.auto.processing_auto import PROCESSOR_MAPPING_NAMES
|
||||
|
|
@ -483,6 +496,7 @@ def _construct_vlm_processor_fallback(tokenizer_name, model_type, token, trust_r
|
|||
tokenizer_name,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
# Load tokenizer via PreTrainedTokenizerFast (bypasses tokenizer_class check)
|
||||
tok = PreTrainedTokenizerFast.from_pretrained(
|
||||
|
|
@ -490,14 +504,35 @@ def _construct_vlm_processor_fallback(tokenizer_name, model_type, token, trust_r
|
|||
padding_side = "left",
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
# Read tokenizer_config.json for model-specific special tokens
|
||||
# Read tokenizer_config.json for special tokens: prefer the local file (offline
|
||||
# / local checkpoint dir), else hf_hub_download with local_files_only forwarded.
|
||||
try:
|
||||
from huggingface_hub import hf_hub_download
|
||||
import json as _json
|
||||
|
||||
config_path = hf_hub_download(tokenizer_name, "tokenizer_config.json", token = token)
|
||||
with open(config_path, "r", encoding = "utf-8") as f:
|
||||
tok_config = json.load(f)
|
||||
tok_config = None
|
||||
_local_cfg = os.path.join(tokenizer_name, "tokenizer_config.json")
|
||||
if os.path.isdir(tokenizer_name):
|
||||
# Local dir: read directly. A missing file raises a clear FileNotFoundError
|
||||
# rather than letting hf_hub_download treat the path as a repo id.
|
||||
if os.path.exists(_local_cfg):
|
||||
with open(_local_cfg, "r", encoding = "utf-8") as f:
|
||||
tok_config = _json.load(f)
|
||||
else:
|
||||
raise FileNotFoundError(
|
||||
f"tokenizer_config.json not found in local directory: {tokenizer_name}"
|
||||
)
|
||||
else:
|
||||
from huggingface_hub import hf_hub_download
|
||||
config_path = hf_hub_download(
|
||||
tokenizer_name,
|
||||
"tokenizer_config.json",
|
||||
token = token,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
with open(config_path, "r", encoding = "utf-8") as f:
|
||||
tok_config = _json.load(f)
|
||||
# Set model-specific special tokens and their IDs
|
||||
for key in (
|
||||
"image_token",
|
||||
|
|
@ -512,8 +547,8 @@ def _construct_vlm_processor_fallback(tokenizer_name, model_type, token, trust_r
|
|||
token_id = tok.convert_tokens_to_ids(tok_config[key])
|
||||
if not hasattr(tok, id_key):
|
||||
setattr(tok, id_key, token_id)
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as _e:
|
||||
_fb_err = _e # remember (non-fatal here); surfaced only if no processor is built
|
||||
|
||||
# Find the processor class - try model_type first, then top-level config model_type
|
||||
proc_class_name = PROCESSOR_MAPPING_NAMES.get(model_type)
|
||||
|
|
@ -525,10 +560,11 @@ def _construct_vlm_processor_fallback(tokenizer_name, model_type, token, trust_r
|
|||
tokenizer_name,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
proc_class_name = PROCESSOR_MAPPING_NAMES.get(config.model_type)
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as _e:
|
||||
_fb_err = _e # surface a network/cache miss so the offline retry can fire
|
||||
|
||||
if proc_class_name is not None:
|
||||
import transformers
|
||||
|
|
@ -540,10 +576,10 @@ def _construct_vlm_processor_fallback(tokenizer_name, model_type, token, trust_r
|
|||
tok, "chat_template", None
|
||||
):
|
||||
processor.chat_template = tok.chat_template
|
||||
return processor
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
return processor, None
|
||||
except Exception as _e:
|
||||
_fb_err = _e
|
||||
return None, _fb_err
|
||||
|
||||
|
||||
def _get_total_transformer_layers(model):
|
||||
|
|
@ -577,6 +613,7 @@ def _get_total_transformer_layers(model):
|
|||
|
||||
class FastBaseModel:
|
||||
@staticmethod
|
||||
@_offline_aware_load
|
||||
def from_pretrained(
|
||||
model_name = "unsloth/Llama-3.2-1B-Instruct",
|
||||
max_seq_length = 2048,
|
||||
|
|
@ -614,6 +651,10 @@ class FastBaseModel:
|
|||
if auto_config is None and user_config is not None:
|
||||
auto_config = user_config
|
||||
|
||||
# Offline snapshot for the loads below; not popped, so the weight load still
|
||||
# reads local_files_only from **kwargs. See _get_effective_local_files_only.
|
||||
local_files_only = _get_effective_local_files_only(kwargs)
|
||||
|
||||
if unsloth_vllm_standby and os.environ.get("UNSLOTH_VLLM_STANDBY", "0") != "1":
|
||||
raise RuntimeError(
|
||||
"Unsloth: UNSLOTH_VLLM_STANDBY is True, but UNSLOTH_VLLM_STANDBY is not set to 1!"
|
||||
|
|
@ -633,6 +674,7 @@ class FastBaseModel:
|
|||
model_name,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
if text_only and hasattr(auto_config, "vision_config"):
|
||||
parent_config = auto_config
|
||||
|
|
@ -811,6 +853,7 @@ class FastBaseModel:
|
|||
model_name,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
model_class = resolve_model_class(auto_model, auto_config)
|
||||
attn_impl = resolve_attention_implementation(
|
||||
|
|
@ -918,6 +961,7 @@ class FastBaseModel:
|
|||
model_name,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
if hasattr(auto_config, "quantization_config"):
|
||||
from transformers.quantizers.auto import (
|
||||
|
|
@ -971,6 +1015,7 @@ class FastBaseModel:
|
|||
model_name,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
_set_attn_impl(auto_config, config_attn_impl)
|
||||
model_config = auto_config
|
||||
|
|
@ -978,151 +1023,152 @@ class FastBaseModel:
|
|||
verify_fp8_support_if_applicable(model_config)
|
||||
|
||||
raise_handler = RaiseUninitialized()
|
||||
if not fast_inference:
|
||||
# Prevent load_in_fp8 from being forwarded into HF internal model loading
|
||||
load_in_fp8 = kwargs.pop("load_in_fp8", None)
|
||||
# Transformers 5.x @strict config classes reject unexpected kwargs.
|
||||
# Move config-level attributes onto the config object directly.
|
||||
_num_labels = kwargs.pop("num_labels", None)
|
||||
if _num_labels is not None:
|
||||
set_task_config_attr(model_config, "num_labels", _num_labels)
|
||||
for _cfg_key in ("id2label", "label2id", "problem_type"):
|
||||
_cfg_val = kwargs.pop(_cfg_key, None)
|
||||
try:
|
||||
if not fast_inference:
|
||||
# Prevent load_in_fp8 from being forwarded into HF internal model loading
|
||||
load_in_fp8 = kwargs.pop("load_in_fp8", None)
|
||||
# Transformers 5.x @strict config classes reject unexpected kwargs.
|
||||
# Move config-level attributes onto the config object directly.
|
||||
_num_labels = kwargs.pop("num_labels", None)
|
||||
if _num_labels is not None:
|
||||
set_task_config_attr(model_config, "num_labels", _num_labels)
|
||||
for _cfg_key in ("id2label", "label2id", "problem_type"):
|
||||
_cfg_val = kwargs.pop(_cfg_key, None)
|
||||
if _cfg_val is not None:
|
||||
set_task_config_attr(model_config, _cfg_key, _cfg_val)
|
||||
_cfg_val = kwargs.pop("max_position_embeddings", None)
|
||||
if _cfg_val is not None:
|
||||
set_task_config_attr(model_config, _cfg_key, _cfg_val)
|
||||
_cfg_val = kwargs.pop("max_position_embeddings", None)
|
||||
if _cfg_val is not None:
|
||||
setattr(model_config, "max_position_embeddings", _cfg_val)
|
||||
model = auto_model.from_pretrained(
|
||||
model_name,
|
||||
config = model_config,
|
||||
device_map = device_map,
|
||||
# torch_dtype = torch_dtype, # Transformers removed torch_dtype
|
||||
# quantization_config = bnb_config,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
# attn_implementation = attn_implementation,
|
||||
**kwargs,
|
||||
)
|
||||
# Attach dispatch hooks for bnb multi-device loads.
|
||||
_attach_bnb_multidevice_hooks(
|
||||
model,
|
||||
load_in_4bit = load_in_4bit,
|
||||
load_in_8bit = load_in_8bit,
|
||||
offload_embedding = offload_embedding,
|
||||
fast_inference = fast_inference,
|
||||
)
|
||||
if hasattr(model, "generate"):
|
||||
model.fast_generate = make_fast_generate_wrapper(model.generate)
|
||||
model.fast_generate_batches = error_out_no_vllm
|
||||
if offload_embedding:
|
||||
if bool(os.environ.get("WSL_DISTRO_NAME") or os.environ.get("WSL_INTEROP")):
|
||||
# WSL doesn't work with offloaded embeddings
|
||||
pass
|
||||
elif os.name == "nt":
|
||||
# Windows doesn't work with offloaded embeddings
|
||||
pass
|
||||
else:
|
||||
embed_tokens = model.get_input_embeddings()
|
||||
nbytes = embed_tokens.weight.numel() * embed_tokens.weight.itemsize
|
||||
ngb = round(nbytes / 1024 / 1024 / 1024, 2)
|
||||
print(f"Unsloth: Offloading embeddings to RAM to save {ngb} GB.")
|
||||
embed_tokens.to("cpu")
|
||||
setattr(model_config, "max_position_embeddings", _cfg_val)
|
||||
model = auto_model.from_pretrained(
|
||||
model_name,
|
||||
config = model_config,
|
||||
device_map = device_map,
|
||||
# torch_dtype = torch_dtype, # Transformers removed torch_dtype
|
||||
# quantization_config = bnb_config,
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
# attn_implementation = attn_implementation,
|
||||
**kwargs,
|
||||
)
|
||||
# Attach dispatch hooks for bnb multi-device loads.
|
||||
_attach_bnb_multidevice_hooks(
|
||||
model,
|
||||
load_in_4bit = load_in_4bit,
|
||||
load_in_8bit = load_in_8bit,
|
||||
offload_embedding = offload_embedding,
|
||||
fast_inference = fast_inference,
|
||||
)
|
||||
if hasattr(model, "generate"):
|
||||
model.fast_generate = make_fast_generate_wrapper(model.generate)
|
||||
model.fast_generate_batches = error_out_no_vllm
|
||||
if offload_embedding:
|
||||
if bool(os.environ.get("WSL_DISTRO_NAME") or os.environ.get("WSL_INTEROP")):
|
||||
# WSL doesn't work with offloaded embeddings
|
||||
pass
|
||||
elif os.name == "nt":
|
||||
# Windows doesn't work with offloaded embeddings
|
||||
pass
|
||||
else:
|
||||
embed_tokens = model.get_input_embeddings()
|
||||
nbytes = embed_tokens.weight.numel() * embed_tokens.weight.itemsize
|
||||
ngb = round(nbytes / 1024 / 1024 / 1024, 2)
|
||||
print(f"Unsloth: Offloading embeddings to RAM to save {ngb} GB.")
|
||||
embed_tokens.to("cpu")
|
||||
|
||||
# Add hooks to move inputs to CPU and back to CUDA
|
||||
# [TODO] Doesn't seem to work!
|
||||
# def pre_hook(module, args):
|
||||
# args[0]._old_device = args[0].device
|
||||
# return (args[0].to("cpu", non_blocking = True))
|
||||
# def post_hook(module, args, output):
|
||||
# old_device = getattr(args[0], "_old_device", "cuda")
|
||||
# return output.to(old_device, non_blocking = True)
|
||||
# embed_tokens.register_forward_pre_hook(pre_hook, prepend = True)
|
||||
# embed_tokens.register_forward_hook (post_hook, prepend = True)
|
||||
# Must free GPU memory otherwise will not free!
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
else:
|
||||
from unsloth_zoo.vllm_utils import (
|
||||
load_vllm,
|
||||
get_vllm_state_dict,
|
||||
convert_vllm_to_huggingface,
|
||||
generate_batches,
|
||||
get_lora_supported_ranks,
|
||||
)
|
||||
|
||||
if full_finetuning:
|
||||
max_lora_rank = max(get_lora_supported_ranks())
|
||||
raise NotImplementedError(
|
||||
"Unsloth: `fast_inference=True` cannot be used together with `full_finetuning=True`.\n"
|
||||
"Reason: fast_inference is optimized for inference-only workflows and "
|
||||
"does not currently support full fine-tuning.\n"
|
||||
"Workaround: disable fast_inference, or use parameter-efficient fine-tuning "
|
||||
f"(e.g. LoRA with rank r={max_lora_rank})."
|
||||
# Add hooks to move inputs to CPU and back to CUDA
|
||||
# [TODO] Doesn't seem to work!
|
||||
# def pre_hook(module, args):
|
||||
# args[0]._old_device = args[0].device
|
||||
# return (args[0].to("cpu", non_blocking = True))
|
||||
# def post_hook(module, args, output):
|
||||
# old_device = getattr(args[0], "_old_device", "cuda")
|
||||
# return output.to(old_device, non_blocking = True)
|
||||
# embed_tokens.register_forward_pre_hook(pre_hook, prepend = True)
|
||||
# embed_tokens.register_forward_hook (post_hook, prepend = True)
|
||||
# Must free GPU memory otherwise will not free!
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
else:
|
||||
from unsloth_zoo.vllm_utils import (
|
||||
load_vllm,
|
||||
get_vllm_state_dict,
|
||||
convert_vllm_to_huggingface,
|
||||
generate_batches,
|
||||
get_lora_supported_ranks,
|
||||
)
|
||||
|
||||
model_config.model_name = model_name
|
||||
if full_finetuning:
|
||||
max_lora_rank = max(get_lora_supported_ranks())
|
||||
raise NotImplementedError(
|
||||
"Unsloth: `fast_inference=True` cannot be used together with `full_finetuning=True`.\n"
|
||||
"Reason: fast_inference is optimized for inference-only workflows and "
|
||||
"does not currently support full fine-tuning.\n"
|
||||
"Workaround: disable fast_inference, or use parameter-efficient fine-tuning "
|
||||
f"(e.g. LoRA with rank r={max_lora_rank})."
|
||||
)
|
||||
|
||||
if fast_inference:
|
||||
fast_inference, model_name = fast_inference_setup(model_name, model_config)
|
||||
model_config.model_name = model_name
|
||||
|
||||
fp8_mode = None
|
||||
if load_in_fp8 != False:
|
||||
fp8_mode = _get_fp8_mode_and_check_settings(
|
||||
load_in_fp8,
|
||||
fast_inference,
|
||||
full_finetuning,
|
||||
load_in_4bit,
|
||||
load_in_8bit,
|
||||
load_in_16bit,
|
||||
if fast_inference:
|
||||
fast_inference, model_name = fast_inference_setup(model_name, model_config)
|
||||
|
||||
fp8_mode = None
|
||||
if load_in_fp8 != False:
|
||||
fp8_mode = _get_fp8_mode_and_check_settings(
|
||||
load_in_fp8,
|
||||
fast_inference,
|
||||
full_finetuning,
|
||||
load_in_4bit,
|
||||
load_in_8bit,
|
||||
load_in_16bit,
|
||||
)
|
||||
|
||||
allowed_args = inspect.getfullargspec(load_vllm).args
|
||||
load_vllm_kwargs = dict(
|
||||
model_name = model_name,
|
||||
config = model_config,
|
||||
gpu_memory_utilization = gpu_memory_utilization,
|
||||
max_seq_length = max_seq_length,
|
||||
dtype = dtype,
|
||||
float8_kv_cache = float8_kv_cache,
|
||||
enable_lora = vllm_enable_lora,
|
||||
max_lora_rank = max_lora_rank,
|
||||
disable_log_stats = disable_log_stats,
|
||||
use_bitsandbytes = load_in_4bit,
|
||||
unsloth_vllm_standby = unsloth_vllm_standby,
|
||||
is_vision_model = is_vlm_config,
|
||||
fp8_mode = fp8_mode,
|
||||
)
|
||||
for allowed_arg in allowed_args:
|
||||
if allowed_arg not in load_vllm_kwargs and allowed_arg in kwargs:
|
||||
load_vllm_kwargs[allowed_arg] = kwargs[allowed_arg]
|
||||
|
||||
allowed_args = inspect.getfullargspec(load_vllm).args
|
||||
load_vllm_kwargs = dict(
|
||||
model_name = model_name,
|
||||
config = model_config,
|
||||
gpu_memory_utilization = gpu_memory_utilization,
|
||||
max_seq_length = max_seq_length,
|
||||
dtype = dtype,
|
||||
float8_kv_cache = float8_kv_cache,
|
||||
enable_lora = vllm_enable_lora,
|
||||
max_lora_rank = max_lora_rank,
|
||||
disable_log_stats = disable_log_stats,
|
||||
use_bitsandbytes = load_in_4bit,
|
||||
unsloth_vllm_standby = unsloth_vllm_standby,
|
||||
is_vision_model = is_vlm_config,
|
||||
fp8_mode = fp8_mode,
|
||||
)
|
||||
for allowed_arg in allowed_args:
|
||||
if allowed_arg not in load_vllm_kwargs and allowed_arg in kwargs:
|
||||
load_vllm_kwargs[allowed_arg] = kwargs[allowed_arg]
|
||||
# Load vLLM first
|
||||
llm = load_vllm(**load_vllm_kwargs)
|
||||
|
||||
# Load vLLM first
|
||||
llm = load_vllm(**load_vllm_kwargs)
|
||||
# Convert to HF format
|
||||
_, quant_state_dict = get_vllm_state_dict(
|
||||
llm,
|
||||
config = model_config,
|
||||
is_vision_model = is_vlm_config,
|
||||
load_in_fp8 = load_in_fp8,
|
||||
)
|
||||
model = convert_vllm_to_huggingface(
|
||||
quant_state_dict,
|
||||
model_config,
|
||||
dtype,
|
||||
bnb_config,
|
||||
is_vision_model = is_vlm_config,
|
||||
)
|
||||
model.vllm_engine = llm
|
||||
llm.shared_weights = True
|
||||
model.fast_generate = model.vllm_engine.generate
|
||||
model.fast_generate_batches = functools.partial(generate_batches, model.vllm_engine)
|
||||
|
||||
# Convert to HF format
|
||||
_, quant_state_dict = get_vllm_state_dict(
|
||||
llm,
|
||||
config = model_config,
|
||||
is_vision_model = is_vlm_config,
|
||||
load_in_fp8 = load_in_fp8,
|
||||
)
|
||||
model = convert_vllm_to_huggingface(
|
||||
quant_state_dict,
|
||||
model_config,
|
||||
dtype,
|
||||
bnb_config,
|
||||
is_vision_model = is_vlm_config,
|
||||
)
|
||||
model.vllm_engine = llm
|
||||
llm.shared_weights = True
|
||||
model.fast_generate = model.vllm_engine.generate
|
||||
model.fast_generate_batches = functools.partial(generate_batches, model.vllm_engine)
|
||||
|
||||
raise_handler.remove()
|
||||
|
||||
# Return old flag
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = old_hf_transfer
|
||||
finally:
|
||||
raise_handler.remove()
|
||||
# Return old flag
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = old_hf_transfer
|
||||
|
||||
# Check float32 norm weights
|
||||
if os.environ.get("UNSLOTH_HIGH_PRECISION_LAYERNORM", "0") == "1":
|
||||
|
|
@ -1171,70 +1217,101 @@ class FastBaseModel:
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
_processor_load_error = None
|
||||
if (whisper_language and whisper_task) or auto_model.__name__.endswith(
|
||||
"ForConditionalGeneration"
|
||||
):
|
||||
try:
|
||||
tokenizer = auto_processor.from_pretrained(
|
||||
tokenizer_name,
|
||||
padding_side = "left",
|
||||
token = token,
|
||||
language = whisper_language,
|
||||
task = whisper_task,
|
||||
trust_remote_code = trust_remote_code,
|
||||
)
|
||||
except Exception as e:
|
||||
_processor_load_error = e
|
||||
tokenizer = None
|
||||
else:
|
||||
try:
|
||||
tokenizer = auto_processor.from_pretrained(
|
||||
tokenizer_name,
|
||||
padding_side = "left",
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
)
|
||||
except Exception as e:
|
||||
_processor_load_error = e
|
||||
tokenizer = get_auto_processor(
|
||||
tokenizer_name,
|
||||
padding_side = "left",
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
)
|
||||
|
||||
# If processor loading failed (e.g., tokenizer class not found),
|
||||
# or if AutoProcessor silently degraded to a text-only tokenizer
|
||||
# instead of returning a full VLM processor (issue #4085),
|
||||
# try constructing the processor manually from separate components.
|
||||
_processor_is_degraded = (
|
||||
is_vlm and tokenizer is not None and not hasattr(tokenizer, "image_processor")
|
||||
)
|
||||
if (tokenizer is None or _processor_is_degraded) and is_vlm:
|
||||
_fallback = _construct_vlm_processor_fallback(
|
||||
tokenizer_name,
|
||||
model_type_arch,
|
||||
token,
|
||||
trust_remote_code,
|
||||
)
|
||||
if _fallback is not None:
|
||||
tokenizer = _fallback
|
||||
# Missing torchvision silently degrades the VLM processor to a text-only
|
||||
# tokenizer; surface the real cause instead of the later collator error (#4202).
|
||||
if tokenizer is None or not hasattr(tokenizer, "image_processor"):
|
||||
if _missing_torchvision_error(_processor_load_error):
|
||||
raise ImportError(
|
||||
f"Unsloth: Could not load the vision processor for `{tokenizer_name}` "
|
||||
"because torchvision is not installed. transformers requires torchvision "
|
||||
"for this model's vision (image/video) processors. Please install it, "
|
||||
"e.g. `pip install torchvision`."
|
||||
# Functional load chain (AutoProcessor -> get_auto_processor -> manual VLM
|
||||
# fallback); offline is already forced upstream. Surfaces the error for the retry.
|
||||
def _acquire_processor(lfo):
|
||||
_err = None # underlying load failure (used by the entry-point retry)
|
||||
if (whisper_language and whisper_task) or auto_model.__name__.endswith(
|
||||
"ForConditionalGeneration"
|
||||
):
|
||||
try:
|
||||
_tok = auto_processor.from_pretrained(
|
||||
tokenizer_name,
|
||||
padding_side = "left",
|
||||
token = token,
|
||||
language = whisper_language,
|
||||
task = whisper_task,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = lfo,
|
||||
)
|
||||
import sys
|
||||
print(
|
||||
f"Unsloth: Warning - VLM processor fallback returned None for model_type={model_type_arch}",
|
||||
file = sys.stderr,
|
||||
except Exception as _e:
|
||||
_tok = None
|
||||
_err = _e
|
||||
else:
|
||||
try:
|
||||
_tok = auto_processor.from_pretrained(
|
||||
tokenizer_name,
|
||||
padding_side = "left",
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = lfo,
|
||||
)
|
||||
except Exception as _e:
|
||||
_err = _e
|
||||
try:
|
||||
_tok = get_auto_processor(
|
||||
tokenizer_name,
|
||||
padding_side = "left",
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = lfo,
|
||||
)
|
||||
except Exception:
|
||||
# Swallow so the manual fallback / entry-point retry can run.
|
||||
_tok = None
|
||||
|
||||
# Build the processor manually if it failed to load or silently degraded to
|
||||
# a text-only tokenizer (no image_processor) for a VLM (issue #4085).
|
||||
_processor_is_degraded = (
|
||||
is_vlm and _tok is not None and not hasattr(_tok, "image_processor")
|
||||
)
|
||||
if (_tok is None or _processor_is_degraded) and is_vlm:
|
||||
try:
|
||||
_fallback, _fb_err = _construct_vlm_processor_fallback(
|
||||
tokenizer_name,
|
||||
model_type_arch,
|
||||
token,
|
||||
trust_remote_code,
|
||||
local_files_only = lfo,
|
||||
)
|
||||
except Exception as _fe:
|
||||
_fallback, _fb_err = None, _fe
|
||||
if _fallback is not None:
|
||||
_tok = _fallback
|
||||
elif _err is None or (_fb_err is not None and _is_offline_related_error(_fb_err)):
|
||||
# Prefer a network fallback error over a permanent primary one so the
|
||||
# offline retry still fires.
|
||||
_err = _fb_err
|
||||
return _tok, _err
|
||||
|
||||
def _is_degraded_vlm(_t):
|
||||
# VLM that loaded only a text-only tokenizer (no image_processor).
|
||||
return is_vlm and _t is not None and not hasattr(_t, "image_processor")
|
||||
|
||||
tokenizer, _primary_err = _acquire_processor(local_files_only)
|
||||
# Online network failure/degrade: raise so @_offline_aware_load retries from cache.
|
||||
# Permanent / missing-file errors propagate; when already offline keep what we got.
|
||||
if (
|
||||
(tokenizer is None or _is_degraded_vlm(tokenizer))
|
||||
and not local_files_only
|
||||
and _is_offline_related_error(_primary_err)
|
||||
):
|
||||
raise _primary_err
|
||||
# Missing torchvision silently degrades a VLM processor to text-only; surface the
|
||||
# real cause instead of a later collator error (#4202), incl. on a silent degrade.
|
||||
if is_vlm and (tokenizer is None or not hasattr(tokenizer, "image_processor")):
|
||||
if _missing_torchvision_error(_primary_err):
|
||||
raise ImportError(
|
||||
f"Unsloth: Could not load the vision processor for `{tokenizer_name}` "
|
||||
"because torchvision is not installed. transformers requires torchvision "
|
||||
"for this model's vision (image/video) processors. Please install it, "
|
||||
"e.g. `pip install torchvision`."
|
||||
)
|
||||
import sys
|
||||
print(
|
||||
f"Unsloth: Warning - VLM processor fallback returned None for model_type={model_type_arch}",
|
||||
file = sys.stderr,
|
||||
)
|
||||
# Backwards compat: if processor has no chat_template (e.g. old saves without
|
||||
# chat_template.jinja) but the inner tokenizer does, copy it to the processor.
|
||||
if (
|
||||
|
|
@ -1271,8 +1348,7 @@ class FastBaseModel:
|
|||
try:
|
||||
model, tokenizer = patch_tokenizer(model, tokenizer)
|
||||
except Exception as _patch_err:
|
||||
# Some VLM processors (e.g., ERNIE VL) may fail during tokenizer patching.
|
||||
# Try loading tokenizer separately via AutoTokenizer as fallback.
|
||||
# Some VLM processors (e.g. ERNIE VL) fail patching; fall back to AutoTokenizer.
|
||||
try:
|
||||
from transformers import AutoTokenizer as _AutoTokenizer
|
||||
|
||||
|
|
@ -1281,6 +1357,7 @@ class FastBaseModel:
|
|||
padding_side = "left",
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = local_files_only,
|
||||
)
|
||||
model, _fallback_tok = patch_tokenizer(model, _fallback_tok)
|
||||
# Re-attach as processor wrapper if original was a processor
|
||||
|
|
@ -1288,8 +1365,10 @@ class FastBaseModel:
|
|||
tokenizer.tokenizer = _fallback_tok
|
||||
else:
|
||||
tokenizer = _fallback_tok
|
||||
except Exception:
|
||||
# If fallback also fails, raise the original error
|
||||
except Exception as _fb_err:
|
||||
# Online network failure: propagate for the offline retry; else raise the patch error.
|
||||
if not local_files_only and _is_offline_related_error(_fb_err):
|
||||
raise
|
||||
raise _patch_err
|
||||
model = post_patch_loss_function(model)
|
||||
|
||||
|
|
@ -1298,29 +1377,44 @@ class FastBaseModel:
|
|||
model.config.update({"unsloth_version": __version__})
|
||||
patch_saving_functions(model, vision = True)
|
||||
if tokenizer is None:
|
||||
# Last resort: try loading tokenizer via AutoTokenizer, then PreTrainedTokenizerFast
|
||||
try:
|
||||
# Last resort: AutoTokenizer, then PreTrainedTokenizerFast (raise on network failure to retry).
|
||||
def _last_resort_tokenizer(lfo):
|
||||
from transformers import AutoTokenizer as _AutoTokenizer
|
||||
tokenizer = _AutoTokenizer.from_pretrained(
|
||||
tokenizer_name,
|
||||
padding_side = "left",
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
)
|
||||
except Exception:
|
||||
try:
|
||||
from transformers import PreTrainedTokenizerFast
|
||||
tokenizer = PreTrainedTokenizerFast.from_pretrained(
|
||||
return _AutoTokenizer.from_pretrained(
|
||||
tokenizer_name,
|
||||
padding_side = "left",
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = lfo,
|
||||
)
|
||||
except Exception:
|
||||
del model
|
||||
raise RuntimeError(
|
||||
"Unsloth: The tokenizer is weirdly not loaded? Please check if there is one."
|
||||
from transformers import PreTrainedTokenizerFast
|
||||
return PreTrainedTokenizerFast.from_pretrained(
|
||||
tokenizer_name,
|
||||
padding_side = "left",
|
||||
token = token,
|
||||
trust_remote_code = trust_remote_code,
|
||||
local_files_only = lfo,
|
||||
)
|
||||
|
||||
_last_resort_err = None
|
||||
try:
|
||||
tokenizer = _last_resort_tokenizer(local_files_only)
|
||||
except Exception as _e:
|
||||
_last_resort_err = _e
|
||||
# Online network failure: let the entry point retry forced-offline.
|
||||
if not local_files_only and _is_offline_related_error(_e):
|
||||
raise
|
||||
if tokenizer is None:
|
||||
del model
|
||||
raise RuntimeError(
|
||||
"Unsloth: Could not load the tokenizer/processor. If you are "
|
||||
"offline, make sure the tokenizer files exist in the checkpoint "
|
||||
"folder or were previously downloaded to the Hugging Face cache, "
|
||||
"or set HF_HUB_OFFLINE=1 to force local loading. "
|
||||
"Otherwise please check that the model has a tokenizer."
|
||||
) from _last_resort_err
|
||||
patch_saving_functions(tokenizer, vision = True)
|
||||
|
||||
# Fix gradient accumulation. See issue #4982.
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue