From 3d15865bbc569ce1e3847caa9191390e53d55ac9 Mon Sep 17 00:00:00 2001 From: danielhanchen Date: Sun, 4 Jan 2026 12:21:39 +0000 Subject: [PATCH 1/6] rl.py fixes: buffer reset, safer attribute access, typo fix 1. Auto-reset gradient checkpointing buffers after trainer.train() - Import and call reset_unsloth_gradient_checkpointing_buffers() in prepare_for_training_mode wrapper to free memory after training while keeping buffers ready for subsequent runs 2. Replace eval/exec with safer getattr/setattr - eval(f"trl.trainer.{trainer}") -> getattr(trl.trainer, trainer) - exec(f"...{unwrap} = ...") -> setattr(current_trainer, unwrap, ...) - exec(f"Trainer.prediction_step=...") -> direct assignment 3. Fix psutil.cpu_count() potentially returning None - Change psutil.cpu_count()+4 to (psutil.cpu_count() or 1)+4 - Prevents TypeError on systems where cpu_count() returns None 4. Fix typo: oriignal_is_vlm_text -> original_is_vlm_text --- unsloth/models/rl.py | 18 ++++++++++++------ 1 file changed, 12 insertions(+), 6 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 22189f459c..9ea57e32d3 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -199,15 +199,15 @@ def PatchRL(FastLanguageModel): unwrap = "unwrap_model_for_generation" for trainer in trainers: try: - current_trainer = eval(f"trl.trainer.{trainer}") + current_trainer = getattr(trl.trainer, trainer) except: continue if hasattr(current_trainer, unwrap): try: - exec(f"trl.trainer.{trainer}.{unwrap} = unsloth_{unwrap}") + setattr(current_trainer, unwrap, unsloth_unwrap_model_for_generation) except: continue - exec(f"Trainer.prediction_step=unsloth_prediction_step") + Trainer.prediction_step = unsloth_prediction_step selective_log_softmax = RL_REPLACEMENTS["selective_log_softmax"] @@ -234,6 +234,7 @@ from transformers.training_args import ParallelMode # Also patches W&B since multiple runs must use wandb.finish() import functools from types import MethodType +from unsloth_zoo.gradient_checkpointing import reset_unsloth_gradient_checkpointing_buffers def prepare_for_training_mode(f): @functools.wraps(f) def wrapper(self, *args, **kwargs): @@ -244,6 +245,11 @@ def prepare_for_training_mode(f): # Return inference mode if hasattr(self, 'model') and hasattr(self.model, "for_inference"): self.model.for_inference() + # Reset gradient checkpointing buffers to free memory while staying ready for next run + try: + reset_unsloth_gradient_checkpointing_buffers() + except: + pass # Patch W&B to enable logging on future runs, otherwise it'll overwrite the first run try: import wandb @@ -817,7 +823,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): num_proc_check = ( "if dataset_num_proc is None:\n" " import psutil\n" - " dataset_num_proc = min(max(psutil.cpu_count()+4, 2), 64)\n" + " dataset_num_proc = min(max((psutil.cpu_count() or 1)+4, 2), 64)\n" " memory_gb_left = psutil.virtual_memory().available / (1024**3)\n" " if memory_gb_left <= 4: dataset_num_proc = 1 # Too risky, so set to 1\n" " elif memory_gb_left <= 6: dataset_num_proc = min(2, dataset_num_proc)\n" @@ -994,10 +1000,10 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): # Temporary patch _is_vlm to False # as of 0.22 it only exists in sfttrainer - oriignal_is_vlm_text = "self._is_vlm = True" + original_is_vlm_text = "self._is_vlm = True" new_is_vlm_text = "self._is_vlm = False" RLTrainer_source = RLTrainer_source.replace( - oriignal_is_vlm_text, new_is_vlm_text + original_is_vlm_text, new_is_vlm_text ) # Remove multiple doc strings From 08d619fca1a61774db4f2225b96723c9ffa1573c Mon Sep 17 00:00:00 2001 From: danielhanchen Date: Sun, 4 Jan 2026 12:57:10 +0000 Subject: [PATCH 2/6] Handle older unsloth-zoo without reset_unsloth_gradient_checkpointing_buffers --- unsloth/models/rl.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 9ea57e32d3..88aeeda8a1 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -234,7 +234,10 @@ from transformers.training_args import ParallelMode # Also patches W&B since multiple runs must use wandb.finish() import functools from types import MethodType -from unsloth_zoo.gradient_checkpointing import reset_unsloth_gradient_checkpointing_buffers +try: + from unsloth_zoo.gradient_checkpointing import reset_unsloth_gradient_checkpointing_buffers +except: + def reset_unsloth_gradient_checkpointing_buffers(): pass def prepare_for_training_mode(f): @functools.wraps(f) def wrapper(self, *args, **kwargs): From d31ec48a94482fa1265ebc670866acf700d8226a Mon Sep 17 00:00:00 2001 From: danielhanchen Date: Sun, 4 Jan 2026 12:58:45 +0000 Subject: [PATCH 3/6] Fix psutil.cpu_count() potentially returning None in save.py --- unsloth/save.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/unsloth/save.py b/unsloth/save.py index 714df7682a..071e032c53 100644 --- a/unsloth/save.py +++ b/unsloth/save.py @@ -879,12 +879,12 @@ def install_llama_cpp_make_non_blocking(): IS_CMAKE = False if check == 0: # Uses old MAKE - n_jobs = max(int(psutil.cpu_count() * 1.5), 1) + n_jobs = max(int((psutil.cpu_count() or 1) * 1.5), 1) full_command = ["make", "all", "-j" + str(n_jobs), "-C", "llama.cpp"] IS_CMAKE = False else: # Uses new CMAKE - n_jobs = max(int(psutil.cpu_count()), 1) # Use less CPUs since 1.5x faster + n_jobs = max(int(psutil.cpu_count() or 1), 1) # Use less CPUs since 1.5x faster check = os.system( f"cmake llama.cpp -B llama.cpp/build -DBUILD_SHARED_LIBS=OFF -DGGML_CUDA=OFF {CURL_FLAG}" ) @@ -994,13 +994,13 @@ def install_llama_cpp_old(version = -10): # Try using MAKE commands = [ "make clean -C llama.cpp", - f"make all -j{psutil.cpu_count()*2} -C llama.cpp", + f"make all -j{(psutil.cpu_count() or 1)*2} -C llama.cpp", ] if try_execute(commands) == "CMAKE": # Instead use CMAKE commands = [ f"cmake llama.cpp -B llama.cpp/build -DBUILD_SHARED_LIBS=OFF -DGGML_CUDA=OFF {CURL_FLAG}", - f"cmake --build llama.cpp/build --config Release -j{psutil.cpu_count()*2} --clean-first --target {' '.join(LLAMA_CPP_TARGETS)}", + f"cmake --build llama.cpp/build --config Release -j{(psutil.cpu_count() or 1)*2} --clean-first --target {' '.join(LLAMA_CPP_TARGETS)}", "cp llama.cpp/build/bin/llama-* llama.cpp", "rm -rf llama.cpp/build", ] @@ -1040,14 +1040,14 @@ def install_llama_cpp_blocking(use_cuda = False): "make clean -C llama.cpp", # https://github.com/ggerganov/llama.cpp/issues/7062 # Weirdly GPU conversion for GGUF breaks?? - # f"{use_cuda} make all -j{psutil.cpu_count()*2} -C llama.cpp", - f"make all -j{psutil.cpu_count()*2} -C llama.cpp", + # f"{use_cuda} make all -j{(psutil.cpu_count() or 1)*2} -C llama.cpp", + f"make all -j{(psutil.cpu_count() or 1)*2} -C llama.cpp", ] if try_execute(commands) == "CMAKE": # Instead use CMAKE commands = [ f"cmake llama.cpp -B llama.cpp/build -DBUILD_SHARED_LIBS=OFF -DGGML_CUDA=OFF {CURL_FLAG}", - f"cmake --build llama.cpp/build --config Release -j{psutil.cpu_count()*2} --clean-first --target {' '.join(LLAMA_CPP_TARGETS)}", + f"cmake --build llama.cpp/build --config Release -j{(psutil.cpu_count() or 1)*2} --clean-first --target {' '.join(LLAMA_CPP_TARGETS)}", "cp llama.cpp/build/bin/llama-* llama.cpp", "rm -rf llama.cpp/build", ] From 402e7d6285a3610350953c5f2d954fc5c8ddd3d2 Mon Sep 17 00:00:00 2001 From: danielhanchen Date: Sun, 4 Jan 2026 13:03:06 +0000 Subject: [PATCH 4/6] Respect user quantization_config --- unsloth/models/loader.py | 152 ++++++++++++++++++++++++++++++--------- unsloth/models/vision.py | 7 +- 2 files changed, 122 insertions(+), 37 deletions(-) diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index 247c72f43f..645c23d50b 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -151,8 +151,41 @@ class FastLanguageModel(FastLlamaModel): *args, **kwargs, ): + # Respect user-provided quantization_config (e.g. BitsAndBytesConfig) + quantization_config = kwargs.get("quantization_config", None) + if quantization_config is not None: + if getattr(quantization_config, "load_in_4bit", False): + load_in_4bit = True + load_in_8bit = False + if getattr(quantization_config, "load_in_8bit", False): + load_in_8bit = True + load_in_4bit = False + + load_in_4bit_kwargs = load_in_4bit + load_in_8bit_kwargs = load_in_8bit + if quantization_config is not None: + load_in_4bit_kwargs = False + load_in_8bit_kwargs = False + # Login to allow private models token = hf_login(token) + # Align dtype with bnb_4bit_compute_dtype if provided and dtype is unset. + if dtype is None and quantization_config is not None: + bnb_compute_dtype = None + if isinstance(quantization_config, dict): + if quantization_config.get("load_in_4bit", False): + bnb_compute_dtype = quantization_config.get( + "bnb_4bit_compute_dtype", None + ) + else: + if getattr(quantization_config, "load_in_4bit", False): + bnb_compute_dtype = getattr( + quantization_config, "bnb_4bit_compute_dtype", None + ) + if isinstance(bnb_compute_dtype, str): + bnb_compute_dtype = getattr(torch, bnb_compute_dtype, None) + if isinstance(bnb_compute_dtype, torch.dtype): + dtype = bnb_compute_dtype if load_in_8bit or full_finetuning or qat_scheme is not None: return FastModel.from_pretrained( model_name = model_name, @@ -546,7 +579,7 @@ class FastLanguageModel(FastLlamaModel): model_name = model_name, max_seq_length = max_seq_length, dtype = _get_dtype(dtype), - load_in_4bit = load_in_4bit, + load_in_4bit = load_in_4bit_kwargs, token = token, device_map = device_map, rope_scaling = rope_scaling, @@ -583,22 +616,30 @@ class FastLanguageModel(FastLlamaModel): ) if load_in_4bit: - # Fix up bitsandbytes config - compute_dtype = dtype_from_config(model.config) - quantization_config = { - # Sometimes compute_dtype is not a string!! - "bnb_4bit_compute_dtype": compute_dtype, - "bnb_4bit_quant_type": "nf4", - "bnb_4bit_use_double_quant": True, - "llm_int8_enable_fp32_cpu_offload": False, - "llm_int8_has_fp16_weight": False, - "llm_int8_skip_modules": None, - "llm_int8_threshold": 6.0, - "load_in_4bit": True, - "load_in_8bit": False, - "quant_method": "bitsandbytes", - } - model.config.update({"quantization_config": quantization_config}) + # Fix up bitsandbytes config, but respect user-provided quantization_config + if quantization_config is None: + compute_dtype = dtype_from_config(model.config) + quantization_config = { + # Sometimes compute_dtype is not a string!! + "bnb_4bit_compute_dtype": compute_dtype, + "bnb_4bit_quant_type": "nf4", + "bnb_4bit_use_double_quant": True, + "llm_int8_enable_fp32_cpu_offload": False, + "llm_int8_has_fp16_weight": False, + "llm_int8_skip_modules": None, + "llm_int8_threshold": 6.0, + "load_in_4bit": True, + "load_in_8bit": False, + "quant_method": "bitsandbytes", + } + model.config.update({"quantization_config": quantization_config}) + else: + if hasattr(quantization_config, "to_dict"): + model.config.update( + {"quantization_config": quantization_config.to_dict()} + ) + elif isinstance(quantization_config, dict): + model.config.update({"quantization_config": quantization_config}) if load_in_fp8 != False: _tag_model_with_fp8_torchao_config(model, fp8_mode) @@ -690,12 +731,45 @@ class FastModel(FastBaseModel): *args, **kwargs, ): + # Respect user-provided quantization_config (e.g. BitsAndBytesConfig) + quantization_config = kwargs.get("quantization_config", None) + if quantization_config is not None: + if getattr(quantization_config, "load_in_4bit", False): + load_in_4bit = True + load_in_8bit = False + if getattr(quantization_config, "load_in_8bit", False): + load_in_8bit = True + load_in_4bit = False + + load_in_4bit_kwargs = load_in_4bit + load_in_8bit_kwargs = load_in_8bit + if quantization_config is not None: + load_in_4bit_kwargs = False + load_in_8bit_kwargs = False + # Login to allow private models token = hf_login(token) if whisper_language is not None: assert type(whisper_language) is str if whisper_task is not None: assert type(whisper_task) is str + # Align dtype with bnb_4bit_compute_dtype if provided and dtype is unset. + if dtype is None and quantization_config is not None: + bnb_compute_dtype = None + if isinstance(quantization_config, dict): + if quantization_config.get("load_in_4bit", False): + bnb_compute_dtype = quantization_config.get( + "bnb_4bit_compute_dtype", None + ) + else: + if getattr(quantization_config, "load_in_4bit", False): + bnb_compute_dtype = getattr( + quantization_config, "bnb_4bit_compute_dtype", None + ) + if isinstance(bnb_compute_dtype, str): + bnb_compute_dtype = getattr(torch, bnb_compute_dtype, None) + if isinstance(bnb_compute_dtype, torch.dtype): + dtype = bnb_compute_dtype SUPPORTS_BFLOAT16 = is_bfloat16_supported() if dtype is None: dtype = torch.float16 if not SUPPORTS_BFLOAT16 else torch.bfloat16 @@ -1173,8 +1247,8 @@ class FastModel(FastBaseModel): model_name = model_name, max_seq_length = max_seq_length, dtype = _get_dtype(dtype), - load_in_4bit = load_in_4bit, - load_in_8bit = load_in_8bit, + load_in_4bit = load_in_4bit_kwargs, + load_in_8bit = load_in_8bit_kwargs, load_in_16bit = load_in_16bit, full_finetuning = full_finetuning, token = token, @@ -1220,22 +1294,30 @@ class FastModel(FastBaseModel): ) if load_in_4bit: - # Fix up bitsandbytes config - compute_dtype = dtype_from_config(model.config) - quantization_config = { - # Sometimes compute_dtype is not a string!! - "bnb_4bit_compute_dtype": compute_dtype, - "bnb_4bit_quant_type": "nf4", - "bnb_4bit_use_double_quant": True, - "llm_int8_enable_fp32_cpu_offload": False, - "llm_int8_has_fp16_weight": False, - "llm_int8_skip_modules": None, - "llm_int8_threshold": 6.0, - "load_in_4bit": True, - "load_in_8bit": False, - "quant_method": "bitsandbytes", - } - model.config.update({"quantization_config": quantization_config}) + # Fix up bitsandbytes config, but respect user-provided quantization_config + if quantization_config is None: + compute_dtype = dtype_from_config(model.config) + quantization_config = { + # Sometimes compute_dtype is not a string!! + "bnb_4bit_compute_dtype": compute_dtype, + "bnb_4bit_quant_type": "nf4", + "bnb_4bit_use_double_quant": True, + "llm_int8_enable_fp32_cpu_offload": False, + "llm_int8_has_fp16_weight": False, + "llm_int8_skip_modules": None, + "llm_int8_threshold": 6.0, + "load_in_4bit": True, + "load_in_8bit": False, + "quant_method": "bitsandbytes", + } + model.config.update({"quantization_config": quantization_config}) + else: + if hasattr(quantization_config, "to_dict"): + model.config.update( + {"quantization_config": quantization_config.to_dict()} + ) + elif isinstance(quantization_config, dict): + model.config.update({"quantization_config": quantization_config}) if load_in_fp8 != False: _tag_model_with_fp8_torchao_config(model, fp8_mode) diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index 6c5356e0b9..6de942d7d2 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -529,6 +529,7 @@ class FastBaseModel: del kwargs["attn_implementation"] bnb_config = None + user_quantization_config = kwargs.get("quantization_config", None) if full_finetuning and (load_in_4bit or load_in_8bit): print( "Unsloth: You selected full finetuning support, but 4bit / 8bit is enabled - disabling LoRA / QLoRA." @@ -596,7 +597,8 @@ class FastBaseModel: ): pass else: - kwargs["quantization_config"] = bnb_config + if user_quantization_config is None: + kwargs["quantization_config"] = bnb_config else: if auto_config is None: auto_config = AutoConfig.from_pretrained( @@ -641,7 +643,8 @@ class FastBaseModel: ) except: pass - kwargs["quantization_config"] = quantization_config + if user_quantization_config is None: + kwargs["quantization_config"] = quantization_config # Check if using forced float32 - we load it in bfloat16, then cast to float16! torch_dtype = dtype From bfa225b00c5faba087bae3f75d8aa55803cacb52 Mon Sep 17 00:00:00 2001 From: danielhanchen Date: Sun, 4 Jan 2026 13:14:03 +0000 Subject: [PATCH 5/6] Handle dict quantization_config flags --- unsloth/models/loader.py | 20 ++++++++++++++++---- 1 file changed, 16 insertions(+), 4 deletions(-) diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index 645c23d50b..4488ad9a07 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -154,10 +154,16 @@ class FastLanguageModel(FastLlamaModel): # Respect user-provided quantization_config (e.g. BitsAndBytesConfig) quantization_config = kwargs.get("quantization_config", None) if quantization_config is not None: - if getattr(quantization_config, "load_in_4bit", False): + if isinstance(quantization_config, dict): + q_load_in_4bit = quantization_config.get("load_in_4bit", False) + q_load_in_8bit = quantization_config.get("load_in_8bit", False) + else: + q_load_in_4bit = getattr(quantization_config, "load_in_4bit", False) + q_load_in_8bit = getattr(quantization_config, "load_in_8bit", False) + if q_load_in_4bit: load_in_4bit = True load_in_8bit = False - if getattr(quantization_config, "load_in_8bit", False): + if q_load_in_8bit: load_in_8bit = True load_in_4bit = False @@ -734,10 +740,16 @@ class FastModel(FastBaseModel): # Respect user-provided quantization_config (e.g. BitsAndBytesConfig) quantization_config = kwargs.get("quantization_config", None) if quantization_config is not None: - if getattr(quantization_config, "load_in_4bit", False): + if isinstance(quantization_config, dict): + q_load_in_4bit = quantization_config.get("load_in_4bit", False) + q_load_in_8bit = quantization_config.get("load_in_8bit", False) + else: + q_load_in_4bit = getattr(quantization_config, "load_in_4bit", False) + q_load_in_8bit = getattr(quantization_config, "load_in_8bit", False) + if q_load_in_4bit: load_in_4bit = True load_in_8bit = False - if getattr(quantization_config, "load_in_8bit", False): + if q_load_in_8bit: load_in_8bit = True load_in_4bit = False From e22ca346cf687cc79393f0a122ce4c15ba965e63 Mon Sep 17 00:00:00 2001 From: danielhanchen Date: Sun, 4 Jan 2026 13:18:15 +0000 Subject: [PATCH 6/6] Keep 4bit flag for fast_inference --- unsloth/models/loader.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index 4488ad9a07..eb3b21e206 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -167,12 +167,6 @@ class FastLanguageModel(FastLlamaModel): load_in_8bit = True load_in_4bit = False - load_in_4bit_kwargs = load_in_4bit - load_in_8bit_kwargs = load_in_8bit - if quantization_config is not None: - load_in_4bit_kwargs = False - load_in_8bit_kwargs = False - # Login to allow private models token = hf_login(token) # Align dtype with bnb_4bit_compute_dtype if provided and dtype is unset. @@ -581,6 +575,12 @@ class FastLanguageModel(FastLlamaModel): if fast_inference: fast_inference, model_name = fast_inference_setup(model_name, model_config) + load_in_4bit_kwargs = load_in_4bit + load_in_8bit_kwargs = load_in_8bit + if quantization_config is not None and not fast_inference: + load_in_4bit_kwargs = False + load_in_8bit_kwargs = False + model, tokenizer = dispatch_model.from_pretrained( model_name = model_name, max_seq_length = max_seq_length, @@ -753,12 +753,6 @@ class FastModel(FastBaseModel): load_in_8bit = True load_in_4bit = False - load_in_4bit_kwargs = load_in_4bit - load_in_8bit_kwargs = load_in_8bit - if quantization_config is not None: - load_in_4bit_kwargs = False - load_in_8bit_kwargs = False - # Login to allow private models token = hf_login(token) if whisper_language is not None: @@ -1255,6 +1249,12 @@ class FastModel(FastBaseModel): if auto_model is None: auto_model = AutoModelForVision2Seq if is_vlm else AutoModelForCausalLM + load_in_4bit_kwargs = load_in_4bit + load_in_8bit_kwargs = load_in_8bit + if quantization_config is not None and not fast_inference: + load_in_4bit_kwargs = False + load_in_8bit_kwargs = False + model, tokenizer = FastBaseModel.from_pretrained( model_name = model_name, max_seq_length = max_seq_length,