Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 19 additions & 13 deletions bitsandbytes/backends/cpu/ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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 _(
Expand Down
44 changes: 44 additions & 0 deletions tests/test_cpu_ops.py
Original file line number Diff line number Diff line change
@@ -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