Only retry the unsloth import where it can succeed

The gate still let the retry run on hosts unsloth does not support, which is
where it is most harmful: a 7 GB macOS runner lost the Studio server 26 s into a
load, and the Linux runner was torn down mid-generation. Neither MPS nor plain
CPU can complete the import, so the retry there pays the cost and fails anyway.

Require an accelerator unsloth actually supports (CUDA/ROCm via torch.cuda, or
XPU), with UNSLOTH_ALLOW_CPU as the documented override, and hoist the predicate
to module level so it is tested directly rather than through the import system.
On a CPU-only host the retry no longer fires at all; on CUDA the clean-environment
case it was added for still passes 29/29.
This commit is contained in:
Daniel Han 2026-07-27 07:23:22 +00:00
commit 30094835d5
2 changed files with 160 additions and 19 deletions

View file

@ -27,6 +27,33 @@ from typing import Any, Callable, Optional
_HELPERS: Optional[dict] = None
def _retry_could_help(exc: BaseException) -> bool:
"""Whether importing ``unsloth`` could turn ``exc`` into a successful retry.
See ``_helpers`` for why each condition is here; the short version is that the import is
expensive and must not run where it cannot succeed."""
import importlib.util
import os
import sys
if not isinstance(exc, ImportError) or "unsloth" in sys.modules:
return False
torch = sys.modules.get("torch")
if torch is None:
return False
if os.environ.get("UNSLOTH_ALLOW_CPU", "").strip().lower() not in ("1", "true", "yes", "on"):
try:
xpu = getattr(torch, "xpu", None)
if not (torch.cuda.is_available() or (xpu is not None and xpu.is_available())):
return False
except Exception: # noqa: BLE001 — an unprobeable device is not one unsloth can use
return False
try:
return importlib.util.find_spec("unsloth") is not None
except Exception: # noqa: BLE001 — an unimportable package cannot set the sentinel either
return False
def _helpers() -> Optional[dict]:
"""``{"patch": patch_function, "restore": restore_original}``, or None when unavailable.
@ -37,12 +64,17 @@ def _helpers() -> Optional[dict]:
False). So on failure, import ``unsloth`` and retry once, which is also the import order Unsloth
documents.
The retry is gated, because it is not free: importing ``unsloth`` pulls torch in behind it and
costs ~940 MB of RSS in a process that had neither, only to fail anyway on a host with no
accelerator, which is enough to matter on a small CI runner mid-generation. So it runs only when
The retry is gated, because it is not free: importing ``unsloth`` costs ~940 MB of RSS measured
in a process that had not already loaded torch, and on a host unsloth does not support it pays
that and fails anyway. Ungated it took two cross-platform CI runners down -- a Linux job that had
generated fine at ~900 s died 19 s in with SIGTERM and every ``if: always()`` step skipped (the
runner torn down, not a step failing), and a 7 GB macOS runner lost the server 26 s into a load.
So it runs only when it can actually succeed:
* ``torch`` is already imported -- true of the server and of anything that patches a real
module, and the condition that keeps the retry from being the thing that loads torch,
* ``torch`` is already imported -- true of the server and of anything patching a real module,
and the condition that stops the retry from being what loads torch,
* an accelerator ``unsloth`` supports is present (CUDA/ROCm or XPU; ``UNSLOTH_ALLOW_CPU``
overrides). MPS and plain CPU are not, and there the import raises after paying,
* ``unsloth`` is installed but not yet imported (if it were, the sentinel would be set and the
first attempt would have worked),
* and the first failure was the ImportError that guard raises."""
@ -54,19 +86,6 @@ def _helpers() -> Optional[dict]:
from unsloth_zoo.temporary_patches.utils import patch_function, restore_original
return {"patch": patch_function, "restore": restore_original}
def _retry_could_help(exc: BaseException) -> bool:
import importlib.util
import sys
if not isinstance(exc, ImportError) or "unsloth" in sys.modules:
return False
if "torch" not in sys.modules:
return False
try:
return importlib.util.find_spec("unsloth") is not None
except Exception: # noqa: BLE001 — an unimportable package cannot set the sentinel either
return False
for attempt in (0, 1):
try:
_HELPERS = _load()

View file

@ -0,0 +1,122 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Unit tests for the vetted patch entry point (``diffusion_patch_backend.py``).
Focused on the gate around the ``unsloth`` retry: it exists so a process that never imported
unsloth (the test suite, a worker) still installs patches instead of silently running unpatched,
but it must never fire where the import cannot succeed, because it is expensive enough there to
take a small CI runner down.
"""
from __future__ import annotations
import sys
import types
import pytest
import core.inference.diffusion_patch_backend as pb
_SENTINEL_ERROR = ImportError("Please install Unsloth via `pip install unsloth`!")
@pytest.fixture(autouse = True)
def _reset_memo(monkeypatch):
pb._HELPERS = None
monkeypatch.delenv("UNSLOTH_ALLOW_CPU", raising = False)
yield
pb._HELPERS = None
def _torch(*, cuda = False, xpu = False):
return types.SimpleNamespace(
cuda = types.SimpleNamespace(is_available = lambda: cuda),
xpu = types.SimpleNamespace(is_available = lambda: xpu),
)
def _modules(monkeypatch, *, torch = None, unsloth = False):
"""Stub sys.modules so the gate sees a chosen torch / unsloth state."""
mods = dict(sys.modules)
mods.pop("unsloth", None)
mods.pop("torch", None)
if torch is not None:
mods["torch"] = torch
if unsloth:
mods["unsloth"] = types.ModuleType("unsloth")
monkeypatch.setattr(sys, "modules", mods)
def test_retry_skipped_without_a_supported_accelerator(monkeypatch):
# A CPU-only or MPS host cannot import unsloth, so paying ~940 MB of RSS to find that out is
# pure cost. Ungated this took down a Linux CI runner mid-generation and a 7 GB macOS one
# during load.
_modules(monkeypatch, torch = _torch())
assert pb._retry_could_help(_SENTINEL_ERROR) is False
def test_retry_skipped_when_torch_is_not_loaded(monkeypatch):
# The retry must never be the thing that loads torch into a process that had avoided it.
_modules(monkeypatch, torch = None)
assert pb._retry_could_help(_SENTINEL_ERROR) is False
@pytest.mark.parametrize("device", ["cuda", "xpu"])
def test_retry_runs_on_an_accelerator_unsloth_supports(monkeypatch, device):
# The case the retry exists for: a GPU host whose process has simply not imported unsloth yet.
_modules(monkeypatch, torch = _torch(**{device: True}))
assert pb._retry_could_help(_SENTINEL_ERROR) is True
def test_retry_runs_on_cpu_when_explicitly_allowed(monkeypatch):
monkeypatch.setenv("UNSLOTH_ALLOW_CPU", "1")
_modules(monkeypatch, torch = _torch())
assert pb._retry_could_help(_SENTINEL_ERROR) is True
def test_retry_skipped_when_unsloth_is_already_imported(monkeypatch):
# Then the sentinel would already be set and the first attempt would have worked, so the
# failure is something else and re-importing cannot fix it.
_modules(monkeypatch, torch = _torch(cuda = True), unsloth = True)
assert pb._retry_could_help(_SENTINEL_ERROR) is False
def test_retry_skipped_for_a_non_import_failure(monkeypatch):
# A broken patch_function is not fixed by importing unsloth.
_modules(monkeypatch, torch = _torch(cuda = True))
assert pb._retry_could_help(RuntimeError("boom")) is False
def test_retry_skipped_when_the_device_probe_raises(monkeypatch):
# An unprobeable device is not one unsloth can use, so fail closed rather than pay the import.
broken = types.SimpleNamespace(
cuda = types.SimpleNamespace(is_available = lambda: (_ for _ in ()).throw(RuntimeError())),
xpu = None,
)
_modules(monkeypatch, torch = broken)
assert pb._retry_could_help(_SENTINEL_ERROR) is False
def test_helpers_memoises_the_unavailable_result(monkeypatch):
# Resolution can import unsloth, so it must be attempted at most once per process.
attempts: list[int] = []
def _boom():
attempts.append(1)
raise _SENTINEL_ERROR
monkeypatch.setattr(pb, "_retry_could_help", lambda exc: False)
monkeypatch.setitem(sys.modules, "unsloth_zoo.temporary_patches.utils", None)
_modules(monkeypatch, torch = _torch())
assert pb._helpers() is None
assert pb._helpers() is None
def test_apply_and_revert_are_no_ops_when_helpers_are_unavailable(monkeypatch):
# The contract the callers rely on: never raise, just report that nothing was patched.
monkeypatch.setattr(pb, "_helpers", lambda: None)
target = types.SimpleNamespace(fn = lambda: 1)
assert pb.apply_patch(target, "fn", lambda: 2) is False
assert pb.revert_patch(target, "fn") is False
assert target.fn() == 1