From 30094835d5df16e18a26bfc67da09e8e24ae48b3 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 27 Jul 2026 07:23:22 +0000 Subject: [PATCH] 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. --- .../core/inference/diffusion_patch_backend.py | 55 +++++--- .../tests/test_diffusion_patch_backend.py | 122 ++++++++++++++++++ 2 files changed, 159 insertions(+), 18 deletions(-) create mode 100644 studio/backend/tests/test_diffusion_patch_backend.py diff --git a/studio/backend/core/inference/diffusion_patch_backend.py b/studio/backend/core/inference/diffusion_patch_backend.py index 29cf96e5d0..a144df9e52 100644 --- a/studio/backend/core/inference/diffusion_patch_backend.py +++ b/studio/backend/core/inference/diffusion_patch_backend.py @@ -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() diff --git a/studio/backend/tests/test_diffusion_patch_backend.py b/studio/backend/tests/test_diffusion_patch_backend.py new file mode 100644 index 0000000000..3f038b8ccc --- /dev/null +++ b/studio/backend/tests/test_diffusion_patch_backend.py @@ -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