diff --git a/tests/test_import_fixes_drift.py b/tests/test_import_fixes_drift.py index 373830ae4a..6c75a4b886 100644 --- a/tests/test_import_fixes_drift.py +++ b/tests/test_import_fixes_drift.py @@ -584,3 +584,54 @@ def test_accelerate_utils_imports_module_present(): "accelerate.utils.imports.is_wandb_available is gone; " "disable_broken_wandb cannot patch the source module." ) + + +# =========================================================================== +# bitsandbytes -- ROCm arch / warp-size detection shape +# =========================================================================== + + +def test_bitsandbytes_rocm_detection_helpers_recognizable(): + """``fix_bitsandbytes_rocm_arch_detection`` swaps bnb's ROCm helpers + only when they shell out via subprocess and never consult torch device + props; a third shape is declined by design, silently restoring Windows + ROCm noise. Fail so the sniff gets updated. Reads source, no import.""" + spec = importlib.util.find_spec("bitsandbytes") + if spec is None: + pytest.skip("bitsandbytes not installed -- nothing to drift-check.") + cuda_specs_path = None + for location in spec.submodule_search_locations or []: + candidate = os.path.join(location, "cuda_specs.py") + if os.path.isfile(candidate): + cuda_specs_path = candidate + break + if cuda_specs_path is None: + pytest.skip("bitsandbytes has no cuda_specs.py (pre-ROCm version).") + + import ast + + with open(cuda_specs_path, "r", encoding = "utf-8") as f: + source = f.read() + helpers = [ + node + for node in ast.walk(ast.parse(source)) + if isinstance(node, ast.FunctionDef) + and node.name in ("get_rocm_gpu_arch", "get_rocm_warpsize") + ] + if not helpers: + pytest.skip("bitsandbytes cuda_specs has no ROCm detection helpers.") + for node in helpers: + segment = ast.get_source_segment(source, node) or "" + recognized = ( + "subprocess" in segment + or "get_device_properties" in segment + or "gcnArchName" in segment + ) + if not recognized: + pytest.fail( + f"DRIFT DETECTED: bitsandbytes.cuda_specs.{node.name} uses " + "neither subprocess nor torch device properties; " + "fix_bitsandbytes_rocm_arch_detection's shape sniff will " + "decline to patch it and Windows ROCm import-time noise / " + "wrong ROCM_GPU_ARCH may return." + ) diff --git a/unsloth/_gpu_init.py b/unsloth/_gpu_init.py index ede97ce68e..5da1e27a9d 100644 --- a/unsloth/_gpu_init.py +++ b/unsloth/_gpu_init.py @@ -30,6 +30,7 @@ from .import_fixes import ( disable_broken_causal_conv1d, disable_broken_vllm, configure_amdgpu_asic_id_table_path, + fix_bitsandbytes_rocm_arch_detection, torchvision_compatibility_check, fix_diffusers_warnings, fix_huggingface_hub, @@ -67,6 +68,8 @@ except Exception: # Configure libdrm ids table path early so ROCm can resolve AMD GPU names. configure_amdgpu_asic_id_table_path() +# Must precede `import unsloth_zoo` below, which imports bnb on ROCm. +fix_bitsandbytes_rocm_arch_detection() disable_broken_causal_conv1d() disable_broken_vllm() fix_message_factory_issue() @@ -75,6 +78,7 @@ torchvision_compatibility_check() fix_diffusers_warnings() fix_huggingface_hub() del configure_amdgpu_asic_id_table_path +del fix_bitsandbytes_rocm_arch_detection del disable_broken_causal_conv1d del disable_broken_vllm del fix_message_factory_issue diff --git a/unsloth/import_fixes.py b/unsloth/import_fixes.py index 0622f73965..695a6a577a 100644 --- a/unsloth/import_fixes.py +++ b/unsloth/import_fixes.py @@ -1823,6 +1823,340 @@ def configure_amdgpu_asic_id_table_path(): return None +# --------------------------------------------------------------------------- +# bitsandbytes Windows ROCm fix: cextension.py runs get_rocm_gpu_arch() +# (bnb >= 0.47) and get_rocm_warpsize() (0.49.x) at import, shelling out to +# rocminfo / hipinfo.exe via PATH. Neither is on PATH on Windows (AMD torch +# wheels put hipInfo.exe in venv Scripts), so every import logs ERROR + +# WARNING, ROCM_GPU_ARCH becomes "unknown", and warp size defaults to 64: +# wrong on RDNA (wave 32), breaking 4-bit blocksizes and +# ALLOW_PREQUANTIZED_MODELS. Upstream fix unmerged (bitsandbytes#1969), so a +# MetaPathFinder swaps both helpers for torch-device-props-first versions +# right after bitsandbytes.cuda_specs executes, before cextension reads +# them. Must run before `import unsloth_zoo` (imports bnb on ROCm). +# --------------------------------------------------------------------------- + +_BNB_CUDA_SPECS_MODULE = "bitsandbytes.cuda_specs" +_BNB_ROCM_FIX_FINDER_SENTINEL = "_unsloth_bnb_rocm_fix_finder" +_BNB_ROCM_FIX_FUNCTION_FLAG = "__unsloth_bnb_rocm_fix__" + + +def _torch_rocm_device_props(): + """Device-0 props on a ROCm torch build with a visible GPU, else None. + Never raises; bnb's own import initializes the device context anyway.""" + try: + import torch + + if not getattr(getattr(torch, "version", None), "hip", None): + return None + if not torch.cuda.is_available(): + return None + return torch.cuda.get_device_properties(0) + except Exception: + return None + + +def _iter_hipinfo_paths(): + """Yield existing hipInfo.exe paths: PATH, interpreter scripts dir (venv + and conda layouts), then HIP SDK / AMD installer locations.""" + import shutil + import sysconfig + + candidates = [] + try: + resolved = shutil.which("hipinfo.exe") + if resolved: + candidates.append(resolved) + except Exception: + pass + try: + scripts_dir = sysconfig.get_path("scripts") + if scripts_dir: + candidates.append(os.path.join(scripts_dir, "hipInfo.exe")) + except Exception: + pass + executable_dir = os.path.dirname(sys.executable or "") + if executable_dir: + candidates.append(os.path.join(executable_dir, "hipInfo.exe")) + candidates.append(os.path.join(executable_dir, "Scripts", "hipInfo.exe")) + for env_key in ("HIP_PATH", "ROCM_PATH"): + root = os.environ.get(env_key, "").strip() + if root: + candidates.append(os.path.join(root, "bin", "hipInfo.exe")) + rocm_root = os.path.join(os.environ.get("ProgramFiles", r"C:\Program Files"), "AMD", "ROCm") + try: + if os.path.isdir(rocm_root): + for version_dir in sorted(os.listdir(rocm_root), reverse = True): + candidates.append(os.path.join(rocm_root, version_dir, "bin", "hipInfo.exe")) + except Exception: + pass + + seen = set() + for candidate in candidates: + try: + key = os.path.normcase(os.path.normpath(candidate)) + if key in seen: + continue + seen.add(key) + if os.path.isfile(candidate): + yield candidate + except Exception: + continue + + +def _run_hipinfo(hipinfo_path): + """Run hipInfo.exe and return its stdout, or "" on any failure.""" + import subprocess + try: + result = subprocess.run( + [hipinfo_path], + capture_output = True, + text = True, + timeout = 15, + creationflags = getattr(subprocess, "CREATE_NO_WINDOW", 0), + ) + return result.stdout or "" + except Exception as e: + _log_rocm_detection(f"Unsloth: `{hipinfo_path}` failed: {e}") + return "" + + +def _unsloth_get_rocm_gpu_arch(): + """Replaces bnb's get_rocm_gpu_arch: torch device props first (no + subprocess), then hipInfo.exe by absolute path, then a quiet "unknown".""" + try: + import torch + if not getattr(getattr(torch, "version", None), "hip", None): + return "unknown" + except Exception: + return "unknown" + props = _torch_rocm_device_props() + if props is not None: + try: + # gcnArchName may carry feature flags, e.g. "gfx90a:sramecc+:xnack-" + arch = str(props.gcnArchName).split(":")[0].strip() + if arch.startswith("gfx"): + return arch + except Exception: + pass + for hipinfo_path in _iter_hipinfo_paths(): + match = re.search(r"gcnArchName:\s+gfx([a-zA-Z\d]+)", _run_hipinfo(hipinfo_path)) + if match: + return "gfx" + match.group(1) + _log_rocm_detection( + "Unsloth: Could not detect the ROCm GPU architecture - bitsandbytes will see `unknown`." + ) + return "unknown" + + +def _unsloth_get_rocm_warpsize(): + """Replaces bnb 0.49.x get_rocm_warpsize: upstream defaults to 64 when + rocminfo is missing, wrong on RDNA (wave 32).""" + try: + import torch + if not getattr(getattr(torch, "version", None), "hip", None): + return 32 # upstream behavior: NVIDIA warp size is always 32 + except Exception: + return 64 # upstream behavior: default to 64 on failure + props = _torch_rocm_device_props() + if props is not None: + # torch 2.11 ROCm exposes warp_size; some builds used warpSize. + for attribute_name in ("warp_size", "warpSize"): + warp_size = getattr(props, attribute_name, None) + if isinstance(warp_size, int) and warp_size in (32, 64): + return warp_size + for hipinfo_path in _iter_hipinfo_paths(): + match = re.search(r"^\s*warpSize:\s+(\d+)", _run_hipinfo(hipinfo_path), re.MULTILINE) + if match and int(match.group(1)) in (32, 64): + return int(match.group(1)) + _log_rocm_detection( + "Unsloth: Could not detect the ROCm warp size - defaulting to 64 " + "(bitsandbytes' own default)." + ) + return 64 + + +setattr(_unsloth_get_rocm_gpu_arch, _BNB_ROCM_FIX_FUNCTION_FLAG, True) +setattr(_unsloth_get_rocm_warpsize, _BNB_ROCM_FIX_FUNCTION_FLAG, True) + + +def _bnb_rocm_helper_is_broken(function): + """True only for upstream's subprocess-only detectors; co_names works + where getsource fails. Versions consulting torch props are untouched.""" + if function is None or not callable(function): + return False + if getattr(function, _BNB_ROCM_FIX_FUNCTION_FLAG, False): + return False # Already ours. + try: + function = inspect.unwrap(function) + except Exception: + pass + code = getattr(function, "__code__", None) + co_names = getattr(code, "co_names", ()) if code is not None else () + if not co_names: + return False # C function or opaque wrapper -- do not touch. + if "get_device_properties" in co_names or "gcnArchName" in co_names: + return False # Fixed upstream -- no-op. + return "subprocess" in co_names + + +def _patch_bnb_cuda_specs_module(module): + """Swap broken ROCm detection helpers on an executed cuda_specs module. + Returns True when the module ends up patched (now or previously).""" + patched = False + for attribute_name, replacement in ( + ("get_rocm_gpu_arch", _unsloth_get_rocm_gpu_arch), + ("get_rocm_warpsize", _unsloth_get_rocm_warpsize), + ): + original = getattr(module, attribute_name, None) + if getattr(original, _BNB_ROCM_FIX_FUNCTION_FLAG, False): + patched = True # Already ours. + continue + if not _bnb_rocm_helper_is_broken(original): + continue + setattr(module, attribute_name, replacement) + patched = True + logger.info( + f"Unsloth: Patched bitsandbytes.cuda_specs.{attribute_name} - " + f"avoids PATH-dependent subprocess GPU detection on Windows ROCm." + ) + return patched + + +class _BnbCudaSpecsPatchLoader(importlib.abc.Loader): + __slots__ = ("_loader",) + + def __init__(self, loader): + self._loader = loader + + def create_module(self, spec): + create_module = getattr(self._loader, "create_module", None) + if create_module is None: + return None + return create_module(spec) + + def exec_module(self, module): + self._loader.exec_module(module) + # Patch after the module body ran, before cextension calls it. The + # finder stays on sys.meta_path (same lifecycle as the blockers + # above) so importlib.reload(bitsandbytes.cuda_specs) re-patches. + try: + _patch_bnb_cuda_specs_module(module) + except Exception as e: + _log_rocm_detection(f"Unsloth: bitsandbytes ROCm detection patch failed: {e}") + + def __getattr__(self, name): + # Delegate get_source / get_filename etc. so introspection works. + return getattr(self._loader, name) + + +class _BnbCudaSpecsPatchFinder(importlib.abc.MetaPathFinder): + __slots__ = (_BNB_ROCM_FIX_FINDER_SENTINEL,) + + def __init__(self): + setattr(self, _BNB_ROCM_FIX_FINDER_SENTINEL, True) + + def find_spec( + self, + fullname, + path = None, + target = None, + ): + if fullname != _BNB_CUDA_SPECS_MODULE: + return None + # Delegate to remaining finders (editable installs, frozen apps) + # and wrap the loader that would actually be used. + spec = None + for finder in sys.meta_path: + if finder is self or getattr(finder, _BNB_ROCM_FIX_FINDER_SENTINEL, False): + continue + finder_find_spec = getattr(finder, "find_spec", None) + if finder_find_spec is None: + continue + try: + spec = finder_find_spec(fullname, path, target) + except Exception: + spec = None + if spec is not None: + break + if spec is None or spec.loader is None: + return None + if not hasattr(spec.loader, "exec_module"): + return None # Legacy loader -- let the stock machinery handle it. + spec.loader = _BnbCudaSpecsPatchLoader(spec.loader) + return spec + + +def _repair_imported_bitsandbytes_rocm_constants(): + """bnb imported before unsloth: noise already fired, but fix detectors + and cached constants, incl. by-value ROCM_WARP_SIZE_64 copies.""" + cuda_specs = sys.modules.get(_BNB_CUDA_SPECS_MODULE) + if cuda_specs is None: + return + if not _patch_bnb_cuda_specs_module(cuda_specs): + return + + try: + arch = cuda_specs.get_rocm_gpu_arch() + except Exception: + arch = "unknown" + warp_size_64 = None + get_rocm_warpsize = getattr(cuda_specs, "get_rocm_warpsize", None) + if callable(get_rocm_warpsize): + try: + warp_size_64 = get_rocm_warpsize() == 64 + except Exception: + warp_size_64 = None + + for module_name, module in list(sys.modules.items()): + if module is None or module is cuda_specs: + continue + if module_name != "bitsandbytes" and not module_name.startswith("bitsandbytes."): + continue + try: + if arch != "unknown" and getattr(module, "ROCM_GPU_ARCH", None) == "unknown": + module.ROCM_GPU_ARCH = arch + if warp_size_64 is not None and isinstance( + getattr(module, "ROCM_WARP_SIZE_64", None), bool + ): + module.ROCM_WARP_SIZE_64 = warp_size_64 + except Exception: + continue + logger.info("Unsloth: Repaired bitsandbytes ROCm arch / warp-size constants in place.") + + +def fix_bitsandbytes_rocm_arch_detection(): + """Fix bnb's import-time ROCm arch / warp-size detection on Windows + (see header above). No-op on non-Windows, non-ROCm, missing or + upstream-fixed bnb. Idempotent. Opt out: UNSLOTH_DISABLE_BNB_ROCM_FIX=1.""" + if os.environ.get("UNSLOTH_DISABLE_BNB_ROCM_FIX", "0") == "1": + return + if sys.platform != "win32": + return + if not _is_rocm_torch_build(): + return + + # Already imported: prevention impossible, repair in place instead. + if _BNB_CUDA_SPECS_MODULE in sys.modules: + try: + _repair_imported_bitsandbytes_rocm_constants() + except Exception: + pass + return + + try: + if importlib.util.find_spec("bitsandbytes") is None: + return + except Exception: + return + + for finder in sys.meta_path: + if getattr(finder, _BNB_ROCM_FIX_FINDER_SENTINEL, False): + return # Already installed -- idempotent. + sys.meta_path.insert(0, _BnbCudaSpecsPatchFinder()) + _log_rocm_detection("Unsloth: Installed the bitsandbytes ROCm arch detection patch hook.") + + def _is_causal_conv1d_name(module_name: str) -> bool: return module_name == _CAUSAL_CONV1D_PREFIX or module_name.startswith( _CAUSAL_CONV1D_PREFIX + "."