diff --git a/.github/ISSUE_TEMPLATE/bug---issue.md b/.github/ISSUE_TEMPLATE/bug---issue.md index 397d725f95..83e0fd73a9 100644 --- a/.github/ISSUE_TEMPLATE/bug---issue.md +++ b/.github/ISSUE_TEMPLATE/bug---issue.md @@ -18,4 +18,4 @@ assignees: '' Put Minimal code to reproduce error here ###Remove Hugging Face token### ``` -🦥 You can also ask via our Reddit page: https://www.reddit.com/r/unsloth/ +🦥 You can also ask via our Reddit page: https://reddit.com/r/unsloth/ diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 545c7899aa..cc188674b7 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,6 +1,6 @@ repos: - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.14.10 + rev: v0.14.11 hooks: - id: ruff args: diff --git a/README.md b/README.md index ae1fccfbba..ff8dcdeef6 100644 --- a/README.md +++ b/README.md @@ -53,8 +53,9 @@ Use our official [Unsloth Docker image](https://hub.docker.com/r/unsloth/unsloth For RTX 50x, B200, 6000 GPUs: `pip install unsloth`. Read our [Blackwell Guide](https://unsloth.ai/docs/basics/fine-tuning-llms-with-blackwell-rtx-50-series-and-unsloth) and [DGX Spark Guide](https://unsloth.ai/docs/basics/fine-tuning-llms-with-nvidia-dgx-spark-and-unsloth) for more details. ## 🦥 Unsloth News +- New 7x longer context reinforcement learning vs. all other setups, via our new batching algorithms. [Blog](https://unsloth.ai/docs/new/grpo-long-context) - New RoPE & MLP **Triton Kernels** & **Padding Free + Packing**: 3x faster training & 30% less VRAM. [Blog](https://unsloth.ai/docs/new/3x-faster-training-packing) -- **New Mistral**: Run Ministral 3 or Devstral 2 and fine-tune with vision/RL sodoku notebooks. [Guide](https://unsloth.ai/docs/models/ministral-3) • [Notebooks](https://unsloth.ai/docs/models/ministral-3#fine-tuning-ministral-3) +- **Mistral 3**: Run Ministral 3 or Devstral 2 and fine-tune with vision/RL sodoku notebooks. [Guide](https://unsloth.ai/docs/models/ministral-3) • [Notebooks](https://unsloth.ai/docs/models/ministral-3#fine-tuning-ministral-3) - **500K Context**: Training a 20B model with >500K context is now possible on an 80GB GPU. [Blog](https://unsloth.ai/docs/new/500k-context-length-fine-tuning) - **FP8 Reinforcement Learning**: You can now do FP8 GRPO on consumer GPUs. [Blog](https://unsloth.ai/docs/new/fp8-reinforcement-learning) • [Notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Qwen3_8B_FP8_GRPO.ipynb) - **DeepSeek-OCR**: Fine-tune to improve language understanding by 89%. [Guide](https://unsloth.ai/docs/models/deepseek-ocr-how-to-run-and-fine-tune) • [Notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Deepseek_OCR_(3B).ipynb) diff --git a/pyproject.toml b/pyproject.toml index 7fa249e64c..05fb690da5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -51,7 +51,7 @@ huggingfacenotorch = [ "sentencepiece>=0.2.0", "datasets>=3.4.1,!=4.0.*,!=4.1.0,<4.4.0", "accelerate>=0.34.1", - "peft>=0.7.1,!=0.11.0", + "peft>=0.18.0,!=0.11.0", "huggingface_hub>=0.34.0", "hf_transfer", "diffusers", @@ -60,7 +60,7 @@ huggingfacenotorch = [ ] huggingface = [ "unsloth[huggingfacenotorch]", - "unsloth_zoo>=2026.1.2", + "unsloth_zoo>=2026.1.3", "torchvision", "unsloth[triton]", ] @@ -523,7 +523,7 @@ colab-ampere-torch220 = [ "flash-attn>=2.6.3 ; ('linux' in sys_platform)", ] colab-new = [ - "unsloth_zoo>=2026.1.2", + "unsloth_zoo>=2026.1.3", "packaging", "tyro", "transformers>=4.51.3,!=4.52.0,!=4.52.1,!=4.52.2,!=4.52.3,!=4.53.0,!=4.54.0,!=4.55.0,!=4.55.1,!=4.57.0,<=4.57.3", @@ -542,7 +542,7 @@ colab-new = [ colab-no-deps = [ "accelerate>=0.34.1", "trl>=0.18.2,!=0.19.0,<=0.24.0", - "peft>=0.7.1", + "peft>=0.18.0", "xformers ; ('linux' in sys_platform or sys_platform == 'win32') and (platform_machine == 'AMD64' or platform_machine == 'x86_64')", "bitsandbytes>=0.45.5,!=0.46.0,!=0.48.0", "protobuf", diff --git a/tests/test_raw_text.py b/tests/test_raw_text.py new file mode 100644 index 0000000000..9f2e8cda4e --- /dev/null +++ b/tests/test_raw_text.py @@ -0,0 +1,172 @@ +#!/usr/bin/env python3 +""" +Minimal test for raw text training implementation. +Tests basic functionality without heavy dependencies. +""" + +import sys +import os +import tempfile +from pathlib import Path +import importlib.util + + +# Mock the datasets module since it's not installed +class MockDataset: + def __init__(self, data_dict): + self.data = data_dict + self.column_names = list(data_dict.keys()) + + def __len__(self): + return len(next(iter(self.data.values()))) + + def __getitem__(self, idx): + if isinstance(idx, str): + # Allow accessing columns by name like dataset['text'] + return self.data[idx] + elif isinstance(idx, int): + # Allow accessing individual rows by index + return {key: values[idx] for key, values in self.data.items()} + else: + raise TypeError(f"Invalid index type: {type(idx)}") + + @classmethod + def from_dict(cls, data_dict): + return cls(data_dict) + + +# Mock datasets module +datasets_mock = type(sys)("datasets") +datasets_mock.Dataset = MockDataset +sys.modules["datasets"] = datasets_mock + +# Import the raw_text module directly to avoid unsloth/__init__.py dependencies +current_dir = os.path.dirname(__file__) +raw_text_path = os.path.join( + os.path.dirname(current_dir), "unsloth", "dataprep", "raw_text.py" +) + +spec = importlib.util.spec_from_file_location("raw_text", raw_text_path) +raw_text_module = importlib.util.module_from_spec(spec) +spec.loader.exec_module(raw_text_module) + +RawTextDataLoader = raw_text_module.RawTextDataLoader +TextPreprocessor = raw_text_module.TextPreprocessor + + +def test_raw_text_loader(): + """Test basic RawTextDataLoader functionality.""" + + # Mock tokenizer for testing + class MockTokenizer: + def __init__(self): + self.eos_token = "" + self.eos_token_id = 2 # Mock EOS token ID + + def __call__(self, text, return_tensors = None, add_special_tokens = False): + words = text.split() + token_ids = list(range(len(words))) + + if return_tensors == "pt": + # Mock tensor-like object + class MockTensor: + def __init__(self, data): + self.data = data + + def __getitem__(self, idx): + return self.data + + def __len__(self): + return len(self.data) + + def tolist(self): + return self.data + + return {"input_ids": [MockTensor(token_ids)]} + return {"input_ids": token_ids} + + def decode(self, token_ids, skip_special_tokens = False): + return " ".join([f"word_{i}" for i in token_ids]) + + # Create test file + test_content = "This is a test file for raw text training. " * 10 + with tempfile.NamedTemporaryFile(mode = "w", suffix = ".txt", delete = False) as f: + f.write(test_content) + test_file = f.name + + try: + # Test loader + tokenizer = MockTokenizer() + loader = RawTextDataLoader(tokenizer, chunk_size = 5, stride = 2) + + # Test loading with text output (legacy mode) + text_dataset = loader.load_from_file(test_file, return_tokenized = False) + assert len(text_dataset) > 0, "Should create at least one chunk" + assert "text" in text_dataset.column_names, "Dataset should have 'text' column" + + # Test loading with tokenized output (new efficient mode) + tokenized_dataset = loader.load_from_file(test_file, return_tokenized = True) + assert len(tokenized_dataset) > 0, "Should create at least one tokenized chunk" + assert ( + "input_ids" in tokenized_dataset.column_names + ), "Dataset should have 'input_ids' column" + assert ( + "attention_mask" in tokenized_dataset.column_names + ), "Dataset should have 'attention_mask' column" + + # Verify tokenized data structure + first_sample = tokenized_dataset[0] + assert isinstance(first_sample["input_ids"], list), "input_ids should be a list" + assert isinstance( + first_sample["attention_mask"], list + ), "attention_mask should be a list" + assert len(first_sample["input_ids"]) == len( + first_sample["attention_mask"] + ), "input_ids and attention_mask should have same length" + + # Verify labels field exists (for causal LM training) + assert ( + "labels" in tokenized_dataset.column_names + ), "Dataset should have 'labels' column" + assert ( + first_sample["labels"] == first_sample["input_ids"] + ), "labels should match input_ids" + + # Test constructor validation + try: + bad_loader = RawTextDataLoader(tokenizer, chunk_size = 0, stride = 2) + assert False, "Should raise ValueError for chunk_size=0" + except ValueError as e: + assert "chunk_size must be positive" in str(e) + + try: + bad_loader = RawTextDataLoader(tokenizer, chunk_size = 5, stride = 10) + assert False, "Should raise ValueError for stride >= chunk_size" + except ValueError as e: + assert "stride" in str(e) and "chunk_size" in str(e) + + # Test preprocessor + preprocessor = TextPreprocessor() + clean_text = preprocessor.clean_text(" messy text \n\n\n ") + assert "messy text" in clean_text, "Should clean text properly" + + # Test validation + stats = preprocessor.validate_dataset(text_dataset) + assert stats["total_samples"] > 0, "Should count samples" + assert "warnings" in stats, "Should include warnings" + + print("✅ All tests passed!") + return True + + except Exception as e: + print(f"❌ Test failed: {e}") + return False + + finally: + # Cleanup + os.unlink(test_file) + + +if __name__ == "__main__": + success = test_raw_text_loader() + sys.exit(0 if success else 1) diff --git a/tests/utils/test_qat.py b/tests/utils/test_qat.py index 79251cf2ff..1083712d78 100644 --- a/tests/utils/test_qat.py +++ b/tests/utils/test_qat.py @@ -4,12 +4,19 @@ from typing import Dict import pytest import torch -from torchao.quantization.qat import FakeQuantizedLinear -from torchao.quantization.qat.fake_quantizer import ( - FakeQuantizerBase, - Float8FakeQuantizer, - Int4WeightPreshuffledFakeQuantizer, -) + +try: + from torchao.quantization.qat import FakeQuantizedLinear + from torchao.quantization.qat.fake_quantizer import ( + FakeQuantizerBase, + Float8FakeQuantizer, + Int4WeightFakeQuantizer, + IntxFakeQuantizer, + ) +except ImportError: + print( + "Missing torchao import, please install or upgrade torchao with: pip install 'torchao>=0.15.0'" + ) class _CountingFakeQuantizer(torch.nn.Module): @@ -49,14 +56,20 @@ def _test_linear_is_fake_quantized(linear: torch.nn.Linear, qat_scheme: str): """ Verify that the given linear contains fake quantizers according to the `qat_scheme`. """ + weight_only = False if qat_scheme == "fp8-int4": act_fq_class = Float8FakeQuantizer - weight_fq_class = Int4WeightPreshuffledFakeQuantizer + weight_fq_class = Int4WeightFakeQuantizer min_in_features = 128 elif qat_scheme == "fp8-fp8": act_fq_class = Float8FakeQuantizer weight_fq_class = Float8FakeQuantizer min_in_features = -1 + elif qat_scheme == "int8": + act_fq_class = None + weight_fq_class = IntxFakeQuantizer + min_in_features = 128 + weight_only = True else: raise ValueError(f"Unknown qat_scheme: {qat_scheme}") @@ -64,7 +77,8 @@ def _test_linear_is_fake_quantized(linear: torch.nn.Linear, qat_scheme: str): base_layer = getattr(linear, "base_layer", linear) if base_layer.in_features >= min_in_features: assert isinstance(base_layer, FakeQuantizedLinear) - assert isinstance(base_layer.activation_fake_quantizer, act_fq_class) + if not weight_only: + assert isinstance(base_layer.activation_fake_quantizer, act_fq_class) assert isinstance(base_layer.weight_fake_quantizer, weight_fq_class) # Check lora A and B (only for full_finetuning=False) @@ -73,11 +87,13 @@ def _test_linear_is_fake_quantized(linear: torch.nn.Linear, qat_scheme: str): lora_B = linear.lora_B.default if lora_A.in_features >= min_in_features: assert isinstance(lora_A, FakeQuantizedLinear) - assert isinstance(lora_A.activation_fake_quantizer, act_fq_class) + if not weight_only: + assert isinstance(lora_A.activation_fake_quantizer, act_fq_class) assert isinstance(lora_A.weight_fake_quantizer, weight_fq_class) if lora_B.in_features >= min_in_features: assert isinstance(lora_B, FakeQuantizedLinear) - assert isinstance(lora_B.activation_fake_quantizer, act_fq_class) + if not weight_only: + assert isinstance(lora_B.activation_fake_quantizer, act_fq_class) assert isinstance(lora_B.weight_fake_quantizer, weight_fq_class) @@ -85,10 +101,12 @@ def _test_fake_quantizers_are_called( model: torch.nn.Module, example_inputs: Dict, full_finetuning: bool, + qat_scheme: str, ): """ Verify that the fake quantizers are actually called when the model is called. """ + weight_only = qat_scheme == "int8" def _swap_fake_quantizers(model: torch.nn.Module): for name, child in model.named_children(): @@ -99,7 +117,8 @@ def _test_fake_quantizers_are_called( for name, child in model.named_children(): if full_finetuning: if isinstance(child, FakeQuantizedLinear): - assert child.activation_fake_quantizer.count == 1 + if not weight_only: + assert child.activation_fake_quantizer.count == 1 assert child.weight_fake_quantizer.count == 1 else: # For LoRA, we only fake quantize the input activations once per block: @@ -107,12 +126,14 @@ def _test_fake_quantizers_are_called( # For mlp, we only fake quantize the gate_proj's input activations if name == "self_attn": base_layer = child.q_proj.base_layer - assert hasattr(base_layer, "activation_fake_quantizer") - assert base_layer.activation_fake_quantizer.count == 1 + if not weight_only: + assert hasattr(base_layer, "activation_fake_quantizer") + assert base_layer.activation_fake_quantizer.count == 1 elif name == "mlp": base_layer = child.gate_proj.base_layer - assert hasattr(base_layer, "activation_fake_quantizer") - assert base_layer.activation_fake_quantizer.count == 1 + if not weight_only: + assert hasattr(base_layer, "activation_fake_quantizer") + assert base_layer.activation_fake_quantizer.count == 1 elif isinstance(child, FakeQuantizedLinear): # Weight fake quantizers should always be called assert child.weight_fake_quantizer.count == 1 @@ -124,7 +145,7 @@ def _test_fake_quantizers_are_called( model.apply(_assert_fake_quantizers_are_called) -def _test_model_fake_quantize(qat_scheme: bool, full_finetuning: bool): +def _test_model_fake_quantize(qat_scheme: str, full_finetuning: bool): """ Test that all linear layers in the model are fake quantized according to the `qat_scheme`. """ @@ -141,16 +162,16 @@ def _test_model_fake_quantize(qat_scheme: bool, full_finetuning: bool): _test_linear_is_fake_quantized(layer.mlp.up_proj, qat_scheme) _test_linear_is_fake_quantized(layer.mlp.down_proj, qat_scheme) inputs = tokenizer("How are you?", return_tensors = "pt") - _test_fake_quantizers_are_called(model, inputs, full_finetuning) + _test_fake_quantizers_are_called(model, inputs, full_finetuning, qat_scheme) # TODO: there are bad interactions across tests right now, need to figure out # how to disable model caching before re-enabling this test -@pytest.mark.parametrize("qat_scheme", ["fp8-int4", "fp8-fp8"]) -def _test_full_model_fake_quantize(qat_scheme: bool): +@pytest.mark.parametrize("qat_scheme", ["fp8-int4", "fp8-fp8", "int8"]) +def _test_full_model_fake_quantize(qat_scheme: str): _test_model_fake_quantize(qat_scheme, full_finetuning = True) -@pytest.mark.parametrize("qat_scheme", ["fp8-int4", "fp8-fp8"]) -def test_lora_model_fake_quantize(qat_scheme: bool): +@pytest.mark.parametrize("qat_scheme", ["fp8-int4", "fp8-fp8", "int8"]) +def test_lora_model_fake_quantize(qat_scheme: str): _test_model_fake_quantize(qat_scheme, full_finetuning = False) diff --git a/unsloth-cli.py b/unsloth-cli.py index 0222afe0c7..612da11eb2 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -41,6 +41,7 @@ def run(args): from unsloth import is_bfloat16_supported from unsloth.models.loader_utils import prepare_device_map import logging + from unsloth import RawTextDataLoader logging.getLogger("hf-to-gguf").setLevel(logging.WARNING) @@ -99,15 +100,36 @@ def run(args): texts.append(text) return {"text": texts} - use_modelscope = strtobool(os.environ.get("UNSLOTH_USE_MODELSCOPE", "False")) - if use_modelscope: - from modelscope import MsDataset + def load_dataset_smart(args): + from transformers.utils import strtobool - dataset = MsDataset.load(args.dataset, split = "train") - else: - # Load and format dataset - dataset = load_dataset(args.dataset, split = "train") - dataset = dataset.map(formatting_prompts_func, batched = True) + if args.raw_text_file: + # Use raw text loader + loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) + dataset = loader.load_from_file(args.raw_text_file) + elif args.dataset.endswith((".txt", ".md", ".json", ".jsonl")): + # Auto-detect local raw text files + loader = RawTextDataLoader(tokenizer) + dataset = loader.load_from_file(args.dataset) + else: + # Check for modelscope usage + use_modelscope = strtobool( + os.environ.get("UNSLOTH_USE_MODELSCOPE", "False") + ) + if use_modelscope: + from modelscope import MsDataset + + dataset = MsDataset.load(args.dataset, split = "train") + else: + # Existing HuggingFace dataset logic + dataset = load_dataset(args.dataset, split = "train") + + # Apply formatting for structured datasets + dataset = dataset.map(formatting_prompts_func, batched = True) + return dataset + + # Load dataset using smart loader + dataset = load_dataset_smart(args) print("Data is formatted and ready!") # Configure training arguments @@ -437,5 +459,15 @@ if __name__ == "__main__": help = "Token for pushing the model to Hugging Face hub", ) + parser.add_argument( + "--raw_text_file", type = str, help = "Path to raw text file for training" + ) + parser.add_argument( + "--chunk_size", type = int, default = 2048, help = "Size of text chunks for training" + ) + parser.add_argument( + "--stride", type = int, default = 512, help = "Overlap between chunks" + ) + args = parser.parse_args() run(args) diff --git a/unsloth/__init__.py b/unsloth/__init__.py index 5b571cd456..89824a25d1 100644 --- a/unsloth/__init__.py +++ b/unsloth/__init__.py @@ -279,6 +279,9 @@ from .save import * from .chat_templates import * from .tokenizer_utils import * from .trainer import * + +# Export dataprep utilities for CLI and downstream users +from .dataprep.raw_text import RawTextDataLoader, TextPreprocessor from unsloth_zoo.rl_environments import ( check_python_modules, create_locked_down_function, diff --git a/unsloth/dataprep/__init__.py b/unsloth/dataprep/__init__.py index b36122eb74..048f9b8010 100644 --- a/unsloth/dataprep/__init__.py +++ b/unsloth/dataprep/__init__.py @@ -13,3 +13,4 @@ # limitations under the License. from .synthetic import * +from .raw_text import * diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py new file mode 100644 index 0000000000..ba010edabb --- /dev/null +++ b/unsloth/dataprep/raw_text.py @@ -0,0 +1,348 @@ +# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +import re +import json +import csv +from typing import List, Dict, Any, Union, Optional +from datasets import Dataset +from pathlib import Path + +__all__ = [ + "RawTextDataLoader", + "TextPreprocessor", +] + +SUPPORTED_FORMATS = { + ".txt": "plain_text", + ".md": "markdown", + ".json": "json_lines", + ".jsonl": "json_lines", + ".csv": "csv_text_column", +} + + +class RawTextDataLoader: + def __init__(self, tokenizer, chunk_size = 2048, stride = 512, return_tokenized = True): + if chunk_size <= 0: + raise ValueError(f"chunk_size must be positive, got {chunk_size}") + if stride >= chunk_size: + raise ValueError( + f"stride ({stride}) must be smaller than chunk_size ({chunk_size})" + ) + self.tokenizer = tokenizer + self.chunk_size = chunk_size + self.stride = stride + self.return_tokenized = return_tokenized + + def detect_format(self, file_path): + """Auto-detect file format and parse accordingly""" + extension = Path(file_path).suffix.lower() + return SUPPORTED_FORMATS.get(extension, "plain_text") + + def load_from_file(self, file_path, return_tokenized = None): + """Load raw text and convert to dataset""" + if return_tokenized is None: + return_tokenized = self.return_tokenized + file_format = self.detect_format(file_path) + text_content = self._read_file_by_format(file_path, file_format) + if not text_content or not text_content.strip(): + raise ValueError(f"File '{file_path}' is empty or contains only whitespace") + chunks = self.smart_chunk_text( + text_content, self.chunk_size, self.stride, return_tokenized + ) + return self.create_causal_dataset(chunks) + + def load_from_files(self, file_paths, return_tokenized = None): + """Load multiple text files""" + if return_tokenized is None: + return_tokenized = self.return_tokenized + all_chunks = [] + for file_path in file_paths: + file_format = self.detect_format(file_path) + text_content = self._read_file_by_format(file_path, file_format) + chunks = self.smart_chunk_text( + text_content, self.chunk_size, self.stride, return_tokenized + ) + all_chunks.extend(chunks) + return self.create_causal_dataset(all_chunks) + + def chunk_text(self, text, return_tokenized = None): + """Split text into overlapping chunks""" + if return_tokenized is None: + return_tokenized = self.return_tokenized + return self.smart_chunk_text( + text, self.chunk_size, self.stride, return_tokenized + ) + + def create_causal_dataset(self, chunks): + """Create dataset for causal language modeling""" + if chunks and isinstance(chunks[0], dict): + # If chunks are already tokenized (dict with input_ids, attention_mask) + # Reorganize the data structure for Dataset.from_dict + input_ids = [chunk["input_ids"] for chunk in chunks] + attention_mask = [chunk["attention_mask"] for chunk in chunks] + # Labels are same as input_ids for causal LM training + labels = [list(ids) for ids in input_ids] + return Dataset.from_dict( + { + "input_ids": input_ids, + "attention_mask": attention_mask, + "labels": labels, + } + ) + else: + # If chunks are text strings (backward compatibility) + return Dataset.from_dict({"text": chunks}) + + def smart_chunk_text(self, text, chunk_size, stride, return_tokenized = True): + """ + Intelligent chunking that: + 1. Respects sentence/paragraph boundaries + 2. Handles various text formats (.txt, .md, .json, etc.) + 3. Maintains context with stride overlap + 4. Returns tokenized chunks directly (more efficient) or text chunks + """ + # First pass: tokenize the entire text to get accurate token counts + tokenized = self.tokenizer(text, return_tensors = "pt", add_special_tokens = False) + tokens = tokenized["input_ids"] + + # Handle different tokenizer return formats + if hasattr(tokens, "__len__") and len(tokens) > 0: + # If it's a nested structure, get the first element + if hasattr(tokens[0], "__len__"): + tokens = tokens[0] + elif isinstance(tokens, int): + # If tokenizer returns just a count, create a simple range + tokens = list(range(tokens)) + + if len(tokens) <= chunk_size: + # Text is small enough to fit in one chunk + if return_tokenized: + # Add EOS token to the tokens if available + eos_token_id = getattr(self.tokenizer, "eos_token_id", None) + if eos_token_id is not None: + tokens = ( + tokens.tolist() if hasattr(tokens, "tolist") else list(tokens) + ) + tokens.append(eos_token_id) + + # Create attention mask + attention_mask = [1] * len(tokens) + return [{"input_ids": tokens, "attention_mask": attention_mask}] + else: + eos_token = self.tokenizer.eos_token if self.tokenizer.eos_token else "" + return [text + eos_token] + + chunks = [] + start_idx = 0 + + while start_idx < len(tokens): + # Calculate end index for this chunk + end_idx = min(start_idx + chunk_size, len(tokens)) + + # Extract tokens for this chunk + chunk_tokens = tokens[start_idx:end_idx] + + if return_tokenized: + # Convert to list if it's a tensor + chunk_tokens_list = ( + chunk_tokens.tolist() + if hasattr(chunk_tokens, "tolist") + else list(chunk_tokens) + ) + + # Add EOS token if it's the last chunk or chunk is complete + if end_idx == len(tokens) or len(chunk_tokens_list) == chunk_size: + eos_token_id = getattr(self.tokenizer, "eos_token_id", None) + if eos_token_id is not None: + chunk_tokens_list.append(eos_token_id) + + # Create attention mask (all tokens are attended to) + attention_mask = [1] * len(chunk_tokens_list) + + chunks.append( + {"input_ids": chunk_tokens_list, "attention_mask": attention_mask} + ) + else: + # Decode back to text (backward compatibility) + chunk_text = self.tokenizer.decode( + chunk_tokens, skip_special_tokens = True + ) + + # Add EOS token if it's the last chunk or chunk is complete + if end_idx == len(tokens) or len(chunk_tokens) == chunk_size: + eos_token = ( + self.tokenizer.eos_token if self.tokenizer.eos_token else "" + ) + chunk_text += eos_token + + chunks.append(chunk_text) + + # Move to next chunk with stride overlap + if end_idx == len(tokens): + break + start_idx += chunk_size - stride + + return chunks + + def _read_file_by_format(self, file_path, file_format): + """Read file content based on detected format.""" + with open(file_path, "r", encoding = "utf-8") as f: + if file_format == "plain_text" or file_format == "markdown": + return f.read() + elif file_format == "json_lines": + lines = [] + for line in f: + try: + data = json.loads(line.strip()) + text = self._extract_text_from_json(data) + if text: + lines.append(text) + except json.JSONDecodeError: + continue + return "\n\n".join(lines) + elif file_format == "csv_text_column": + reader = csv.DictReader(f) + texts = [] + for row in reader: + text = self._extract_text_from_csv_row(row) + if text: + texts.append(text) + return "\n\n".join(texts) + return "" + + def _extract_text_from_json(self, data): + """Extract text from JSON object using common field names.""" + text_fields = ["text", "content", "message", "body", "description", "prompt"] + for field in text_fields: + if field in data and isinstance(data[field], str): + return data[field] + return "" + + def _extract_text_from_csv_row(self, row): + """Extract text from CSV row using common column names.""" + text_columns = ["text", "content", "message", "body", "description", "prompt"] + for column in text_columns: + if column in row and row[column]: + return row[column] + return "" + + +class TextPreprocessor: + def clean_text(self, text): + """Remove unwanted characters, normalize whitespace""" + text = re.sub(r"\s+", " ", text) + text = re.sub(r"[^\x20-\x7E\n\t]", "", text) + text = text.replace("\r\n", "\n").replace("\r", "\n") + text = re.sub(r"\n{3,}", "\n\n", text) + return text.strip() + + def extract_sections(self, text, patterns): + """Extract specific sections (e.g., code blocks, quotes)""" + sections = [] + for pattern in patterns: + matches = re.findall(pattern, text, re.MULTILINE | re.DOTALL) + sections.extend(matches) + return sections + + def add_structure_tokens(self, text): + """Add special tokens for structure (chapters, sections)""" + text = re.sub( + r"^# (.+)$", r"<|chapter|>\1<|/chapter|>", text, flags = re.MULTILINE + ) + text = re.sub( + r"^## (.+)$", r"<|section|>\1<|/section|>", text, flags = re.MULTILINE + ) + text = re.sub( + r"^### (.+)$", r"<|subsection|>\1<|/subsection|>", text, flags = re.MULTILINE + ) + text = re.sub( + r"```(\w*)\n(.*?)\n```", r"<|code|\1|>\2<|/code|>", text, flags = re.DOTALL + ) + return text + + def validate_dataset(self, dataset): + """ + Check for: + - Minimum/maximum sequence lengths + - Character encoding issues + - Repeated content + - Empty chunks + """ + stats = { + "total_samples": len(dataset), + "empty_samples": 0, + "min_length": float("inf"), + "max_length": 0, + "avg_length": 0, + "repeated_content": 0, + "encoding_issues": 0, + "warnings": [], + } + + texts = dataset["text"] + text_lengths = [] + seen_texts = set() + + for i, text in enumerate(texts): + if not text or len(text.strip()) == 0: + stats["empty_samples"] += 1 + continue + + # Check for encoding issues + try: + text.encode("utf-8") + except UnicodeEncodeError: + stats["encoding_issues"] += 1 + + # Calculate lengths + length = len(text) + text_lengths.append(length) + stats["min_length"] = min(stats["min_length"], length) + stats["max_length"] = max(stats["max_length"], length) + + # Check for repeated content + text_hash = hash(text.strip()) + if text_hash in seen_texts: + stats["repeated_content"] += 1 + else: + seen_texts.add(text_hash) + + # Calculate average length + if text_lengths: + stats["avg_length"] = sum(text_lengths) / len(text_lengths) + stats["min_length"] = ( + stats["min_length"] if stats["min_length"] != float("inf") else 0 + ) + + # Generate warnings + if stats["empty_samples"] > 0: + stats["warnings"].append(f"Found {stats['empty_samples']} empty samples") + + if stats["repeated_content"] > 0: + stats["warnings"].append( + f"Found {stats['repeated_content']} repeated samples" + ) + + if stats["encoding_issues"] > 0: + stats["warnings"].append( + f"Found {stats['encoding_issues']} encoding issues" + ) + + if stats["min_length"] < 10: + stats["warnings"].append("Some samples are very short (< 10 characters)") + + return stats diff --git a/unsloth/import_fixes.py b/unsloth/import_fixes.py index 1e05e462e9..27e5342e20 100644 --- a/unsloth/import_fixes.py +++ b/unsloth/import_fixes.py @@ -94,16 +94,34 @@ class HidePrintMessage: if os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") != "1": import sys - # Apply to stderr for FBGEMM + # Apply to stderr for FBGEMM and CUTLASS errors sys.stderr = HidePrintMessage(sys.stderr) # https://github.com/pytorch/FBGEMM/blob/d99cd96490ec4aabac2ee95b1e76ea4dcfcfa628/fbgemm_gpu/experimental/gemm/triton_gemm/utils.py#L43-L52 sys.stderr.add_filter("TMA benchmarks will be running") + # CUTLASS/FBGEMM MMA instruction error on SM90 vs SM100 (Blackwell) GPUs + # https://github.com/NVIDIA/cutlass/blob/main/include/cutlass/gemm/kernel/sm90_gemm_tma_warpspecialized.hpp + sys.stderr.add_filter("Arch conditional MMA instruction used without targeting") + # CUTLASS arch conditional errors for various architectures + sys.stderr.add_filter("CUTE_INVALID_CONTROL_PATH") + # CUTLASS TMA-related errors when not targeting correct architecture + sys.stderr.add_filter("Trying to use tma without CUTE_ARCH_TMA") # Skipping import of cpp extensions due to incompatible torch version 2.9.0+cu128 for torchao version 0.15.0 logging.getLogger("torchao").setLevel(logging.ERROR) + # Also filter torchao print to stderr about cpp extensions + sys.stderr.add_filter("Skipping import of cpp extensions") # SyntaxWarning: invalid escape sequence '\.' warnings.filterwarnings( "ignore", message = "invalid escape sequence", category = SyntaxWarning ) + # PYTORCH_CUDA_ALLOC_CONF is deprecated warning from torch + warnings.filterwarnings("ignore", message = "PYTORCH_CUDA_ALLOC_CONF is deprecated") + # TF32 precision deprecation warning from torch + warnings.filterwarnings( + "ignore", message = "Please use the new API settings to control TF32" + ) + # Deprecation warnings from torchao + warnings.filterwarnings("ignore", message = "`int4_weight_only` is deprecated") + warnings.filterwarnings("ignore", message = "`int8_weight_only` is deprecated") # Fix up AttributeError: 'MessageFactory' object has no attribute 'GetPrototype' @@ -323,10 +341,14 @@ def check_fbgemm_gpu_version(): except: return # We noticed some SegFault or bad alloc errors on lower versions of fbgemm_gpu. + # Instead of raising an error, disable FBGEMM and fall back to Triton kernels. if Version(fbgemm_gpu_version) < Version("1.4.0"): - raise ImportError( - f"Unsloth: fbgemm_gpu_genai=={fbgemm_gpu_version} detected. It might cause unexpected issues like segmentation faults. Please uninstall the current one by doing `pip uninstall fbgemm-gpu` && `pip install fbgemm-gpu` to install fbgemm-gpu 1.4.0 or newer!" + os.environ["UNSLOTH_HAS_FBGEMM"] = "0" + logger.info( + f"Unsloth: fbgemm_gpu_genai=={fbgemm_gpu_version} is old and may cause issues. " + f"Disabling FBGEMM - using Triton kernels instead." ) + return logger.info(f"Unsloth: fbgemm_gpu_genai=={fbgemm_gpu_version} detected.") diff --git a/unsloth/kernels/fp8.py b/unsloth/kernels/fp8.py index 3093bf61b1..e9f9161709 100644 --- a/unsloth/kernels/fp8.py +++ b/unsloth/kernels/fp8.py @@ -523,6 +523,7 @@ def fp8_fbgemm_block_linear(X, weight, weight_scale, bias = None): def test_has_fbgemm(): # We must manually check if the faster FBGEMM works on the specific GPU # For example RTX 5090 and RTX 4090 does not work + # Also SM100 (Blackwell B200/B100) GPUs fail with CUTLASS SM90 kernels # [TODO] Investigate with TorchAO why FBGEMM fails on consumer GPUs M, N, K = 128, 128, 128 xq = torch.ones(M, K, dtype = torch.float8_e4m3fn, device = "cuda") @@ -537,10 +538,25 @@ def test_has_fbgemm(): has_fbgemm = True del out except Exception as e: - e = str(e) - if "cutlass cannot initialize" in e.lower(): + error_str = str(e).lower() + # Catch any CUTLASS/CUDA errors and disable FBGEMM + # This includes MMA instruction errors, architecture mismatches, kernel launch failures, etc. + cutlass_cuda_errors = ( + "cutlass", + "cuda error", + "cuda runtime error", + "no kernel image", + "arch conditional", + "mma instruction", + "compute capability", + "cute_invalid_control_path", + "tma", + ) + is_cutlass_cuda_error = any(err in error_str for err in cutlass_cuda_errors) + + if is_cutlass_cuda_error: print( - f"Unsloth: FBGEMM on the current GPU cannot load - will switch to Triton kernels" + "Unsloth: FBGEMM on the current GPU cannot load - will switch to Triton kernels" ) else: print( diff --git a/unsloth/kernels/swiglu.py b/unsloth/kernels/swiglu.py index b321f5179e..b3ae9d40e6 100644 --- a/unsloth/kernels/swiglu.py +++ b/unsloth/kernels/swiglu.py @@ -128,7 +128,7 @@ def _DWf_DW_dfg_kernel( def swiglu_DWf_DW_dfg_kernel(DW, e, g): - batch_seq_len, hd = e.shape + batch_seq_len, hd = e.shape # Flattened to 2D, so 1st dim is bsz * seq_len n_elements = e.numel() grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),) with torch_gpu_device(e.device): diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index e6c4a12874..28dbe450ed 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -__version__ = "2026.1.2" +__version__ = "2026.1.3" __all__ = [ "SUPPORTS_BFLOAT16", @@ -175,6 +175,8 @@ warnings.filterwarnings(action = "ignore", category = UserWarning, module = "bit # Stop "Special tokens have been added in the vocabulary, ..." logging.getLogger("transformers.tokenization_utils_base").setLevel(logging.CRITICAL + 1) +TORCHAO_MSG = "Error: torchao not found, please install with `pip install torchao`" + # Ignore logging messages class HideLoggingMessage(logging.Filter): @@ -1106,53 +1108,66 @@ def _get_statistics(statistics = None, force_download = True): global USE_MODELSCOPE USE_MODELSCOPE = os.environ.get("UNSLOTH_USE_MODELSCOPE", "0") == "1" - if statistics is not None: - pass - elif "\nCOLAB_" in keynames and n_cpus == 1: - statistics = "colab" - elif "\nCOLAB_" in keynames: - statistics = "colabpro" - elif "\nKAGGLE_" in keynames: - statistics = "kaggle" - elif "\nRUNPOD_" in keynames: - statistics = "runpod" - elif "\nAWS_" in keynames: - statistics = "aws" - elif "\nAZURE_" in keynames: - statistics = "azure" - # elif "\nK_" in keynames or "\nFUNCTION_" in keynames: statistics = "gcp" - elif "\nINVOCATION_ID" in keynames: - statistics = "lambda" - # else: statistics = "other" - else: - - def try_vllm_check(): - vendor_files = ( - "/sys/class/dmi/id/product_version", - "/sys/class/dmi/id/bios_vendor", - "/sys/class/dmi/id/product_name", - "/sys/class/dmi/id/chassis_asset_tag", - "/sys/class/dmi/id/sys_vendor", - ) + if statistics is None: + # Prefer filesystem markers (harder to misidentify) before env-key matching + try: from pathlib import Path - for vendor_file in vendor_files: - path = Path(vendor_file) - if path.is_file(): - file_content = path.read_text().lower() - if "amazon" in file_content: - return "aws" - elif "microsoft corporation" in file_content: - return "azure" - elif "google" in file_content: - return "gcp" - return "other" + if Path("/kaggle/working").exists(): + statistics = "kaggle" + elif Path("/content").exists() and Path("/opt/colab").exists(): + statistics = "colab" if n_cpus == 1 else "colabpro" + elif Path("/runpod-volume").exists(): + statistics = "runpod" + except Exception: + pass + + # Fallback to env-key detection + if statistics is None: + if "\nKAGGLE_" in keynames: + statistics = "kaggle" + elif "\nCOLAB_" in keynames and n_cpus == 1: + statistics = "colab" + elif "\nCOLAB_" in keynames: + statistics = "colabpro" + elif "\nRUNPOD_" in keynames: + statistics = "runpod" + elif "\nAWS_" in keynames: + statistics = "aws" + elif "\nAZURE_" in keynames: + statistics = "azure" + # elif "\nK_" in keynames or "\nFUNCTION_" in keynames: statistics = "gcp" + elif "\nINVOCATION_ID" in keynames: + statistics = "lambda" + # else: statistics = "other" + else: + + def try_vllm_check(): + vendor_files = ( + "/sys/class/dmi/id/product_version", + "/sys/class/dmi/id/bios_vendor", + "/sys/class/dmi/id/product_name", + "/sys/class/dmi/id/chassis_asset_tag", + "/sys/class/dmi/id/sys_vendor", + ) + + for vendor_file in vendor_files: + path = Path(vendor_file) + if path.is_file(): + file_content = path.read_text().lower() + if "amazon" in file_content: + return "aws" + elif "microsoft corporation" in file_content: + return "azure" + elif "google" in file_content: + return "gcp" + return "other" + + try: + statistics = try_vllm_check() + except Exception: + statistics = "other" - pass - try: - statistics = try_vllm_check() - except: - statistics = "other" if statistics is not None: import tempfile from huggingface_hub import snapshot_download @@ -1184,7 +1199,7 @@ def _get_statistics(statistics = None, force_download = True): "model = FastLanguageModel.from_pretrained('unsloth/gpt-oss-20b')\n" "```" ) - except: + except Exception: # Try no time limit check stats_check() @@ -2198,9 +2213,12 @@ def _prepare_model_for_qat( QAT can be optionally combined with LoRA fine-tuning to for additional throughput improvement. For more details: https://dev-discuss.pytorch.org/t/speeding-up-qat-by-1-89x-with-lora/2700 """ - from torchao.quantization import PerRow, quantize_ - from torchao.quantization.granularity import PerGroup, PerAxis - from torchao.quantization.qat import QATConfig + try: + from torchao.quantization import PerRow, quantize_ + from torchao.quantization.granularity import PerGroup, PerAxis + from torchao.quantization.qat import QATConfig + except ImportError: + raise ImportError(TORCHAO_MSG) # Gemma3 models have issues with int8 embedding quantization due to their # large vocabulary size (262144). Auto-switch to int4 weight-only instead. @@ -2217,8 +2235,10 @@ def _prepare_model_for_qat( if not isinstance(qat_scheme, TorchAOConfig): torchao_config: Optional[TorchAOConfig] = None if qat_scheme == "fp8-int4": - from torchao.quantization import Float8DynamicActivationInt4WeightConfig - + try: + from torchao.quantization import Float8DynamicActivationInt4WeightConfig + except ImportError: + raise ImportError(TORCHAO_MSG) group_size = 128 base_config = Float8DynamicActivationInt4WeightConfig() filter_fn = ( @@ -2230,8 +2250,12 @@ def _prepare_model_for_qat( base_config_and_filter_fns = [(base_config, filter_fn)], ) elif qat_scheme == "fp8-fp8": - from torchao.quantization import Float8DynamicActivationFloat8WeightConfig - + try: + from torchao.quantization import ( + Float8DynamicActivationFloat8WeightConfig, + ) + except ImportError: + raise ImportError(TORCHAO_MSG) base_config = Float8DynamicActivationFloat8WeightConfig( granularity = PerRow() ) @@ -2239,11 +2263,13 @@ def _prepare_model_for_qat( qat_scheme = qat_scheme, base_config_and_filter_fns = [(base_config, None)] ) elif qat_scheme == "int8-int4": - from torchao.quantization import ( - Int8DynamicActivationIntxWeightConfig, - IntxWeightOnlyConfig, - ) - + try: + from torchao.quantization import ( + Int8DynamicActivationIntxWeightConfig, + IntxWeightOnlyConfig, + ) + except ImportError: + raise ImportError(TORCHAO_MSG) torchao_config = TorchAOConfig( qat_scheme = qat_scheme, base_config_and_filter_fns = [ @@ -2263,8 +2289,10 @@ def _prepare_model_for_qat( prequantization_transform = _untie_input_output_embeddings, ) elif qat_scheme == "int4": - from torchao.quantization import Int4WeightOnlyConfig - + try: + from torchao.quantization import Int4WeightOnlyConfig + except ImportError: + raise ImportError(TORCHAO_MSG) group_size = 128 base_config = Int4WeightOnlyConfig(group_size = group_size) filter_fn = ( @@ -2275,6 +2303,22 @@ def _prepare_model_for_qat( qat_scheme = qat_scheme, base_config_and_filter_fns = [(base_config, filter_fn)], ) + elif qat_scheme == "int8": + try: + from torchao.quantization import IntxWeightOnlyConfig + from torchao.quantization.granularity import PerAxis + except ImportError: + raise ImportError(TORCHAO_MSG) + + base_config = IntxWeightOnlyConfig( + weight_dtype = torch.int8, + granularity = PerAxis(0), + ) + filter_fn = lambda m, _: isinstance(m, torch.nn.Linear) + torchao_config = TorchAOConfig( + qat_scheme = qat_scheme, + base_config_and_filter_fns = [(base_config, filter_fn)], + ) else: raise ValueError(f"Unexpected QAT scheme {qat_scheme}") assert torchao_config is not None, f"TorchAOConfig was not set for {qat_scheme}" diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 92d51b73ad..39f2ba1460 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -146,6 +146,59 @@ torch_nn_functional_softmax = torch.nn.functional.softmax # SDPA has GQA internally SDPA_HAS_GQA = "enable_gqa" in scaled_dot_product_attention.__doc__ +from peft.utils.other import ModulesToSaveWrapper + + +def _offload_frozen_module_for_training( + module: ModulesToSaveWrapper, + device_type: str, + offload_device: str = "cpu", +) -> None: + """ + Offload frozen module to CPU and configure trainable copy for mixed precision training. + + This function optimizes memory usage by: + 1. Moving the trainable copy to the target device with appropriate precision + 2. Offloading the original frozen module to CPU/disk to free VRAM + 3. Converting float16 to float32 for compatibility with certain GPUs (e.g., Tesla T4) + + Args: + module: The module to configure. Must be a ModulesToSaveWrapper with a + `modules_to_save` attribute containing trainable and original modules. + device_type: Target device string for training (e.g., "cuda:0", "xpu:0") + offload_device: Device to offload frozen parameters (default: "cpu") + Note: Currently only "cpu" is supported; disk offloading is planned. + + Returns: + None (modifies module in-place) + + Note: + - Float16 weights are automatically promoted to float32 for GPU compatibility + - Original frozen parameters are moved to CPU to reduce active VRAM usage + - Future versions will support disk-based offloading for even larger models + + See Also: + - https://github.com/unslothai/unsloth/pull/1200 (Tesla T4 float32 requirement) + """ + # Early return with explicit None if module doesn't support mixed precision training + if not hasattr(module, "modules_to_save"): + return None + + new_dtype = module.modules_to_save.default.weight.dtype + if new_dtype == torch.float16: + # See https://github.com/unslothai/unsloth/pull/1200 + # Tesla T4 must use float32 and not float16 + new_dtype = torch.float32 + + module.modules_to_save.default.to( + device = device_type, dtype = new_dtype, non_blocking = True + ) + module.modules_to_save.default.requires_grad_(True) + + # [TODO] Move old module to CPU - should be disk! + module.original_module.to(device = offload_device, non_blocking = True) + module.original_module.requires_grad_(False) + # Fix new HF's inference code def _fast_prepare_inputs_for_generation( @@ -2711,46 +2764,16 @@ class FastLlamaModel: "Unsloth: Training embed_tokens in mixed precision to save VRAM" ) - new_dtype = model.get_input_embeddings().modules_to_save.default.weight.dtype - if new_dtype == torch.float16: - # See https://github.com/unslothai/unsloth/pull/1200 - # Tesla T4 must use float32 and not float16 - new_dtype = torch.float32 - - model.get_input_embeddings().modules_to_save.default.to( - device = DEVICE_TYPE_TORCH, dtype = new_dtype, non_blocking = True + _offload_frozen_module_for_training( + model.get_input_embeddings(), DEVICE_TYPE_TORCH ) - model.get_input_embeddings().modules_to_save.default.requires_grad_( - True - ) - - # [TODO] Move old embed_tokens to CPU - should be disk! - model.get_input_embeddings().original_module.to( - device = "cpu", non_blocking = True - ) - model.get_input_embeddings().original_module.requires_grad_(False) if "lm_head" in new_target_modules: print("Unsloth: Training lm_head in mixed precision to save VRAM") - new_dtype = model.get_output_embeddings().modules_to_save.default.weight.dtype - if new_dtype == torch.float16: - # See https://github.com/unslothai/unsloth/pull/1200 - # Tesla T4 must use float32 and not float16 - new_dtype = torch.float32 - - model.get_output_embeddings().modules_to_save.default.to( - device = DEVICE_TYPE_TORCH, dtype = new_dtype, non_blocking = True + _offload_frozen_module_for_training( + model.get_output_embeddings(), DEVICE_TYPE_TORCH ) - model.get_output_embeddings().modules_to_save.default.requires_grad_( - True - ) - - # [TODO] Move old lm_head to CPU - should be disk! - model.get_output_embeddings().original_module.to( - device = "cpu", non_blocking = True - ) - model.get_output_embeddings().original_module.requires_grad_(False) return model else: diff --git a/unsloth/models/loader_utils.py b/unsloth/models/loader_utils.py index fe2a89d893..1e5533c25c 100644 --- a/unsloth/models/loader_utils.py +++ b/unsloth/models/loader_utils.py @@ -408,7 +408,7 @@ def _get_fp8_mode_and_check_settings( if Version(torchao.__version__) < Version("0.15.0"): raise ValueError(error_message) - # If fbgemm_gpu_genai is installed, check if it's >= 1.4.1 + # If fbgemm_gpu_genai is installed and old, disable FBGEMM and use Triton instead if ( importlib.util.find_spec("fbgemm_gpu") is not None and importlib.util.find_spec("fbgemm_gpu.experimental") is not None @@ -416,7 +416,12 @@ def _get_fp8_mode_and_check_settings( import fbgemm_gpu.experimental.gen_ai if Version(fbgemm_gpu.__version__) < Version("1.4.1"): - raise ValueError( - "Unsloth: On the fly `load_in_fp8` is only compatible with fbgemm_gpu_genai 1.4.1+. Try `unsloth/Qwen3-8B` instead." + # Old FBGEMM version - disable and use Triton kernels instead + os.environ["UNSLOTH_HAS_FBGEMM"] = "0" + from unsloth_zoo.log import logger + + logger.info( + f"Unsloth: fbgemm_gpu_genai=={fbgemm_gpu.__version__} is old for FP8 loading. " + f"Using Triton kernels instead." ) return fp8_mode diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index e945a80354..9788207c99 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -231,11 +231,13 @@ def PatchRL(FastLanguageModel): Trainer.prediction_step = unsloth_prediction_step +grpo_selective_log_softmax = RL_REPLACEMENTS["grpo_selective_log_softmax"] selective_log_softmax = RL_REPLACEMENTS["selective_log_softmax"] calculate_pad_tokens_in_prompt = RL_REPLACEMENTS["calculate_pad_tokens_in_prompt"] create_completion_attention_mask = RL_REPLACEMENTS["create_completion_attention_mask"] left_pack_padding = RL_REPLACEMENTS["left_pack_padding"] align_logprobs_with_mask = RL_REPLACEMENTS["align_logprobs_with_mask"] +autotune_batch_and_chunks = RL_REPLACEMENTS["grpo_autotune_batch_and_chunks"] RLTrainer_replacement = ''' import os @@ -247,7 +249,6 @@ import numpy as np from contextlib import nullcontext from torch.nn import functional as F import inspect -import psutil from transformers import DataCollatorForSeq2Seq, DataCollatorForLanguageModeling as TransformersDataCollatorForLanguageModeling from transformers.training_args import ParallelMode @@ -264,17 +265,19 @@ def prepare_for_training_mode(f): def wrapper(self, *args, **kwargs): # Enable training mode _was_training = None + # Get gradient checkpointing setting from training arguments + use_gc = getattr(self.args, 'gradient_checkpointing', True) if hasattr(self, 'model') and hasattr(self.model, "training"): _was_training = self.model.training if hasattr(self, 'model') and hasattr(self.model, "for_training"): - self.model.for_training() + self.model.for_training(use_gradient_checkpointing=use_gc) output = f(self, *args, **kwargs) # Restore previous mode when possible if hasattr(self, 'model') and hasattr(self.model, "for_inference"): if _was_training is False: self.model.for_inference() elif _was_training is True and hasattr(self.model, "for_training"): - self.model.for_training() + self.model.for_training(use_gradient_checkpointing=use_gc) # Reset gradient checkpointing buffers to free memory while staying ready for next run try: reset_unsloth_gradient_checkpointing_buffers() @@ -298,11 +301,13 @@ torch_compile_options = {{ "triton.cudagraphs" : False, }} +{grpo_selective_log_softmax_code} {selective_log_softmax_code} {calculate_pad_tokens_in_prompt_code} {create_completion_attention_mask_code} {left_pack_padding_code} {align_logprobs_with_mask_code} +{autotune_batch_and_chunks_code} {RL_pre} @@ -319,10 +324,20 @@ class Unsloth{RLConfig_name}({RLConfig_name}): default = -1, metadata = {{'help': 'Chunk size to reduce memory usage. -1 is most efficient.'}}, ) + unsloth_logit_chunk_multiplier : Optional[int] = field( + default = None, + metadata = {{'help': 'Multiplier for chunked logit computations.'}}, + ) + unsloth_grpo_mini_batch : Optional[int] = field( + default = None, + metadata = {{'help': 'Mini batch size for GRPO hidden state accumulation. Default is None unless user defines it.'}}, + ) {max_seq_length_pre} def __init__({RLConfig_arguments}, vllm_sampling_params = None, unsloth_num_chunks = -1, + unsloth_logit_chunk_multiplier = None, + unsloth_grpo_mini_batch = None, {max_seq_length_call} **kwargs, ): @@ -330,6 +345,15 @@ class Unsloth{RLConfig_name}({RLConfig_name}): super().__init__({RLConfig_call_args}{RLConfig_kwargs}) self.vllm_sampling_params = vllm_sampling_params self.unsloth_num_chunks = unsloth_num_chunks + if unsloth_grpo_mini_batch is not None: + if self.generation_batch_size >= unsloth_grpo_mini_batch: + self.unsloth_grpo_mini_batch = unsloth_grpo_mini_batch + else: + raise ValueError( + f"Unsloth GRPO mini batch size needs to be less than or equal to the effective generation batch size, " + f"which is self.per_device_train_batch_size * gradient_accumulation_steps." + ) + self.unsloth_logit_chunk_multiplier = unsloth_logit_chunk_multiplier {max_seq_length_post} pass @@ -1027,6 +1051,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): # Selective log softmax and other functions selective_log_softmax_code = inspect.getsource(selective_log_softmax) + grpo_selective_log_softmax_code = inspect.getsource(grpo_selective_log_softmax) calculate_pad_tokens_in_prompt_code = inspect.getsource( calculate_pad_tokens_in_prompt ) @@ -1035,6 +1060,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): ) left_pack_padding_code = inspect.getsource(left_pack_padding) align_logprobs_with_mask_code = inspect.getsource(align_logprobs_with_mask) + autotune_batch_and_chunks_code = inspect.getsource(autotune_batch_and_chunks) # Get final source code RLTrainer_source = RLTrainer_replacement.format( RLTrainer_name = RLTrainer_name, @@ -1056,8 +1082,10 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): max_seq_length_call = max_seq_length_call, max_seq_length_post = max_seq_length_post, selective_log_softmax_code = selective_log_softmax_code, + grpo_selective_log_softmax_code = grpo_selective_log_softmax_code, calculate_pad_tokens_in_prompt_code = calculate_pad_tokens_in_prompt_code, create_completion_attention_mask_code = create_completion_attention_mask_code, + autotune_batch_and_chunks_code = autotune_batch_and_chunks_code, left_pack_padding_code = left_pack_padding_code, align_logprobs_with_mask_code = align_logprobs_with_mask_code, ) @@ -1166,6 +1194,41 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import "model = self._prepare_peft_model(model, peft_config, args)\n", "pass\n" ) + # Skip add_adapter("ref") for reference model computation + # Unsloth: We comment out the "ref" adapter creation because: + # 1. We want to use the original BASE MODEL as the reference model, not the SFT/LoRA model + # 2. PEFT doesn't allow multiple adapters when target_parameters is used (MoE models) + # When "ref" is not in peft_config, GRPO/RLOO fallback uses disable_adapter() + # which gives the base model logits - exactly what we want + add_adapter_block_pattern = ( + r"([ \t]*)" # Capture leading indentation + r"if\s+is_peft_available\(\)\s+and\s+is_peft_model\(model\)\s+and\s+args\.beta\s*!=\s*0\.0\s*:" + r"(.*?)" # Match the entire block until ref_param.data.copy_ + r"ref_param\.data\.copy_\(param\.data\)" + ) + + def comment_out_block(match): + """Comment out each line in the matched block, preserving indentation.""" + full_match = match.group(0) + indent = match.group(1) + lines = full_match.split("\n") + commented_lines = [] + # Add explanation comment first + commented_lines.append( + f"{indent}# Unsloth: Commented out - use base model as reference, not SFT/LoRA model" + ) + # Comment out each line - insert # after leading whitespace to preserve indentation + for line in lines: + if line.strip(): + stripped = line.lstrip() + leading_ws = line[: len(line) - len(stripped)] + commented_lines.append(f"{leading_ws}# {stripped}") + else: + commented_lines.append(line) + return "\n".join(commented_lines) + + init = re.sub(add_adapter_block_pattern, comment_out_block, init, flags = re.DOTALL) + # Set use_vllm if not set if "args.use_vllm" in init and "model" in init and "args" in init: # .*? matches first match. .+? matches final match. diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 5e079335ae..ff36da125d 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -50,7 +50,7 @@ RL_ADDITIONAL_FUNCTIONS = defaultdict(list) torch_compile_options = { "epilogue_fusion": True, - "max_autotune": True, + "max_autotune": False, # I saw speedups, but not sure if this has issues in collab "shape_padding": True, "trace.enabled": False, "triton.cudagraphs": False, @@ -258,18 +258,20 @@ def grpo_trainer__generate_and_score_completions(function_name, function): # The new multi-line string that will replace the line above replacement_lines = """ + max_left_pad = None batch_size = self.args.per_device_train_batch_size if mode == "train" else self.args.per_device_eval_batch_size try: # TRL 0.23.1 and below path if not has_images: # Left pad prompt before calculation old and ref hidden states - prompt_completion_ids = left_pack_padding(prompt_completion_ids, self.processing_class.pad_token_id) - self.model.for_training() + left_pad_tokens_per_prompt = calculate_pad_tokens_in_prompt(prompt_completion_ids, logits_to_keep, self.processing_class.pad_token_id) + max_left_pad = torch.max(left_pad_tokens_per_prompt).item() except: # TRL 0.24.0 and below path if images is None: # Left pad prompt before calculation old and ref hidden states - prompt_completion_ids = left_pack_padding(prompt_completion_ids, self.processing_class.pad_token_id) + left_pad_tokens_per_prompt = calculate_pad_tokens_in_prompt(prompt_completion_ids, logits_to_keep, self.processing_class.pad_token_id) + max_left_pad = torch.max(left_pad_tokens_per_prompt).item() self.model.for_training()""" function = function.replace(line_to_replace, replacement_lines) @@ -346,17 +348,45 @@ def grpo_trainer__generate_and_score_completions(function_name, function): if self.use_vllm:""" function = function.replace(replace_part, new_replacement) + # Important note: we disable TRL's importance sampling logic + # It is disabled because the LLM path moves left padding to the right. + # We must adjust the vLLM sampling_logprob tensor in Unsloth to account for this. + string_to_find = "if self.use_vllm and self.vllm_importance_sampling_correction:" + + replacement_string = ( + "if False and self.use_vllm and self.vllm_importance_sampling_correction:" + ) + + function = function.replace(string_to_find, replacement_string) + string_to_find = """ if "image_sizes" in prompt_inputs: output["image_sizes"] = prompt_inputs["image_sizes"]""" replacement_string = """ if "image_sizes" in prompt_inputs: output["image_sizes"] = prompt_inputs["image_sizes"] - - if self.use_vllm: - try: + if max_left_pad is not None: + output["max_left_pad"] = torch.tensor(prompt_ids.shape[0] * [max_left_pad]).unsqueeze(-1) + try: + if self.use_vllm and getattr(self, "vllm_importance_sampling_correction", False): output["sampling_per_token_logps"] = sampling_per_token_logps - except NameError: - output["sampling_per_token_logps"] = None""" + except NameError: + output["sampling_per_token_logps"] = None""" + + function = function.replace(string_to_find, replacement_string) + + # This path is for TRL 0.24.0 images is a variable exclusive to this version + string_to_find = """ if images is not None: + output["num_images"] = num_images""" + + replacement_string = """ if images is not None: + output["num_images"] = num_images + if max_left_pad is not None: + output["max_left_pad"] = torch.tensor(prompt_ids.shape[0] * [max_left_pad]).unsqueeze(-1) + try: + if self.use_vllm and getattr(self, "vllm_importance_sampling_correction", False): + output["sampling_per_token_logps"] = sampling_per_token_logps + except NameError: + output["sampling_per_token_logps"] = None""" function = function.replace(string_to_find, replacement_string) @@ -532,12 +562,12 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): *args, **kwargs, ): + # All Unsloth code here in this function is licensed under AGPL3 # if True: # os.environ.get('UNSLOTH_USE_NEW_MODEL', '0') == '0': # return None, None # logps, entropies Unsloth efficient GRPO if compute_efficient: return None, None else: - # Otherwise, calculate normally: if not hasattr(self, "_autocast_dtype"): self._autocast_dtype = ( torch.float16 @@ -556,47 +586,199 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): kwargs.get("image_sizes", None), ) - os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" - unwrapped_model = self.accelerator.unwrap_model( model, keep_fp32_wrapper = False ) - with torch.amp.autocast(device_type = "cuda", dtype = self._autocast_dtype): - with _get_inference_mode_context_manager(model): - if pixel_values is None: - attention_mask = input_ids != self.processing_class.pad_token_id - attention_mask = attention_mask.to(attention_mask.dtype) - # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded - logits = unwrapped_model( - input_ids = input_ids, - attention_mask = attention_mask, - pixel_values = pixel_values, - image_grid_thw = image_grid_thw, - pixel_attention_mask = pixel_attention_mask, - image_sizes = image_sizes, - # logits_to_keep = logits_to_keep + 1, - ).logits + lm_head = self.model.get_output_embeddings().weight + + dtype_bytes = ( + 16 if self._autocast_dtype in [torch.float16, torch.bfloat16] else 32 + ) + total_rows = input_ids.shape[0] + seq_len = input_ids.shape[1] + hidden_dim = lm_head.shape[1] + vocab_dim = lm_head.shape[0] + + if self.args.unsloth_grpo_mini_batch is None: + B, multiplier = autotune_batch_and_chunks( + total_rows, + seq_len, + hidden_dim, + vocab_dim, + dtype_bytes, + self.args.unsloth_logit_chunk_multiplier, + ) + B = total_rows // B + else: + B = self.args.unsloth_grpo_mini_batch + + if self.args.unsloth_logit_chunk_multiplier is None: + multiplier = max(4, seq_len // 4096) + else: + multiplier = self.args.unsloth_logit_chunk_multiplier + + all_logprobs_list = [] + if pixel_values is None: + left_pad_tokens_per_prompt = calculate_pad_tokens_in_prompt( + input_ids, logits_to_keep, self.processing_class.pad_token_id + ) + max_left_pad = torch.max(left_pad_tokens_per_prompt).item() + input_ids = left_pack_padding( + input_ids, self.processing_class.pad_token_id + ) + attention_mask = input_ids != self.processing_class.pad_token_id + attention_mask = attention_mask.to(attention_mask.dtype) + else: + max_left_pad = 0 + + # input_ids_chunks = torch.chunk(input_ids, chunks = B, dim = 0) + attention_mask_chunks = torch.chunk(attention_mask, chunks = B, dim = 0) + + def chunk_optional(tensor, chunks): + if tensor is None: + return [None] * chunks + return torch.chunk(tensor, chunks = chunks, dim = 0) + + import math + + total_samples = input_ids.shape[0] + batch_size = math.ceil(total_samples / B) + + input_ids_chunks = [] + attention_mask_chunks = [] + pixel_values_chunks = [] + image_grid_thw_chunks = [] + pixel_attention_mask_chunks = [] + + current_pixel_idx = 0 + # TRL 0.23.0 batching logic + for start in range(0, total_samples, batch_size): + end = start + batch_size + + input_ids_chunks.append(input_ids[start:end]) + attention_mask_chunks.append(attention_mask[start:end]) + + if image_grid_thw is not None and pixel_values is not None: + grid_slice = image_grid_thw[start:end] + image_grid_thw_chunks.append(grid_slice) + + batch_pixel_count = grid_slice.prod(dim = -1).sum().item() + + start_pixel_idx = current_pixel_idx + end_pixel_idx = current_pixel_idx + batch_pixel_count + + pixel_values_chunks.append( + pixel_values[start_pixel_idx:end_pixel_idx] + ) + + if pixel_attention_mask is not None: + pixel_attention_mask_chunks.append( + pixel_attention_mask[start_pixel_idx:end_pixel_idx] + ) else: - logits = unwrapped_model( - input_ids = input_ids, - attention_mask = attention_mask, - pixel_values = pixel_values, - image_grid_thw = image_grid_thw, - pixel_attention_mask = pixel_attention_mask, - image_sizes = image_sizes, - logits_to_keep = logits_to_keep + 1, - ).logits + pixel_attention_mask_chunks.append(None) + current_pixel_idx = end_pixel_idx + + else: + pixel_values_chunks.append(None) + image_grid_thw_chunks.append(None) + pixel_attention_mask_chunks.append(None) + + if image_sizes is not None and not isinstance(image_sizes, torch.Tensor): + image_sizes_chunks = [[size] for size in image_sizes] + else: + image_sizes_chunks = chunk_optional(image_sizes, B) + + temperature = self.temperature + logit_softcapping = getattr(model.config, "final_logit_softcapping", 0) + if logit_softcapping is None: + logit_softcapping = 0 + logit_scale_multiply = getattr(model.config, "logit_scale", 0) + if logit_scale_multiply is None: + logit_scale_multiply = 0 + logit_scale_divide = getattr(model.config, "logits_scaling", 0) + if logit_scale_divide is None: + logit_scale_divide = 0 + + zipped_inputs = zip( + input_ids_chunks, + attention_mask_chunks, + pixel_values_chunks, + image_grid_thw_chunks, + pixel_attention_mask_chunks, + image_sizes_chunks, + ) + os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" + + with _get_inference_mode_context_manager(model): + for ( + input_ids_chunk, + attention_mask_chunk, + pixel_values_chunk, + image_grid_thw_chunk, + pixel_attention_mask_chunk, + image_sizes_chunk, + ) in zipped_inputs: + with torch.amp.autocast( + device_type = "cuda", dtype = self._autocast_dtype + ): + if pixel_values is None: + logits_chunk = unwrapped_model( + input_ids = input_ids_chunk, + attention_mask = attention_mask_chunk, + pixel_values = pixel_values_chunk, + image_grid_thw = image_grid_thw_chunk, + pixel_attention_mask = pixel_attention_mask_chunk, + image_sizes = image_sizes_chunk, + ).logits + + completion_input_ids_chunk = input_ids_chunk[ + :, -(logits_to_keep + max_left_pad) : + ] + logits_chunk = logits_chunk[ + :, -(logits_to_keep + max_left_pad + 1) :, : + ] + logits_chunk = logits_chunk[:, :-1, :] + else: + # Essentially, for VLMs we do not go via the optimized path in models/, + # so we don't encounter the Flash Attn left-padding issue. + logits_chunk = unwrapped_model( + input_ids = input_ids_chunk, + attention_mask = attention_mask_chunk, + pixel_values = pixel_values_chunk, + image_grid_thw = image_grid_thw_chunk, + pixel_attention_mask = pixel_attention_mask_chunk, + image_sizes = image_sizes_chunk, + logits_to_keep = logits_to_keep + 1, + ).logits + + logits_chunk = logits_chunk[:, :-1, :] + completion_input_ids_chunk = input_ids_chunk[ + :, -logits_to_keep: + ] + + logprobs_chunk = chunked_hidden_states_selective_log_softmax( + logits_chunk, + lm_head, + completion_input_ids_chunk, + chunks = input_ids_chunk.shape[0] * multiplier, + logit_scale_multiply = logit_scale_multiply, + logit_scale_divide = logit_scale_divide, + logit_softcapping = logit_softcapping, + temperature = temperature, + ) + # This is needed to avoid race conditions with GPT OSS offload_embbed=True + # However, it seems that this line does not slow down or disrupt models. + torch.cuda.synchronize() + all_logprobs_list.append(logprobs_chunk) + logprobs = torch.cat(all_logprobs_list, dim = 0) entropies = None - if compute_entropy: - from trl.trainer.utils import entropy_from_logits - - entropies = entropy_from_logits(logits) os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "0" - # logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred - return logits.detach(), entropies # logps, entropies + + return logprobs.detach(), entropies # logps, entropies # input_ids = input_ids[:, -logits_to_keep:] # For transformers<=4.48, logits_to_keep argument isn't supported, so here we drop logits ourselves. # See https://github.com/huggingface/trl/issues/2770 @@ -708,14 +890,14 @@ def grpo_trainer_compute_loss(function_name, function): # ref_per_token_logps = per_token_logps = get_logps_func(model, input_ids, attention_mask, logits_to_keep) # else: # ref_per_token_logps = None - ref_hidden_states = inputs.get("ref_per_token_logps", None) + ref_logps = inputs.get("ref_per_token_logps", None) # per_token_kl = torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1 # x - x.detach() allows for preserving gradients from x advantages = inputs["advantages"] # per_token_loss = torch.exp(per_token_logps - per_token_logps.detach()) * advantages.unsqueeze(1) # per_token_loss = -(per_token_loss - self.beta * per_token_kl) # loss = ((per_token_loss * completion_mask).sum(dim=1) / completion_mask.sum(dim=1)).mean() - old_hidden_states = inputs.get("old_per_token_logps", None) + old_logps = inputs.get("old_per_token_logps", None) input_ids = input_ids[:, -logits_to_keep:] @@ -730,24 +912,13 @@ def grpo_trainer_compute_loss(function_name, function): if logit_scale_divide is None: logit_scale_divide = 0 + max_left_pad = inputs.get("max_left_pad", 0) if per_token_logps is not None: - if ref_hidden_states is not None: - ref_hidden_states = ref_hidden_states[ - :, :-1, : - ] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred - if old_hidden_states is not None: - old_hidden_states = old_hidden_states[ - :, :-1, : - ] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred - per_token_logps = per_token_logps[ - :, :-1, : - ] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred - loss, completion_length, mean_kl, delta, flat_is_ratio = ( grpo_compute_loss_slow( - ref_hidden_states, + ref_logps, per_token_logps, - old_hidden_states, + old_logps, input_ids, completion_mask, self.beta, @@ -761,6 +932,7 @@ def grpo_trainer_compute_loss(function_name, function): max_completion_length = self.args.max_completion_length, delta = self.args.delta, temperature = self.args.temperature, + max_left_pad = max_left_pad, logit_softcapping = logit_softcapping, logit_scale_multiply = logit_scale_multiply, logit_scale_divide = logit_scale_divide, @@ -781,8 +953,8 @@ def grpo_trainer_compute_loss(function_name, function): logits_to_keep = logits_to_keep, completion_mask = completion_mask, advantages = advantages, - old_hidden_states = old_hidden_states, - ref_hidden_states = ref_hidden_states, + old_logps = old_logps, + ref_logps = ref_logps, n_chunks = self.args.unsloth_num_chunks, loss_type = self.args.loss_type, importance_sampling_level = self.importance_sampling_level, @@ -791,6 +963,7 @@ def grpo_trainer_compute_loss(function_name, function): max_completion_length = self.args.max_completion_length, delta = self.args.delta, temperature = self.args.temperature, + max_left_pad = max_left_pad, logit_softcapping = logit_softcapping, logit_scale_multiply = logit_scale_multiply, logit_scale_divide = logit_scale_divide, @@ -809,8 +982,8 @@ def grpo_trainer_compute_loss(function_name, function): logits_to_keep = logits_to_keep, completion_mask = completion_mask, advantages = advantages, - old_hidden_states = old_hidden_states, - ref_hidden_states = ref_hidden_states, + old_logps = old_logps, + ref_logps = ref_logps, n_chunks = self.args.unsloth_num_chunks, temperature = self.args.temperature, logit_softcapping = logit_softcapping, @@ -827,7 +1000,11 @@ def grpo_trainer_compute_loss(function_name, function): self._metrics["completion_length"].append(completion_length.item()) self._metrics["kl"].append(mean_kl.item()) - if self.use_vllm and delta is not None: + if ( + self.use_vllm + and delta is not None + and getattr(self, "vllm_importance_sampling_correction", False) + ): mean_delta = ( torch.mean(delta) if delta.numel() > 0 diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index 6de942d7d2..4e03e0a168 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -1273,7 +1273,7 @@ class FastBaseModel: # Since transformers 4.53, must turn on explicitly for module in model.modules(): if hasattr(module, "gradient_checkpointing"): - module.gradient_checkpointing = True + module.gradient_checkpointing = use_gradient_checkpointing # Also re-enable training for embeddings for NEFTune if hasattr(model, "get_input_embeddings"):