Make bitsandbytes optional on ROCm and add bf16 helper

This commit is contained in:
Daniel Han-Chen 2026-02-08 16:20:03 +00:00
commit 7ecb049ac9
8 changed files with 87 additions and 25 deletions

View file

@ -205,6 +205,12 @@ elif DEVICE_TYPE == "xpu":
# set SUPPORTS_BFLOAT16 as torch.xpu.is_bf16_supported()
SUPPORTS_BFLOAT16 = torch.xpu.is_bf16_supported()
# Backwards compatibility: some notebooks import `unsloth.is_bf16_supported`.
# Ensure it exists on all backends (HIP / XPU) and has a stable signature.
if "is_bf16_supported" not in globals():
def is_bf16_supported(including_emulation = False):
return SUPPORTS_BFLOAT16
# For Gradio HF Spaces?
# if "SPACE_AUTHOR_NAME" not in os.environ and "SPACE_REPO_NAME" not in os.environ:
import triton

View file

@ -309,7 +309,10 @@ def fix_vllm_aimv2_issue():
spec = importlib.util.find_spec("vllm")
if spec is None:
return
vllm_version = importlib_version("vllm")
try:
vllm_version = importlib_version("vllm")
except Exception:
return
if Version(vllm_version) < Version("0.10.1"):
vllm_location = spec.origin
if vllm_location is None:

View file

@ -108,11 +108,14 @@ def calculate_settings(
HAS_CUDA_STREAM = False
import bitsandbytes as bnb
# https://github.com/bitsandbytes-foundation/bitsandbytes/pull/1330/files
HAS_CUDA_STREAM = Version(bnb.__version__) > Version("0.43.3")
get_ptr = bnb.functional.get_ptr
try:
import bitsandbytes as bnb
# https://github.com/bitsandbytes-foundation/bitsandbytes/pull/1330/files
HAS_CUDA_STREAM = Version(bnb.__version__) > Version("0.43.3")
get_ptr = bnb.functional.get_ptr
except Exception:
bnb = None
get_ptr = None
if DEVICE_TYPE == "xpu":
HAS_XPU_STREAM = True
@ -180,21 +183,29 @@ else:
CUDA_STREAMS = tuple(CUDA_STREAMS)
del _CUDA_STREAMS
# Bitsandbytes operations
# Bitsandbytes operations (optional)
ctypes_c_int = ctypes.c_int
ctypes_c_int32 = ctypes.c_int32
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
cdequantize_blockwise_fp16_nf4 = bnb.functional.lib.cdequantize_blockwise_fp16_nf4
cdequantize_blockwise_bf16_nf4 = bnb.functional.lib.cdequantize_blockwise_bf16_nf4
if DEVICE_TYPE == "xpu":
# https://github.com/bitsandbytes-foundation/bitsandbytes/blob/c3b8de268fdb55a88f92feada23fc811a1e6877a/bitsandbytes/backends/xpu/ops.py#L115
# for xpu, inference gemv using above link
cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemv_4bit_inference_fp16
cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemv_4bit_inference_bf16
if bnb is not None:
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
cdequantize_blockwise_fp16_nf4 = bnb.functional.lib.cdequantize_blockwise_fp16_nf4
cdequantize_blockwise_bf16_nf4 = bnb.functional.lib.cdequantize_blockwise_bf16_nf4
if DEVICE_TYPE == "xpu":
# https://github.com/bitsandbytes-foundation/bitsandbytes/blob/c3b8de268fdb55a88f92feada23fc811a1e6877a/bitsandbytes/backends/xpu/ops.py#L115
cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemv_4bit_inference_fp16
cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemv_4bit_inference_bf16
else:
cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemm_4bit_inference_naive_fp16
cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemm_4bit_inference_naive_bf16
else:
cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemm_4bit_inference_naive_fp16
cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemm_4bit_inference_naive_bf16
cdequantize_blockwise_fp32 = None
cdequantize_blockwise_fp16_nf4 = None
cdequantize_blockwise_bf16_nf4 = None
cgemm_4bit_inference_naive_fp16 = None
cgemm_4bit_inference_naive_bf16 = None
torch_device_stream = (

View file

@ -723,7 +723,10 @@ if is_openai_available():
# =============================================
# Get Flash Attention v2 if Ampere (RTX 30xx, A100)
import bitsandbytes as bnb
try:
import bitsandbytes as bnb
except Exception:
bnb = None
from transformers import AutoTokenizer
from transformers.utils.import_utils import _is_package_available
@ -1908,7 +1911,11 @@ def patch_tokenizer(model, tokenizer):
def patch_fast_lora():
import peft.tuners.lora.bnb
try:
import peft.tuners.lora.bnb
except Exception as e:
print("Unsloth: bitsandbytes/peft bnb not available - skipping 4bit LoRA patch.", repr(e))
return
peft.tuners.lora.bnb.Linear4bit.forward = fast_lora_forward

View file

@ -30,8 +30,19 @@ from .llama import (
LlamaLinearScalingRotaryEmbedding,
)
from .mistral import *
from bitsandbytes.nn import Linear4bit as Bnb_Linear4bit
from peft.tuners.lora import Linear4bit as Peft_Linear4bit
try:
from bitsandbytes.nn import Linear4bit as Bnb_Linear4bit
except Exception:
Bnb_Linear4bit = None
try:
from peft.tuners.lora import Linear4bit as Peft_Linear4bit
except Exception:
Peft_Linear4bit = None
_BNB_LINEAR_TYPES = tuple(
t for t in (Bnb_Linear4bit, Peft_Linear4bit) if t is not None
)
try:
from transformers.models.granite.modeling_granite import (
@ -575,7 +586,7 @@ class FastGraniteModel(FastLlamaModel):
correct_dtype = lm_head.weight.dtype
for name, module in model.named_modules():
if isinstance(module, (Bnb_Linear4bit, Peft_Linear4bit)):
if _BNB_LINEAR_TYPES and isinstance(module, _BNB_LINEAR_TYPES):
weight = module.weight
quant_state = weight.quant_state

View file

@ -3001,6 +3001,14 @@ class FastLlamaModel:
if not SUPPORTS_RSLORA:
del arguments["use_rslora"]
# PEFT API compatibility: only pass kwargs supported by the installed peft version.
try:
import inspect as _inspect
if "ensure_weight_tying" not in _inspect.signature(LoraConfig.__init__).parameters:
arguments.pop("ensure_weight_tying", None)
except Exception:
arguments.pop("ensure_weight_tying", None)
_saved_temp_tokenizer = model._saved_temp_tokenizer
lora_config = LoraConfig(**arguments)

View file

@ -1143,6 +1143,13 @@ def grpo_trainer_compute_loss(function_name, function):
if x.shape[1] == 1: # when importance_sampling_level == "sequence"
return x.mean()
else:
# Align mask/coef lengths when left-padding adds extra tokens.
if x.shape[1] != completion_mask.shape[1]:
min_len = min(x.shape[1], completion_mask.shape[1])
x = x[:, -min_len:]
cm = completion_mask[:, -min_len:]
denom = cm.sum().clamp(min = 1.0)
return (x * cm).sum() / denom
return (x * completion_mask).sum() / completion_token_count
if advantages.dim() == 1:

View file

@ -23,8 +23,15 @@ from unsloth_zoo.llama_cpp import (
check_llama_cpp,
_download_convert_hf_to_gguf,
)
from bitsandbytes.nn import Linear4bit as Bnb_Linear4bit
from peft.tuners.lora import Linear4bit as Peft_Linear4bit
try:
from bitsandbytes.nn import Linear4bit as Bnb_Linear4bit
except Exception:
Bnb_Linear4bit = None
try:
from peft.tuners.lora import Linear4bit as Peft_Linear4bit
except Exception:
Peft_Linear4bit = None
from peft.tuners.lora import Linear as Peft_Linear
from typing import Optional, Callable, Union, List
import sys
@ -58,6 +65,8 @@ except:
from pathlib import Path
from peft import PeftModelForCausalLM, PeftModel
_MERGE_LORA_LINEAR_TYPES = tuple(t for t in (Bnb_Linear4bit, Peft_Linear4bit, Peft_Linear) if t is not None)
__all__ = [
"print_quantization_methods",
"unsloth_save_model",
@ -188,7 +197,7 @@ def _free_cached_model(model):
def _merge_lora(layer, name):
bias = getattr(layer, "bias", None)
if isinstance(layer, (Bnb_Linear4bit, Peft_Linear4bit, Peft_Linear)):
if _MERGE_LORA_LINEAR_TYPES and isinstance(layer, _MERGE_LORA_LINEAR_TYPES):
# Is LoRA so we need to merge!
W, quant_state, A, B, s, bias = get_lora_parameters_bias(layer)
if quant_state is not None: