Merge branch 'unslothai:main' into FST

This commit is contained in:
electroglyph 2026-01-19 00:56:21 -08:00 committed by GitHub
commit 435ecdb895
19 changed files with 1135 additions and 207 deletions

View file

@ -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/

View file

@ -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:

View file

@ -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)

View file

@ -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",

172
tests/test_raw_text.py Normal file
View file

@ -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 = "</s>"
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)

View file

@ -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)

View file

@ -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)

View file

@ -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,

View file

@ -13,3 +13,4 @@
# limitations under the License.
from .synthetic import *
from .raw_text import *

View file

@ -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

View file

@ -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.")

View file

@ -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(

View file

@ -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):

View file

@ -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}"

View file

@ -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:

View file

@ -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

View file

@ -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.

View file

@ -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

View file

@ -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"):