Keep import unsloth working when bitsandbytes is absent (#7502)

* Keep `import unsloth` working when bitsandbytes is absent

device_type.py already prints "bitsandbytes is not installed - 4bit QLoRA
unallowed, but 16bit and full finetuning works" and clears
ALLOW_BITSANDBYTES / ALLOW_PREQUANTIZED_MODELS, but the import chain then
hard-required the module anyway, so `import unsloth` raised instead.

#7354 made this reachable: the gfx906 install path uninstalls the generic
bitsandbytes wheel (no gfx906 kernels in it), which leaves an MI50 / Radeon VII
host unable to import unsloth at all, not on the 16bit path the message
promises.

- kernels/utils.py: guard the bnb import; bind get_ptr and the five 4bit ctypes
  handles to a stub that raises a clear message if a 4bit path is entered.
  HAS_CUDA_STREAM stays False, which is the correct route.
- save.py, models/granite.py: guard Bnb_Linear4bit and peft's Linear4bit
  (peft exports it only when bnb imported cleanly) with placeholder classes.
  Both names only feed isinstance checks, so nothing matching is exact.
- _gpu_init.py: same degradation on the xpu branch as the cuda branch above.

Verified on a Strix Halo (gfx1151, DEVICE_TYPE=hip, torch 2.11.0+rocm7.13.0)
by blocking bitsandbytes with sys.modules["bitsandbytes"] = None, so
find_spec returns None and the import raises exactly as when the package is
absent. Before: ModuleNotFoundError at kernels/utils.py:136. After: import
succeeds, FastLanguageModel/FastModel import, ALLOW_BITSANDBYTES=False,
ALLOW_PREQUANTIZED=False, and the 4bit stub raises with the real cause. With
bitsandbytes present, every binding is unchanged.

New test walks the `import unsloth` module graph with ast and fails on any
unguarded bitsandbytes (or peft Linear4bit) import; verified it catches the
old code. Targeted suites: 702 passed, 18 skipped.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Address the review on #7502: zoo coupling, non-hip flags, py3.9 collection

Three findings, each reproduced first and negative-controlled after.

1. The fix still needed an unreleased unsloth_zoo (P1). save.py imported
   unsloth_zoo.saving_utils at module scope, and any zoo without the companion
   #953 fix imports bitsandbytes there, so `import unsloth` kept failing for a
   dependency set pyproject.toml allows. Raising the floor was not an option:
   PyPI's newest zoo is 2026.7.6 and #953 is merged but unreleased, so a bump
   would break every install today. Both names it pulled in are used only inside
   functions, so the import is now lazy at those two call sites, matching what
   determine_base_model_source in the same file already does. Verified against a
   real pre-#953 zoo checkout with bitsandbytes blocked: import succeeds, and
   restoring the eager import reproduces the failure at saving_utils.py:70.
   This PR no longer depends on a zoo release.

2. Capability flags were only cleared on hip (P2). device_type.py probed
   bitsandbytes inside its DEVICE_TYPE == "hip" branch, so a cuda or xpu host
   without bnb imported fine but still reported ALLOW_BITSANDBYTES=True, and the
   default load_in_4bit=True path in models/loader.py would select a 4bit
   checkpoint before failing. Clear both flags whenever the module is absent, on
   every backend, via find_spec so a working install pays nothing. A cuda host
   with bnb blocked now reports False/False; with bnb present nothing changes.

3. The new test could not be collected on Python 3.9 (P2). `Path | None` is a
   PEP 604 union and requires-python still allows 3.9, so pytest raised
   TypeError at import. Added `from __future__ import annotations`. Checked in
   real uv venvs on 3.9, 3.10 and 3.13: 2 passed each; removing the future
   import reproduces "unsupported operand type(s) for |" on 3.9 only.

The xpu branch in _gpu_init.py needs no separate flag handling now that the
probe is backend-independent.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Address the second review on #7502: guarded probe, and 8bit in the same guard

1. The capability probe used find_spec while the fallbacks in kernels/utils.py
   and _gpu_init.py treat any import failure as unavailable, so an installed but
   unusable wheel would leave ALLOW_BITSANDBYTES true while the kernels had
   already bound the stub. Probe with the same guarded import instead, so all
   three agree by construction. No new cost on any path: _gpu_init.py already
   imports bnb before device_type is reached on cuda, and device_type's own hip
   block imports it a few lines later.

   Worth recording that the state this prevents is currently unreachable for an
   unrelated reason: a broken wheel takes `import unsloth` down earlier, in
   transformers/integrations/bitsandbytes.py:20 via
   unsloth_zoo/patching_utils.py:680, whichever exception it raises (OSError also
   escapes the zoo moe_utils `except ImportError`). So this is correctness for
   when those imports get guarded, not an observable fix today.

2. Both loader guards printed for load_in_4bit or load_in_8bit but only cleared
   load_in_4bit, so an explicit load_in_8bit=True survived and reached
   Transformers, which builds the bnb quantizer and fails there. Clear both. The
   message no longer says AMD either: the flag now goes false whenever bnb is
   unusable on any backend.

Tests: the probe must not use find_spec, and an ast walk requires every
ALLOW_BITSANDBYTES guard in loader.py to clear both flags, so a third guard
cannot be added with the same omission. Dropping either fix reddens them (1 and
2 failures respectively). 4 passed on 3.9, 3.13 and the ROCm venv; absent and
healthy bnb both stay consistent across hip and cuda.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Drop the importlib import left over from the find_spec probe on #7502

* Address the third review on #7502: exact-name bypass and a forwarded bnb config

Both findings hold up, so both are fixed.

1. use_exact_model_name=True skipped the guard entirely. load_in_4bit defaults
   to True, so on a host without bitsandbytes
   FastLanguageModel.from_pretrained(name, use_exact_model_name=True) kept 4bit
   set and failed downstream. That option suppresses repo-name remapping and
   cannot make bitsandbytes available, so it has no business gating a capability
   check. Ungated at both sites.

2. A user-supplied quantization_config survived the fallback. It sets
   load_in_4bit/8bit at the top of from_pretrained and stays in kwargs, so
   clearing the local flags still let Transformers rebuild the bnb quantizer.
   Now dropped as part of the fallback.

One correction to the second suggestion: it cannot be dropped whenever the
fallback runs. quantization_config also carries GPTQ, AWQ, fp8 and torchao
configs, which have nothing to do with bitsandbytes and must reach the loader
untouched. The pop is gated on the config actually requesting load_in_4bit or
load_in_8bit, reusing the same dict/attr probe from the top of the function.

Behaviour, exercising the real guard block against synthetic inputs with
use_exact_model_name=True and bnb unusable:

  default 4bit, no cfg          4bit=False 8bit=False
  explicit 8bit, no cfg         4bit=False 8bit=False
  BitsAndBytesConfig(4bit/8bit) 4bit=False 8bit=False  config dropped
  dict bnb config               4bit=False 8bit=False  config dropped
  GPTQ config                   4bit=False 8bit=False  config SURVIVES
  fp8 dict                      4bit=False 8bit=False  config SURVIVES

Nothing changes when bitsandbytes works: the whole block is inside
`if not ALLOW_BITSANDBYTES`.

Tests: an ast walk requires neither guard to reference use_exact_model_name in
its test, and requires each to pop quantization_config behind a _wants_bnb
check, so an unconditional pop fails too. Re-gating one guard or removing one
pop reddens a test each. 6 passed on 3.9, 3.13 and the ROCm venv.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Address the fourth review on #7502: FastModel never reached the 16bit path

Both findings are real, and the second one meant this PR did not actually
deliver what it advertises for FastModel or vision loads. Reproduced first.

1. patch_compiling_bitsandbytes() ran unguarded at the top of
   FastModel.from_pretrained, and unsloth_zoo's copy imports bitsandbytes
   unconditionally (patching_utils.py:40). So every FastModel call on a
   bnb-less host died there, whatever the arguments:

     FastModel(load_in_16bit=True)    -> ModuleNotFoundError at patching_utils.py:40
     FastModel(full_finetuning=True)  -> ModuleNotFoundError at patching_utils.py:40

   The FastLanguageModel path already wraps this call in try/except with a
   warning, and its comment even says "Mirror FastModel" - FastModel was the
   unwrapped one. Wrapped it the same way, so behaviour is unchanged wherever
   bitsandbytes imports.

2. The mode-exclusivity check ran before the capability fallback. load_in_4bit
   defaults to True, so load_in_16bit=True made
   int(load_in_4bit) + int(load_in_16bit) == 2 and raised "Can only load in 4bit
   or 8bit or 16bit" before the fallback could clear the unavailable 4bit
   request. Moved the fallback ahead of that check.

After both, the same three calls get past every bitsandbytes gate and reach
model resolution, failing only on the deliberately fake repo name used by the
probe. Nothing changes when bitsandbytes works: the fallback is still inside
`if not ALLOW_BITSANDBYTES`, and the wrapper only swallows an import that
previously crashed the load.

Tests: the mode check must be preceded by an ALLOW_BITSANDBYTES fallback in the
same function, and no call to patch_compiling_bitsandbytes may sit outside a
try. The ordering assertion is scoped to the enclosing function on purpose - my
first version compared line numbers file-wide, so the other loader's guard
satisfied it and the negative control passed when it should have failed. With
the scoping fixed, moving the fallback back after the mode check reddens it, as
does unwrapping the patch call. 8 passed on 3.9, 3.13 and the ROCm venv.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
Daniel Han 2026-07-28 08:03:54 -07:00 committed by GitHub
commit 52a9601032
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 467 additions and 37 deletions

View file

@ -0,0 +1,296 @@
"""`import unsloth` must survive a missing bitsandbytes.
device_type.py already tells the user "bitsandbytes is not installed - 4bit QLoRA
unallowed, but 16bit and full finetuning works", and the gfx906 install path
(#7354) deliberately removes the generic wheel because it carries no gfx906
kernels. Any module-level `import bitsandbytes` on the import chain turns that
into an unimportable package instead.
peft's 4bit LoRA layer is exported only when bnb is importable, so
`from peft.tuners.lora import Linear4bit` fails on the same hosts and is checked
here too.
"""
# Path | None below is a PEP 604 union; the project still supports Python 3.9.
from __future__ import annotations
import ast
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
ROOT_MODULE = "unsloth"
def _module_path(name: str) -> Path | None:
base = REPO_ROOT / Path(*name.split("."))
for candidate in (base.with_suffix(".py"), base / "__init__.py"):
if candidate.is_file():
return candidate
return None
def _bnb_dependent(node: ast.stmt) -> bool:
"""True for an import that raises when bitsandbytes is absent."""
if isinstance(node, ast.Import):
return any(a.name.split(".")[0] == "bitsandbytes" for a in node.names)
if isinstance(node, ast.ImportFrom) and node.level == 0:
module = node.module or ""
if module.split(".")[0] == "bitsandbytes":
return True
# peft re-exports Linear4bit only when bnb imported cleanly.
if module.startswith("peft.tuners.lora"):
return any(a.name == "Linear4bit" for a in node.names)
return False
def _allow_bitsandbytes_gated(test: ast.expr) -> bool:
"""device_type.py sets ALLOW_BITSANDBYTES=False exactly when the import failed,
so a branch keyed on it cannot run without bnb."""
return any(isinstance(n, ast.Name) and n.id == "ALLOW_BITSANDBYTES" for n in ast.walk(test))
def _scan(path: Path, module: str):
"""Yield (lineno, source) for unguarded top-level imports.
Imports inside a `try`, or under an ALLOW_BITSANDBYTES branch, are guarded.
Other `if` bodies are not: the condition may well be true on a host without bnb.
"""
is_package = path.name == "__init__.py"
package = module if is_package else module.rpartition(".")[0]
tree = ast.parse(path.read_text(encoding = "utf-8"))
risky, edges = [], []
def walk(body, guarded):
for node in body:
if isinstance(node, (ast.Import, ast.ImportFrom)):
if not guarded and _bnb_dependent(node):
risky.append((node.lineno, ast.unparse(node)))
if isinstance(node, ast.Import):
edges.extend(a.name for a in node.names)
elif node.level:
parts = package.split(".")
base = ".".join(parts[: len(parts) - (node.level - 1)])
edges.append(f"{base}.{node.module}" if node.module else base)
else:
edges.append(node.module or "")
elif isinstance(node, ast.Try):
walk(node.body, True)
for handler in node.handlers:
walk(handler.body, True)
walk(node.orelse, True)
walk(node.finalbody, guarded)
elif isinstance(node, ast.If):
walk(node.body, guarded or _allow_bitsandbytes_gated(node.test))
walk(node.orelse, guarded)
walk(tree.body, False)
return risky, edges
def test_no_unguarded_bitsandbytes_import_on_the_unsloth_import_chain():
seen, pending, offenders = set(), [(ROOT_MODULE, [])], []
while pending:
module, chain = pending.pop()
if module in seen:
continue
seen.add(module)
path = _module_path(module)
if path is None:
continue
risky, edges = _scan(path, module)
for lineno, source in risky:
rel = path.relative_to(REPO_ROOT).as_posix()
offenders.append(f"{rel}:{lineno} {source}\n via {' -> '.join(chain + [module])}")
pending.extend(
(edge, chain + [module]) for edge in edges if edge.split(".")[0] == ROOT_MODULE
)
assert len(seen) > 20, f"import chain walk collapsed, only reached {seen}"
assert not offenders, (
"`import unsloth` must not hard-require bitsandbytes. Wrap these in "
"try/except and fall back to a placeholder:\n " + "\n ".join(offenders)
)
def test_missing_bnb_leaves_a_callable_that_reports_the_real_cause():
"""The 4bit ctypes handles degrade to a stub, not a NameError later on."""
src = (REPO_ROOT / "unsloth" / "kernels" / "utils.py").read_text(encoding = "utf-8")
assert "def _bnb_required(" in src
assert "get_ptr = _bnb_required" in src
for name in (
"cdequantize_blockwise_fp32",
"cdequantize_blockwise_fp16_nf4",
"cdequantize_blockwise_bf16_nf4",
"cgemm_4bit_inference_naive_fp16",
"cgemm_4bit_inference_naive_bf16",
):
assert f"{name} = _bnb_required" in src, f"{name} has no bnb-less fallback"
def test_capability_flags_come_from_a_guarded_import_not_find_spec():
"""kernels/utils.py and _gpu_init.py treat any import failure as unavailable.
device_type.py must agree, or an installed-but-unusable wheel leaves
ALLOW_BITSANDBYTES true while the kernels fall back to the stub."""
src = (REPO_ROOT / "unsloth" / "device_type.py").read_text(encoding = "utf-8")
head = src.split('if DEVICE_TYPE == "hip":')[0]
assert "import bitsandbytes as _bnb_probe" in head
assert 'find_spec("bitsandbytes")' not in head, "find_spec cannot see a broken wheel"
assert head.count("ALLOW_BITSANDBYTES = False") >= 1
def _bnb_guards():
src = (REPO_ROOT / "unsloth" / "models" / "loader.py").read_text(encoding = "utf-8")
tree = ast.parse(src)
return src, [
node
for node in ast.walk(tree)
if isinstance(node, ast.If)
and any(
isinstance(n, ast.Name) and n.id == "ALLOW_BITSANDBYTES" for n in ast.walk(node.test)
)
]
def test_bitsandbytes_guard_is_not_gated_on_use_exact_model_name():
"""use_exact_model_name suppresses repo-name remapping; it cannot make bnb
available. Gating on it left the default load_in_4bit=True set on a host
without bitsandbytes."""
_, guards = _bnb_guards()
assert len(guards) == 2, f"expected both loader guards, found {len(guards)}"
for guard in guards:
names = {n.id for n in ast.walk(guard.test) if isinstance(n, ast.Name)}
assert (
"use_exact_model_name" not in names
), f"guard at line {guard.lineno} still gates the capability check on naming"
def test_bitsandbytes_guard_drops_a_bnb_quantization_config():
"""A BitsAndBytesConfig in kwargs re-sets the flags downstream, so clearing
load_in_4bit/8bit alone still builds the bnb quantizer in Transformers. A
non-bnb config (GPTQ/AWQ/fp8) must not be touched."""
_, guards = _bnb_guards()
for guard in guards:
# ast.unparse normalises quotes, so match on the call shape instead.
def _is_pop(node):
return (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and node.func.attr == "pop"
and isinstance(node.func.value, ast.Name)
and node.func.value.id == "kwargs"
and node.args
and isinstance(node.args[0], ast.Constant)
and node.args[0].value == "quantization_config"
)
assert any(
_is_pop(n) for n in ast.walk(guard)
), f"guard at line {guard.lineno} leaves the bnb config in kwargs"
# the pop must be conditional on the config actually asking for bnb
pops = [
node
for node in ast.walk(guard)
if isinstance(node, ast.If) and any(_is_pop(n) for n in ast.walk(node))
]
assert pops, f"guard at line {guard.lineno} pops unconditionally"
assert any(
isinstance(n, ast.Name) and n.id == "_wants_bnb"
for node in pops
for n in ast.walk(node.test)
), f"guard at line {guard.lineno} does not gate the pop on a bnb request"
def test_bitsandbytes_guard_clears_8bit_as_well_as_4bit():
"""8bit is bitsandbytes too: leaving load_in_8bit set sends the request to
Transformers, which builds the bnb quantizer and fails there instead."""
src = (REPO_ROOT / "unsloth" / "models" / "loader.py").read_text(encoding = "utf-8")
tree = ast.parse(src)
guards = [
node
for node in ast.walk(tree)
if isinstance(node, ast.If)
and any(
isinstance(n, ast.Name) and n.id == "ALLOW_BITSANDBYTES" for n in ast.walk(node.test)
)
]
assert len(guards) == 2, f"expected both loader guards, found {len(guards)}"
for guard in guards:
cleared = {
target.id
for stmt in guard.body
if isinstance(stmt, ast.Assign)
for target in stmt.targets
if isinstance(target, ast.Name)
and isinstance(stmt.value, ast.Constant)
and stmt.value.value is False
}
assert {
"load_in_4bit",
"load_in_8bit",
} <= cleared, f"guard at line {guard.lineno} clears only {sorted(cleared)}"
def test_capability_fallback_precedes_the_mutually_exclusive_mode_check():
"""load_in_4bit defaults to True, so load_in_16bit=True trips the
"can only load in 4bit or 8bit or 16bit" RuntimeError unless the unavailable
4bit request is cleared first. That check must come after the fallback."""
src, _ = _bnb_guards()
tree = ast.parse(src)
checked = 0
# Scope to the enclosing function: the other loader's guard sits earlier in the
# file and would otherwise satisfy a plain line-number comparison.
for func in ast.walk(tree):
if not isinstance(func, ast.FunctionDef):
continue
raises = [
node.lineno
for node in ast.walk(func)
if isinstance(node, ast.Raise)
and "Can only load in 4bit or 8bit or 16bit" in ast.unparse(node)
]
if not raises:
continue
guards = [
node.lineno
for node in ast.walk(func)
if isinstance(node, ast.If)
and any(
isinstance(n, ast.Name) and n.id == "ALLOW_BITSANDBYTES"
for n in ast.walk(node.test)
)
]
for lineno in raises:
checked += 1
assert any(g < lineno for g in guards), (
f"{func.name}: the mode check at line {lineno} runs before this "
"function's ALLOW_BITSANDBYTES fallback, so load_in_16bit=True on a "
"bnb-less host raises instead of taking the 16bit path"
)
assert checked, "mode-exclusivity check not found"
def test_bitsandbytes_compile_patch_is_never_called_unguarded():
"""unsloth_zoo's patch_compiling_bitsandbytes imports bitsandbytes
unconditionally, so an unwrapped call raises on a bnb-less host before any
fallback can run."""
src = (REPO_ROOT / "unsloth" / "models" / "loader.py").read_text(encoding = "utf-8")
tree = ast.parse(src)
calls = [
node
for node in ast.walk(tree)
if isinstance(node, ast.Call)
and isinstance(node.func, ast.Name)
and node.func.id == "patch_compiling_bitsandbytes"
]
assert calls, "call sites not found"
guarded = {
call.lineno
for node in ast.walk(tree)
if isinstance(node, ast.Try)
for call in ast.walk(node)
if isinstance(call, ast.Call)
and isinstance(call.func, ast.Name)
and call.func.id == "patch_compiling_bitsandbytes"
}
unguarded = sorted({c.lineno for c in calls} - guarded)
assert not unguarded, f"patch_compiling_bitsandbytes called unguarded at {unguarded}"