diff --git a/tests/test_windows_rocm_bnb_version.py b/tests/test_windows_rocm_bnb_version.py new file mode 100644 index 0000000000..c8439b1fcd --- /dev/null +++ b/tests/test_windows_rocm_bnb_version.py @@ -0,0 +1,251 @@ +# Unsloth - 2x faster, 60% less VRAM LLM training and finetuning +# Copyright 2023-present Daniel Han-Chen, Michael Han-Chen & the Unsloth team. All rights reserved. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Lesser General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU Lesser General Public License for more details. + +"""Tests for ``maybe_set_windows_rocm_bnb_version`` (unsloth/import_fixes.py). + +The module is loaded in isolation (stdlib + packaging only), so no torch / +GPU is required and unsloth's GPU init never runs.""" + +from __future__ import annotations + +import importlib.util +import os +import types +from pathlib import Path + +import pytest + +_IMPORT_FIXES_PATH = Path(__file__).resolve().parent.parent / "unsloth" / "import_fixes.py" + + +def _load_import_fixes(): + spec = importlib.util.spec_from_file_location( + "unsloth_import_fixes_under_test", _IMPORT_FIXES_PATH + ) + module = importlib.util.module_from_spec(spec) + assert spec is not None and spec.loader is not None + spec.loader.exec_module(module) + return module + + +@pytest.fixture() +def import_fixes(): + return _load_import_fixes() + + +@pytest.fixture() +def clean_env(monkeypatch): + """Unset the env vars and remove them afterwards (the function writes + os.environ directly, which monkeypatch does not auto-revert).""" + for var in ( + "BNB_ROCM_VERSION", + "UNSLOTH_SKIP_BNB_ROCM_VERSION", + "UNSLOTH_BNB_ROCM_VERSION_SOURCE", + ): + monkeypatch.delenv(var, raising = False) + yield monkeypatch + for var in ( + "BNB_ROCM_VERSION", + "UNSLOTH_SKIP_BNB_ROCM_VERSION", + "UNSLOTH_BNB_ROCM_VERSION_SOURCE", + ): + os.environ.pop(var, None) + + +def _force(import_fixes, monkeypatch, *, win, rocm, detected): + monkeypatch.setattr(import_fixes.sys, "platform", "win32" if win else "linux") + monkeypatch.setattr(import_fixes, "_is_hip_torch_build", lambda: rocm) + monkeypatch.setattr(import_fixes, "_detect_installed_bnb_rocm_version", lambda: detected) + + +# --------------------------------------------------------------------------- +# _detect_installed_bnb_rocm_version +# --------------------------------------------------------------------------- + + +def test_detect_picks_highest_rocm_suffix(import_fixes, tmp_path, monkeypatch): + pkg = tmp_path / "bitsandbytes" + pkg.mkdir() + for name in ( + "libbitsandbytes_rocm72.dll", + "libbitsandbytes_rocm713.dll", # numerically highest -> should win + "libbitsandbytes_cpu.dll", + "__init__.py", + ): + (pkg / name).write_text("") + fake_spec = types.SimpleNamespace(submodule_search_locations = [str(pkg)]) + monkeypatch.setattr(importlib.util, "find_spec", lambda name: fake_spec) + assert import_fixes._detect_installed_bnb_rocm_version() == "713" + + +def test_detect_none_when_only_non_rocm_dlls(import_fixes, tmp_path, monkeypatch): + pkg = tmp_path / "bitsandbytes" + pkg.mkdir() + (pkg / "libbitsandbytes_cpu.dll").write_text("") + (pkg / "libbitsandbytes_cuda124.dll").write_text("") + fake_spec = types.SimpleNamespace(submodule_search_locations = [str(pkg)]) + monkeypatch.setattr(importlib.util, "find_spec", lambda name: fake_spec) + assert import_fixes._detect_installed_bnb_rocm_version() is None + + +def test_detect_none_when_bnb_absent(import_fixes, monkeypatch): + monkeypatch.setattr(importlib.util, "find_spec", lambda name: None) + assert import_fixes._detect_installed_bnb_rocm_version() is None + + +# --------------------------------------------------------------------------- +# maybe_set_windows_rocm_bnb_version +# --------------------------------------------------------------------------- + + +def test_sets_bnb_version_on_windows_rocm(import_fixes, clean_env): + _force(import_fixes, clean_env, win = True, rocm = True, detected = "72") + assert import_fixes.maybe_set_windows_rocm_bnb_version() == "72" + assert os.environ["BNB_ROCM_VERSION"] == "72" + assert os.environ["UNSLOTH_BNB_ROCM_VERSION_SOURCE"] == "detected" + + +def test_noop_off_windows(import_fixes, clean_env): + # Linux ROCm resolves its backend correctly from torch.version.hip. + _force(import_fixes, clean_env, win = False, rocm = True, detected = "72") + assert import_fixes.maybe_set_windows_rocm_bnb_version() is None + assert "BNB_ROCM_VERSION" not in os.environ + + +def test_noop_when_not_rocm_torch(import_fixes, clean_env): + _force(import_fixes, clean_env, win = True, rocm = False, detected = "72") + assert import_fixes.maybe_set_windows_rocm_bnb_version() is None + assert "BNB_ROCM_VERSION" not in os.environ + + +def test_noop_when_no_rocm_dll_installed(import_fixes, clean_env): + # Never force a ROCm backend name when no ROCm DLL ships (avoid breaking a + # non-ROCm bitsandbytes that happens to sit next to a ROCm torch build). + _force(import_fixes, clean_env, win = True, rocm = True, detected = None) + assert import_fixes.maybe_set_windows_rocm_bnb_version() is None + assert "BNB_ROCM_VERSION" not in os.environ + + +def test_respects_user_provided_value(import_fixes, clean_env): + clean_env.setenv("BNB_ROCM_VERSION", "999") + _force(import_fixes, clean_env, win = True, rocm = True, detected = "72") + assert import_fixes.maybe_set_windows_rocm_bnb_version() is None + assert os.environ["BNB_ROCM_VERSION"] == "999" + + +def test_explicit_opt_out(import_fixes, clean_env): + clean_env.setenv("UNSLOTH_SKIP_BNB_ROCM_VERSION", "1") + _force(import_fixes, clean_env, win = True, rocm = True, detected = "72") + assert import_fixes.maybe_set_windows_rocm_bnb_version() is None + assert "BNB_ROCM_VERSION" not in os.environ + + +def test_redetects_sitecustomize_seeded_default(import_fixes, clean_env): + # Studio's installer persists a default via the venv sitecustomize.py; the + # wheel may have changed since, so the seeded value must be redetected. + clean_env.setenv("BNB_ROCM_VERSION", "72") + clean_env.setenv("UNSLOTH_BNB_ROCM_VERSION_SOURCE", "sitecustomize") + _force(import_fixes, clean_env, win = True, rocm = True, detected = "713") + assert import_fixes.maybe_set_windows_rocm_bnb_version() == "713" + assert os.environ["BNB_ROCM_VERSION"] == "713" + assert os.environ["UNSLOTH_BNB_ROCM_VERSION_SOURCE"] == "detected" + + +def test_sitecustomize_default_kept_when_no_dll_found(import_fixes, clean_env): + # A failed redetect must not discard the seeded value. + clean_env.setenv("BNB_ROCM_VERSION", "72") + clean_env.setenv("UNSLOTH_BNB_ROCM_VERSION_SOURCE", "sitecustomize") + _force(import_fixes, clean_env, win = True, rocm = True, detected = None) + assert import_fixes.maybe_set_windows_rocm_bnb_version() is None + assert os.environ["BNB_ROCM_VERSION"] == "72" + assert os.environ["UNSLOTH_BNB_ROCM_VERSION_SOURCE"] == "sitecustomize" + + +def test_user_value_with_non_sitecustomize_marker_untouched(import_fixes, clean_env): + # Only the sitecustomize marker makes a value redetectable. + clean_env.setenv("BNB_ROCM_VERSION", "999") + clean_env.setenv("UNSLOTH_BNB_ROCM_VERSION_SOURCE", "detected") + _force(import_fixes, clean_env, win = True, rocm = True, detected = "72") + assert import_fixes.maybe_set_windows_rocm_bnb_version() is None + assert os.environ["BNB_ROCM_VERSION"] == "999" + + +def test_opt_out_unseats_sitecustomize_seeded_value(import_fixes, clean_env): + # The opt-out must also drop a default our own sitecustomize block seeded, + # so bitsandbytes never sees the override the user disabled. + clean_env.setenv("BNB_ROCM_VERSION", "72") + clean_env.setenv("UNSLOTH_BNB_ROCM_VERSION_SOURCE", "sitecustomize") + clean_env.setenv("UNSLOTH_SKIP_BNB_ROCM_VERSION", "1") + _force(import_fixes, clean_env, win = True, rocm = True, detected = "713") + assert import_fixes.maybe_set_windows_rocm_bnb_version() is None + assert "BNB_ROCM_VERSION" not in os.environ + assert "UNSLOTH_BNB_ROCM_VERSION_SOURCE" not in os.environ + + +def test_opt_out_keeps_explicit_user_value(import_fixes, clean_env): + # Opt-out must never remove a value the user set themselves (no marker). + clean_env.setenv("BNB_ROCM_VERSION", "999") + clean_env.setenv("UNSLOTH_SKIP_BNB_ROCM_VERSION", "1") + _force(import_fixes, clean_env, win = True, rocm = True, detected = "72") + assert import_fixes.maybe_set_windows_rocm_bnb_version() is None + assert os.environ["BNB_ROCM_VERSION"] == "999" + + +def test_empty_string_value_without_marker_is_respected(import_fixes, clean_env): + # "" counts as present: without the sitecustomize marker it is not ours + # to overwrite. + clean_env.setenv("BNB_ROCM_VERSION", "") + _force(import_fixes, clean_env, win = True, rocm = True, detected = "72") + assert import_fixes.maybe_set_windows_rocm_bnb_version() is None + assert os.environ["BNB_ROCM_VERSION"] == "" + + +# --------------------------------------------------------------------------- +# _is_hip_torch_build (the strict gate -- regression for the HIP-SDK-on-a- +# CUDA-box false positive: env hints like HIP_PATH must NOT count) +# --------------------------------------------------------------------------- + + +def _fake_torch(hip): + return types.SimpleNamespace(version = types.SimpleNamespace(hip = hip)) + + +def test_hip_build_true_from_wheel_tag(import_fixes, monkeypatch): + monkeypatch.setattr(import_fixes, "importlib_version", lambda name: "2.11.0+rocm7.13.0") + assert import_fixes._is_hip_torch_build() is True + + +def test_hip_build_true_from_torch_version_hip(import_fixes, monkeypatch): + # Custom/source HIP build without the +rocm tag. + monkeypatch.setattr(import_fixes, "importlib_version", lambda name: "2.11.0") + monkeypatch.setitem(__import__("sys").modules, "torch", _fake_torch("7.2.0")) + assert import_fixes._is_hip_torch_build() is True + + +def test_hip_build_false_for_cuda_torch_despite_rocm_env_hints(import_fixes, monkeypatch): + """HIP SDK env vars set but CUDA torch: the strict gate must say False, + otherwise BNB_ROCM_VERSION gets set and CUDA bitsandbytes raises.""" + monkeypatch.setenv("HIP_PATH", r"C:\Program Files\AMD\ROCm\6.2") + monkeypatch.setenv("ROCM_PATH", r"C:\Program Files\AMD\ROCm\6.2") + monkeypatch.setattr(import_fixes, "importlib_version", lambda name: "2.9.0+cu126") + monkeypatch.setitem(__import__("sys").modules, "torch", _fake_torch(None)) + assert import_fixes._is_hip_torch_build() is False + + +def test_hip_build_false_when_torch_absent(import_fixes, monkeypatch): + def _raise(name): + raise Exception("no torch dist") + + monkeypatch.setattr(import_fixes, "importlib_version", _raise) + monkeypatch.setitem(__import__("sys").modules, "torch", None) + assert import_fixes._is_hip_torch_build() is False diff --git a/unsloth/_gpu_init.py b/unsloth/_gpu_init.py index a7b8ac594c..ede97ce68e 100644 --- a/unsloth/_gpu_init.py +++ b/unsloth/_gpu_init.py @@ -96,6 +96,13 @@ if already_imported: ) del already_imported, critical_modules +# Pin BNB_ROCM_VERSION before bitsandbytes is first imported (`import +# unsloth_zoo` below pulls it in on ROCm hosts). +from .import_fixes import maybe_set_windows_rocm_bnb_version + +maybe_set_windows_rocm_bnb_version() +del maybe_set_windows_rocm_bnb_version + # Multi-GPU is not yet supported (beta available on request). # Fixes https://github.com/unslothai/unsloth/issues/1266 diff --git a/unsloth/import_fixes.py b/unsloth/import_fixes.py index 25f5fdb9d2..0622f73965 100644 --- a/unsloth/import_fixes.py +++ b/unsloth/import_fixes.py @@ -2162,3 +2162,91 @@ def disable_broken_causal_conv1d(): "Unsloth: Detected broken causal_conv1d binary; " "disabling causal_conv1d fast path and continuing import." ) + + +_BNB_ROCM_DLL_RE = re.compile(r"libbitsandbytes_rocm(\d+)\.dll", re.IGNORECASE) + + +def _is_hip_torch_build(): + """True only when torch itself is a HIP/ROCm build. Env hints (HIP_PATH + etc.) do not count: CUDA bitsandbytes raises at import when the ROCm + override is set. Wheel tag first (no torch import); torch.version.hip + fallback for source builds.""" + try: + if "rocm" in str(importlib_version("torch")).lower(): + return True + except Exception: + pass + try: + import torch + return bool(getattr(torch.version, "hip", None)) + except Exception: + return False + + +def _detect_installed_bnb_rocm_version(): + """Highest installed ``libbitsandbytes_rocm.dll`` suffix ("72", "713") + or ``None``. Listing order is unordered, so take the numeric max.""" + try: + spec = importlib.util.find_spec("bitsandbytes") + except Exception: + return None + if spec is None or not spec.submodule_search_locations: + return None + + suffixes = [] + for pkg_dir in spec.submodule_search_locations: + try: + entries = os.listdir(pkg_dir) + except Exception: + continue + for entry in entries: + match = _BNB_ROCM_DLL_RE.fullmatch(entry) + if match is not None: + suffixes.append(match.group(1)) + if not suffixes: + return None + return max(suffixes, key = lambda value: int(value)) + + +def maybe_set_windows_rocm_bnb_version(): + """Pin ``BNB_ROCM_VERSION`` from the installed wheel on Windows + ROCm torch. + + AMD's Windows wheel ships one ``libbitsandbytes_rocm.dll`` whose + suffix can disagree with ``torch.version.hip`` (HIP 7.13 vs rocm72.dll), + breaking the native 4-bit/8-bit paths. Pin the installed suffix before + bitsandbytes is first imported. + + No-op unless ALL of: Windows, a real HIP torch build (env hints like + HIP_PATH do not count), a ROCm DLL installed, and no explicit user value. + Linux is untouched. Values seeded by Studio's venv sitecustomize.py + (marked ``UNSLOTH_BNB_ROCM_VERSION_SOURCE=sitecustomize``) are + redetectable defaults, not overrides; ``UNSLOTH_SKIP_BNB_ROCM_VERSION=1`` + opts out and drops a seeded default. Returns the value set, else None. + """ + if sys.platform != "win32": + return None + if os.environ.get("UNSLOTH_SKIP_BNB_ROCM_VERSION") == "1": + # Real opt-out: drop our seeded default (marker present); explicit + # user values carry no marker and are kept. + if os.environ.get("UNSLOTH_BNB_ROCM_VERSION_SOURCE") == "sitecustomize": + os.environ.pop("BNB_ROCM_VERSION", None) + os.environ.pop("UNSLOTH_BNB_ROCM_VERSION_SOURCE", None) + return None + if "BNB_ROCM_VERSION" in os.environ and ( + os.environ.get("UNSLOTH_BNB_ROCM_VERSION_SOURCE") != "sitecustomize" + ): + return None + if not _is_hip_torch_build(): + return None + version = _detect_installed_bnb_rocm_version() + if version is None: + return None + os.environ["BNB_ROCM_VERSION"] = version + os.environ["UNSLOTH_BNB_ROCM_VERSION_SOURCE"] = "detected" + if UNSLOTH_ENABLE_LOGGING: + logger.info( + f"Unsloth: set BNB_ROCM_VERSION={version} " + "(detected from the installed bitsandbytes ROCm wheel on Windows)." + ) + return version