diff --git a/bitsandbytes/backends/cpu/ops.py b/bitsandbytes/backends/cpu/ops.py index 44fb5dceb..87795f679 100755 --- a/bitsandbytes/backends/cpu/ops.py +++ b/bitsandbytes/backends/cpu/ops.py @@ -16,6 +16,24 @@ _has_avx512 = torch.backends.cpu.get_cpu_capability() == "AVX512" + +def _load_gemm_4bit_forward_kernel(): + try: + from kernels import get_kernel + + return get_kernel( + "kernels-community/quantization-bitsandbytes", + version=1, + backend="cpu", + ).gemm_4bit_forward + except Exception as exc: # pragma: no cover - best effort fallback + logger.warning( + "Failed to load CPU gemm_4bit_forward from kernels-community: %s. Please make sure you already `pip install kernels` and the kernels >= 0.11.1", + exc, + ) + return None + + # torch._int_mm for s8@s8->s32 is supported on CPU from torch 2.4+. # However, we can overflow if we use this without AVX512_VNNI support. # This is fixed in torch 2.6+, so we set this as the minimum to be safe. @@ -240,19 +258,7 @@ def _( return out if has_avx512bf16(): - gemm_4bit_forward_kernel = None - try: - from kernels import get_kernel - - gemm_4bit_forward_kernel = get_kernel( - "kernels-community/quantization-bitsandbytes", version=1 - ).gemm_4bit_forward - except Exception as exc: # pragma: no cover - best effort fallback - gemm_4bit_forward_kernel = None - logger.warning( - "Failed to load CPU gemm_4bit_forward from kernels-community: %s. Please make sure you already `pip install kernels` and the kernels >= 0.11.1", - exc, - ) + gemm_4bit_forward_kernel = _load_gemm_4bit_forward_kernel() @register_kernel("bitsandbytes::gemv_4bit", "cpu") def _( diff --git a/tests/test_cpu_ops.py b/tests/test_cpu_ops.py new file mode 100644 index 000000000..8847b20d9 --- /dev/null +++ b/tests/test_cpu_ops.py @@ -0,0 +1,44 @@ +import logging +import sys +from types import ModuleType, SimpleNamespace + +import bitsandbytes.backends.cpu.ops as cpu_ops + + +def test_load_gemm_4bit_forward_kernel_requests_cpu_backend(monkeypatch): + kernels = ModuleType("kernels") + sentinel = object() + captured = {} + + def fake_get_kernel(repo_id, **kwargs): + captured["repo_id"] = repo_id + captured["kwargs"] = kwargs + return SimpleNamespace(gemm_4bit_forward=sentinel) + + kernels.get_kernel = fake_get_kernel + monkeypatch.setitem(sys.modules, "kernels", kernels) + + kernel = cpu_ops._load_gemm_4bit_forward_kernel() + + assert kernel is sentinel + assert captured == { + "repo_id": "kernels-community/quantization-bitsandbytes", + "kwargs": {"version": 1, "backend": "cpu"}, + } + + +def test_load_gemm_4bit_forward_kernel_logs_and_falls_back(monkeypatch, caplog): + kernels = ModuleType("kernels") + + def fake_get_kernel(*args, **kwargs): + raise FileNotFoundError("missing cpu variant") + + kernels.get_kernel = fake_get_kernel + monkeypatch.setitem(sys.modules, "kernels", kernels) + + with caplog.at_level(logging.WARNING, logger=cpu_ops.__name__): + kernel = cpu_ops._load_gemm_4bit_forward_kernel() + + assert kernel is None + assert "Failed to load CPU gemm_4bit_forward from kernels-community" in caplog.text + assert "missing cpu variant" in caplog.text