From eb4dd11c493a1a7ecf366395ce29d6aa8a786ed1 Mon Sep 17 00:00:00 2001 From: vangmay Date: Tue, 18 Nov 2025 21:46:41 +0800 Subject: [PATCH 01/44] Write file and template for raw_text dataprep --- unsloth/dataprep/raw_text.py | 76 ++++++++++++++++++++++++++++++++++++ 1 file changed, 76 insertions(+) create mode 100644 unsloth/dataprep/raw_text.py diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py new file mode 100644 index 0000000000..93c8c6855b --- /dev/null +++ b/unsloth/dataprep/raw_text.py @@ -0,0 +1,76 @@ +# 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 + +class RawTextDataLoader: + def __init__(self, tokenizer, chunk_size=2048, stride=512): + self.tokenizer = tokenizer + self.chunk_size = chunk_size + self.stride = stride + + def load_from_file(self, file_path): + """Load raw text and convert to dataset""" + + def load_from_files(self, file_paths): + """Load multiple text files""" + + def chunk_text(self, text): + """Split text into overlapping chunks""" + + def create_causal_dataset(self, chunks): + """Create dataset for causal language modeling""" + + def smart_chunk_text(self, text, chunk_size, stride): + """ + Intelligent chunking that: + 1. Respects sentence/paragraph boundaries + 2. Handles various text formats (.txt, .md, .json, etc.) + 3. Maintains context with stride overlap + 4. Adds proper EOS tokens + """ + + def tokenize_and_chunk(self, text): + """ + Tokenize first, then chunk by token count: + 1. More precise length control + 2. Avoids mid-token splits + 3. Handles different languages better + """ + +class TextPreprocessor: + def clean_text(self, text): + """Remove unwanted characters, normalize whitespace""" + + def extract_sections(self, text, patterns): + """Extract specific sections (e.g., code blocks, quotes)""" + + def add_structure_tokens(self, text): + """Add special tokens for structure (chapters, sections)""" + +def validate_dataset(self, dataset): + """ + Check for: + - Minimum/maximum sequence lengths + - Character encoding issues + - Repeated content + - Empty chunks + """ + From 14985c4011519dfc9f5827d8d2a837d891abed0e Mon Sep 17 00:00:00 2001 From: vangmay Date: Tue, 18 Nov 2025 21:53:20 +0800 Subject: [PATCH 02/44] Add implementation to cli --- unsloth-cli.py | 48 ++++++++++++++++++++++++++++++++++++ unsloth/dataprep/raw_text.py | 25 +++++++++++++++++++ 2 files changed, 73 insertions(+) diff --git a/unsloth-cli.py b/unsloth-cli.py index fb6e392662..aac0e7f7e1 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -42,6 +42,7 @@ def run(args): from transformers import TrainingArguments from unsloth import is_bfloat16_supported import logging + from unsloth import RawTextDataLoader logging.getLogger("hf-to-gguf").setLevel(logging.WARNING) @@ -98,6 +99,21 @@ def run(args): texts.append(text) return {"text": texts} + def load_dataset_smart(args): + 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: + # Existing HuggingFace dataset logic + dataset = load_dataset(args.dataset, split="train") + dataset = dataset.map(formatting_prompts_func, batched=True) + return dataset + use_modelscope = strtobool(os.environ.get("UNSLOTH_USE_MODELSCOPE", "False")) if use_modelscope: from modelscope import MsDataset @@ -389,5 +405,37 @@ if __name__ == "__main__": "--hub_token", type = str, 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" + ) + + TRAINING_MODES = { + 'instruction': 'Standard instruction-following', + 'causal': 'Causal language modeling (raw text)', + 'completion': 'Text completion tasks' + } + + parser.add_argument( + "--training_mode", + type=str, + default="instruction", + choices=list(TRAINING_MODES.keys()), + help="Training mode for the model" + ) + args = parser.parse_args() run(args) diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index 93c8c6855b..e36978bf1c 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -20,12 +20,28 @@ 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): self.tokenizer = tokenizer self.chunk_size = chunk_size self.stride = stride + def detect_format(self, file_path): + """Auto-detect file format and parse accordingly""" + def load_from_file(self, file_path): """Load raw text and convert to dataset""" @@ -64,6 +80,15 @@ class TextPreprocessor: def add_structure_tokens(self, text): """Add special tokens for structure (chapters, sections)""" + + def validate_dataset(self, dataset): + """ + Check for: + - Minimum/maximum sequence lengths + - Character encoding issues + - Repeated content + - Empty chunks + """ def validate_dataset(self, dataset): """ From 84b12be161080315251700d9bab4f5587efa1558 Mon Sep 17 00:00:00 2001 From: vangmay Date: Tue, 18 Nov 2025 21:59:01 +0800 Subject: [PATCH 03/44] Add support for multiple files --- unsloth/dataprep/raw_text.py | 56 ++++++++++++++++++++++++++++++++++++ 1 file changed, 56 insertions(+) diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index e36978bf1c..8b705b3a96 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -41,12 +41,26 @@ class RawTextDataLoader: 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): """Load raw text and convert to dataset""" + 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 self.create_causal_dataset(chunks) def load_from_files(self, file_paths): """Load multiple text files""" + 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) + all_chunks.extend(chunks) + return self.create_causal_dataset(all_chunks) + def chunk_text(self, text): """Split text into overlapping chunks""" @@ -71,6 +85,48 @@ class RawTextDataLoader: 3. Handles different languages better """ + 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""" From d4060b627a68b5c82ed825899be826f0679534dc Mon Sep 17 00:00:00 2001 From: vangmay Date: Tue, 18 Nov 2025 22:00:07 +0800 Subject: [PATCH 04/44] Write chunking logic --- unsloth/dataprep/raw_text.py | 47 ++++++++++++++++++++++++++++++++++++ 1 file changed, 47 insertions(+) diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index 8b705b3a96..44582c76a6 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -64,9 +64,11 @@ class RawTextDataLoader: def chunk_text(self, text): """Split text into overlapping chunks""" + return self.smart_chunk_text(text, self.chunk_size, self.stride) def create_causal_dataset(self, chunks): """Create dataset for causal language modeling""" + return Dataset.from_dict({"text": chunks}) def smart_chunk_text(self, text, chunk_size, stride): """ @@ -76,6 +78,51 @@ class RawTextDataLoader: 3. Maintains context with stride overlap 4. Adds proper EOS tokens """ + # 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 + 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] + + # Decode back to text + 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 tokenize_and_chunk(self, text): """ From d404ff653c2d24e68d3c23f8b3ee5e30eca33e42 Mon Sep 17 00:00:00 2001 From: vangmay Date: Tue, 18 Nov 2025 22:01:36 +0800 Subject: [PATCH 05/44] Add logic to clean and extract text sections --- unsloth/dataprep/raw_text.py | 17 ++++++++++++++++- 1 file changed, 16 insertions(+), 1 deletion(-) diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index 44582c76a6..b31120cc59 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -177,12 +177,27 @@ class RawTextDataLoader: 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): """ From afdbeff8e957e1f61dd6a5fd9617b4f9c42d623c Mon Sep 17 00:00:00 2001 From: vangmay Date: Tue, 18 Nov 2025 22:02:35 +0800 Subject: [PATCH 06/44] Add validation code --- unsloth/dataprep/raw_text.py | 67 +++++++++++++++++++++++++++++++----- 1 file changed, 58 insertions(+), 9 deletions(-) diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index b31120cc59..dd6d6b96d0 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -207,13 +207,62 @@ class TextPreprocessor: - Repeated content - Empty chunks """ - -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 From e473c0dbdf968bdf22065ad4f5ba0bdc8a6caff4 Mon Sep 17 00:00:00 2001 From: vangmay Date: Tue, 18 Nov 2025 22:36:38 +0800 Subject: [PATCH 07/44] Write simple test --- tests/test_raw_text.py | 121 +++++++++++++++++++++++++++++++++++++++++ 1 file changed, 121 insertions(+) create mode 100644 tests/test_raw_text.py diff --git a/tests/test_raw_text.py b/tests/test_raw_text.py new file mode 100644 index 0000000000..503fbeb4c3 --- /dev/null +++ b/tests/test_raw_text.py @@ -0,0 +1,121 @@ +#!/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 = "" + + 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) + 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 + dataset = loader.load_from_file(test_file) + assert len(dataset) > 0, "Should create at least one chunk" + assert 'text' in dataset.column_names, "Dataset should have 'text' column" + + # 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(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) \ No newline at end of file From d738a087f928611a60ed3098f95a185c187266f7 Mon Sep 17 00:00:00 2001 From: vangmay Date: Tue, 18 Nov 2025 22:44:48 +0800 Subject: [PATCH 08/44] Add module to init --- unsloth/dataprep/__init__.py | 1 + 1 file changed, 1 insertion(+) diff --git a/unsloth/dataprep/__init__.py b/unsloth/dataprep/__init__.py index b36122eb74..b6840f247f 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 * From 69e89677423f2ad062728d47f9aa318946bb6558 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 20 Nov 2025 12:51:17 +0000 Subject: [PATCH 09/44] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/test_raw_text.py | 64 ++++++------ unsloth-cli.py | 34 +++---- unsloth/dataprep/raw_text.py | 183 +++++++++++++++++++---------------- 3 files changed, 146 insertions(+), 135 deletions(-) diff --git a/tests/test_raw_text.py b/tests/test_raw_text.py index 503fbeb4c3..9bbfee92ac 100644 --- a/tests/test_raw_text.py +++ b/tests/test_raw_text.py @@ -10,15 +10,16 @@ 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'] @@ -28,19 +29,22 @@ class MockDataset: 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 = type(sys)("datasets") datasets_mock.Dataset = MockDataset -sys.modules['datasets'] = datasets_mock +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') +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) @@ -49,73 +53,75 @@ 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 = "" - - def __call__(self, text, return_tensors=None, add_special_tokens=False): + + 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) + return {"input_ids": [MockTensor(token_ids)]} return {"input_ids": token_ids} - - def decode(self, token_ids, skip_special_tokens=False): + + 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: + 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) - + loader = RawTextDataLoader(tokenizer, chunk_size = 5, stride = 2) + # Test loading dataset = loader.load_from_file(test_file) assert len(dataset) > 0, "Should create at least one chunk" - assert 'text' in dataset.column_names, "Dataset should have 'text' column" - + assert "text" in dataset.column_names, "Dataset should have 'text' column" + # 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(dataset) - assert stats['total_samples'] > 0, "Should count samples" - assert 'warnings' in stats, "Should include warnings" - + 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) \ No newline at end of file + sys.exit(0 if success else 1) diff --git a/unsloth-cli.py b/unsloth-cli.py index aac0e7f7e1..e454a47c44 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -104,14 +104,14 @@ def run(args): # 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')): + 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: # Existing HuggingFace dataset logic - dataset = load_dataset(args.dataset, split="train") - dataset = dataset.map(formatting_prompts_func, batched=True) + dataset = load_dataset(args.dataset, split = "train") + dataset = dataset.map(formatting_prompts_func, batched = True) return dataset use_modelscope = strtobool(os.environ.get("UNSLOTH_USE_MODELSCOPE", "False")) @@ -406,35 +406,27 @@ if __name__ == "__main__": ) parser.add_argument( - "--raw_text_file", - type=str, - help="Path to raw text file for training" + "--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" + "--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" + "--stride", type = int, default = 512, help = "Overlap between chunks" ) TRAINING_MODES = { - 'instruction': 'Standard instruction-following', - 'causal': 'Causal language modeling (raw text)', - 'completion': 'Text completion tasks' + "instruction": "Standard instruction-following", + "causal": "Causal language modeling (raw text)", + "completion": "Text completion tasks", } parser.add_argument( "--training_mode", - type=str, - default="instruction", - choices=list(TRAINING_MODES.keys()), - help="Training mode for the model" + type = str, + default = "instruction", + choices = list(TRAINING_MODES.keys()), + help = "Training mode for the model", ) args = parser.parse_args() diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index dd6d6b96d0..d809eb312a 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -26,23 +26,24 @@ __all__ = [ ] SUPPORTED_FORMATS = { - '.txt': 'plain_text', - '.md': 'markdown', - '.json': 'json_lines', - '.jsonl': 'json_lines', - '.csv': 'csv_text_column' + ".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): + def __init__(self, tokenizer, chunk_size = 2048, stride = 512): self.tokenizer = tokenizer - self.chunk_size = chunk_size + self.chunk_size = chunk_size self.stride = stride - + 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') + return SUPPORTED_FORMATS.get(extension, "plain_text") def load_from_file(self, file_path): """Load raw text and convert to dataset""" @@ -50,7 +51,7 @@ class RawTextDataLoader: text_content = self._read_file_by_format(file_path, file_format) chunks = self.smart_chunk_text(text_content, self.chunk_size, self.stride) return self.create_causal_dataset(chunks) - + def load_from_files(self, file_paths): """Load multiple text files""" all_chunks = [] @@ -61,11 +62,10 @@ class RawTextDataLoader: all_chunks.extend(chunks) return self.create_causal_dataset(all_chunks) - def chunk_text(self, text): """Split text into overlapping chunks""" return self.smart_chunk_text(text, self.chunk_size, self.stride) - + def create_causal_dataset(self, chunks): """Create dataset for causal language modeling""" return Dataset.from_dict({"text": chunks}) @@ -79,51 +79,50 @@ class RawTextDataLoader: 4. Adds proper EOS tokens """ # First pass: tokenize the entire text to get accurate token counts - tokenized = self.tokenizer(text, return_tensors="pt", add_special_tokens=False) + 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 hasattr(tokens, "__len__") and len(tokens) > 0: # If it's a nested structure, get the first element - if hasattr(tokens[0], '__len__'): + 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 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] - + # Decode back to text - chunk_text = self.tokenizer.decode(chunk_tokens, skip_special_tokens=True) - + 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 tokenize_and_chunk(self, text): """ Tokenize first, then chunk by token count: @@ -134,10 +133,10 @@ class RawTextDataLoader: 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': + 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': + elif file_format == "json_lines": lines = [] for line in f: try: @@ -147,42 +146,43 @@ class RawTextDataLoader: lines.append(text) except json.JSONDecodeError: continue - return '\n\n'.join(lines) - elif file_format == 'csv_text_column': + 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 "\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'] + 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'] + 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) + 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 = [] @@ -193,12 +193,20 @@ class TextPreprocessor: 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) + 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: @@ -208,61 +216,66 @@ class TextPreprocessor: - 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': [] + "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'] + + 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 + stats["empty_samples"] += 1 continue - + # Check for encoding issues try: - text.encode('utf-8') + text.encode("utf-8") except UnicodeEncodeError: - stats['encoding_issues'] += 1 - + 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) - + 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 + 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 + 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 From f80fe573e6a44e1337cfd2bae7022fba53437d44 Mon Sep 17 00:00:00 2001 From: vangmay Date: Thu, 20 Nov 2025 20:53:22 +0800 Subject: [PATCH 10/44] Integrate smart dataset loader --- unsloth-cli.py | 25 ++++++++++++++----------- 1 file changed, 14 insertions(+), 11 deletions(-) diff --git a/unsloth-cli.py b/unsloth-cli.py index aac0e7f7e1..044ad93f31 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -100,6 +100,8 @@ def run(args): return {"text": texts} def load_dataset_smart(args): + from transformers.utils import strtobool + if args.raw_text_file: # Use raw text loader loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) @@ -109,20 +111,21 @@ def run(args): loader = RawTextDataLoader(tokenizer) dataset = loader.load_from_file(args.dataset) else: - # Existing HuggingFace dataset logic - dataset = load_dataset(args.dataset, split="train") + # 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 - 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: - # Load and format dataset - dataset = load_dataset(args.dataset, split = "train") - dataset = dataset.map(formatting_prompts_func, batched = True) + # Load dataset using smart loader + dataset = load_dataset_smart(args) print("Data is formatted and ready!") # Configure training arguments From 12474b14d9234f42815ad47c758f979d6680b491 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 20 Nov 2025 12:57:53 +0000 Subject: [PATCH 11/44] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth-cli.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/unsloth-cli.py b/unsloth-cli.py index 44232e19d7..79b3ef5296 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -101,7 +101,7 @@ def run(args): def load_dataset_smart(args): from transformers.utils import strtobool - + if args.raw_text_file: # Use raw text loader loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) @@ -112,16 +112,19 @@ def run(args): dataset = loader.load_from_file(args.dataset) else: # Check for modelscope usage - use_modelscope = strtobool(os.environ.get("UNSLOTH_USE_MODELSCOPE", "False")) + use_modelscope = strtobool( + os.environ.get("UNSLOTH_USE_MODELSCOPE", "False") + ) if use_modelscope: from modelscope import MsDataset - dataset = MsDataset.load(args.dataset, split="train") + + dataset = MsDataset.load(args.dataset, split = "train") else: # Existing HuggingFace dataset logic - dataset = load_dataset(args.dataset, split="train") - + dataset = load_dataset(args.dataset, split = "train") + # Apply formatting for structured datasets - dataset = dataset.map(formatting_prompts_func, batched=True) + dataset = dataset.map(formatting_prompts_func, batched = True) return dataset # Load dataset using smart loader From e57b1c73fde87f382f3123d02a84a164f934d6bb Mon Sep 17 00:00:00 2001 From: vangmay Date: Thu, 20 Nov 2025 21:08:33 +0800 Subject: [PATCH 12/44] Make the chunk function efficient --- tests/test_raw_text.py | 24 ++++++++-- unsloth-cli.py | 20 ++++++--- unsloth/dataprep/raw_text.py | 85 ++++++++++++++++++++++++++++-------- 3 files changed, 99 insertions(+), 30 deletions(-) diff --git a/tests/test_raw_text.py b/tests/test_raw_text.py index 9bbfee92ac..88dac2604d 100644 --- a/tests/test_raw_text.py +++ b/tests/test_raw_text.py @@ -61,6 +61,7 @@ def test_raw_text_loader(): 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() @@ -77,6 +78,9 @@ def test_raw_text_loader(): def __len__(self): return len(self.data) + + def tolist(self): + return self.data return {"input_ids": [MockTensor(token_ids)]} return {"input_ids": token_ids} @@ -95,10 +99,22 @@ def test_raw_text_loader(): tokenizer = MockTokenizer() loader = RawTextDataLoader(tokenizer, chunk_size = 5, stride = 2) - # Test loading - dataset = loader.load_from_file(test_file) - assert len(dataset) > 0, "Should create at least one chunk" - assert "text" in dataset.column_names, "Dataset should have 'text' column" + # Test loading with text output (legacy mode) + text_dataset = loader.load_from_file(test_file, return_tensors=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_tensors=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" # Test preprocessor preprocessor = TextPreprocessor() diff --git a/unsloth-cli.py b/unsloth-cli.py index 79b3ef5296..efdc0b3f4d 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -103,13 +103,19 @@ def run(args): from transformers.utils import strtobool if args.raw_text_file: - # Use raw text loader + # Use raw text loader - returns pre-tokenized data loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) - dataset = loader.load_from_file(args.raw_text_file) + dataset = loader.load_from_file(args.raw_text_file, return_tensors=True) + # Mark dataset as pre-tokenized to skip text formatting + dataset._is_pretokenized = True + return dataset elif args.dataset.endswith((".txt", ".md", ".json", ".jsonl")): - # Auto-detect local raw text files - loader = RawTextDataLoader(tokenizer) - dataset = loader.load_from_file(args.dataset) + # Auto-detect local raw text files - returns pre-tokenized data + loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) + dataset = loader.load_from_file(args.dataset, return_tensors=True) + # Mark dataset as pre-tokenized to skip text formatting + dataset._is_pretokenized = True + return dataset else: # Check for modelscope usage use_modelscope = strtobool( @@ -123,9 +129,9 @@ def run(args): # Existing HuggingFace dataset logic dataset = load_dataset(args.dataset, split = "train") - # Apply formatting for structured datasets + # Apply formatting for structured datasets (text-based) dataset = dataset.map(formatting_prompts_func, batched = True) - return dataset + return dataset # Load dataset using smart loader dataset = load_dataset_smart(args) diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index d809eb312a..50612c1b86 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -35,48 +35,66 @@ SUPPORTED_FORMATS = { class RawTextDataLoader: - def __init__(self, tokenizer, chunk_size = 2048, stride = 512): + def __init__(self, tokenizer, chunk_size = 2048, stride = 512, return_tokenized = True): 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): + 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) - chunks = self.smart_chunk_text(text_content, self.chunk_size, self.stride) + 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): + 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) + 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): + def chunk_text(self, text, return_tokenized=None): """Split text into overlapping chunks""" - return self.smart_chunk_text(text, self.chunk_size, self.stride) + 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""" - return Dataset.from_dict({"text": chunks}) + 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] + return Dataset.from_dict({ + "input_ids": input_ids, + "attention_mask": attention_mask + }) + else: + # If chunks are text strings (backward compatibility) + return Dataset.from_dict({"text": chunks}) - def smart_chunk_text(self, text, chunk_size, stride): + 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. Adds proper EOS tokens + 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) @@ -93,8 +111,19 @@ class RawTextDataLoader: if len(tokens) <= chunk_size: # Text is small enough to fit in one chunk - eos_token = self.tokenizer.eos_token if self.tokenizer.eos_token else "" - return [text + eos_token] + 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 @@ -106,15 +135,33 @@ class RawTextDataLoader: # Extract tokens for this chunk chunk_tokens = tokens[start_idx:end_idx] - # Decode back to text - chunk_text = self.tokenizer.decode(chunk_tokens, skip_special_tokens = True) + 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) - # 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 + # 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) - chunks.append(chunk_text) + # 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): From 2ee16e7ee25be4313c2ba573e93fcdf889502376 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 20 Nov 2025 13:09:07 +0000 Subject: [PATCH 13/44] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/test_raw_text.py | 24 +++++++++----- unsloth-cli.py | 4 +-- unsloth/dataprep/raw_text.py | 62 ++++++++++++++++++++++-------------- 3 files changed, 56 insertions(+), 34 deletions(-) diff --git a/tests/test_raw_text.py b/tests/test_raw_text.py index 88dac2604d..5306c68fa6 100644 --- a/tests/test_raw_text.py +++ b/tests/test_raw_text.py @@ -78,7 +78,7 @@ def test_raw_text_loader(): def __len__(self): return len(self.data) - + def tolist(self): return self.data @@ -100,21 +100,29 @@ def test_raw_text_loader(): loader = RawTextDataLoader(tokenizer, chunk_size = 5, stride = 2) # Test loading with text output (legacy mode) - text_dataset = loader.load_from_file(test_file, return_tensors=False) + text_dataset = loader.load_from_file(test_file, return_tensors = 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_tensors=True) + tokenized_dataset = loader.load_from_file(test_file, return_tensors = 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" - + 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" + 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" # Test preprocessor preprocessor = TextPreprocessor() diff --git a/unsloth-cli.py b/unsloth-cli.py index efdc0b3f4d..5135d00493 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -105,14 +105,14 @@ def run(args): if args.raw_text_file: # Use raw text loader - returns pre-tokenized data loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) - dataset = loader.load_from_file(args.raw_text_file, return_tensors=True) + dataset = loader.load_from_file(args.raw_text_file, return_tensors = True) # Mark dataset as pre-tokenized to skip text formatting dataset._is_pretokenized = True return dataset elif args.dataset.endswith((".txt", ".md", ".json", ".jsonl")): # Auto-detect local raw text files - returns pre-tokenized data loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) - dataset = loader.load_from_file(args.dataset, return_tensors=True) + dataset = loader.load_from_file(args.dataset, return_tensors = True) # Mark dataset as pre-tokenized to skip text formatting dataset._is_pretokenized = True return dataset diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index 50612c1b86..d6a9a6a047 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -46,16 +46,18 @@ class RawTextDataLoader: extension = Path(file_path).suffix.lower() return SUPPORTED_FORMATS.get(extension, "plain_text") - def load_from_file(self, file_path, return_tokenized=None): + 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) - chunks = self.smart_chunk_text(text_content, self.chunk_size, self.stride, return_tokenized) + 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): + def load_from_files(self, file_paths, return_tokenized = None): """Load multiple text files""" if return_tokenized is None: return_tokenized = self.return_tokenized @@ -63,15 +65,19 @@ class RawTextDataLoader: 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) + 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): + 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) + 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""" @@ -80,15 +86,14 @@ class RawTextDataLoader: # 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] - return Dataset.from_dict({ - "input_ids": input_ids, - "attention_mask": attention_mask - }) + return Dataset.from_dict( + {"input_ids": input_ids, "attention_mask": attention_mask} + ) 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): + def smart_chunk_text(self, text, chunk_size, stride, return_tokenized = True): """ Intelligent chunking that: 1. Respects sentence/paragraph boundaries @@ -113,11 +118,13 @@ class RawTextDataLoader: # 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) + 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 = ( + 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}] @@ -137,28 +144,35 @@ class RawTextDataLoader: 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) - + 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) + 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 - }) + + 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) + 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 "" + eos_token = ( + self.tokenizer.eos_token if self.tokenizer.eos_token else "" + ) chunk_text += eos_token chunks.append(chunk_text) From 4596b67dc9c12745f6342e5a262b6eab321d4a90 Mon Sep 17 00:00:00 2001 From: vangmay Date: Thu, 20 Nov 2025 21:40:45 +0800 Subject: [PATCH 14/44] remove old function --- unsloth/dataprep/raw_text.py | 8 -------- 1 file changed, 8 deletions(-) diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index d6a9a6a047..b880e97338 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -184,14 +184,6 @@ class RawTextDataLoader: return chunks - def tokenize_and_chunk(self, text): - """ - Tokenize first, then chunk by token count: - 1. More precise length control - 2. Avoids mid-token splits - 3. Handles different languages better - """ - 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: From 979b506069e43495b3d61a65820f41a3b433c2ea Mon Sep 17 00:00:00 2001 From: vangmay Date: Tue, 25 Nov 2025 21:01:43 +0800 Subject: [PATCH 15/44] Remove training mode arg --- unsloth-cli.py | 34 +++++++--------------------------- 1 file changed, 7 insertions(+), 27 deletions(-) diff --git a/unsloth-cli.py b/unsloth-cli.py index 5135d00493..ef8167ad39 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -103,19 +103,13 @@ def run(args): from transformers.utils import strtobool if args.raw_text_file: - # Use raw text loader - returns pre-tokenized data + # Use raw text loader loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) - dataset = loader.load_from_file(args.raw_text_file, return_tensors = True) - # Mark dataset as pre-tokenized to skip text formatting - dataset._is_pretokenized = True - return dataset + dataset = loader.load_from_file(args.raw_text_file) elif args.dataset.endswith((".txt", ".md", ".json", ".jsonl")): - # Auto-detect local raw text files - returns pre-tokenized data - loader = RawTextDataLoader(tokenizer, args.chunk_size, args.stride) - dataset = loader.load_from_file(args.dataset, return_tensors = True) - # Mark dataset as pre-tokenized to skip text formatting - dataset._is_pretokenized = True - return dataset + # Auto-detect local raw text files + loader = RawTextDataLoader(tokenizer) + dataset = loader.load_from_file(args.dataset) else: # Check for modelscope usage use_modelscope = strtobool( @@ -129,9 +123,9 @@ def run(args): # Existing HuggingFace dataset logic dataset = load_dataset(args.dataset, split = "train") - # Apply formatting for structured datasets (text-based) + # Apply formatting for structured datasets dataset = dataset.map(formatting_prompts_func, batched = True) - return dataset + return dataset # Load dataset using smart loader dataset = load_dataset_smart(args) @@ -427,19 +421,5 @@ if __name__ == "__main__": "--stride", type = int, default = 512, help = "Overlap between chunks" ) - TRAINING_MODES = { - "instruction": "Standard instruction-following", - "causal": "Causal language modeling (raw text)", - "completion": "Text completion tasks", - } - - parser.add_argument( - "--training_mode", - type = str, - default = "instruction", - choices = list(TRAINING_MODES.keys()), - help = "Training mode for the model", - ) - args = parser.parse_args() run(args) From aa36dabd81dcb6abac9620eaaee8e9bc672ce81c Mon Sep 17 00:00:00 2001 From: vangmay Date: Wed, 10 Dec 2025 10:15:56 +0530 Subject: [PATCH 16/44] Fix RawTextDataLoader import issue --- unsloth/__init__.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/unsloth/__init__.py b/unsloth/__init__.py index 8b48ce3ba0..b77af7c242 100644 --- a/unsloth/__init__.py +++ b/unsloth/__init__.py @@ -247,6 +247,8 @@ 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, From 87b924b28ea0190195f2617ec7c500694cad1255 Mon Sep 17 00:00:00 2001 From: vangmay Date: Wed, 10 Dec 2025 10:17:23 +0530 Subject: [PATCH 17/44] Fix Incorrect non-relative import in dataprep package --- unsloth/dataprep/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/unsloth/dataprep/__init__.py b/unsloth/dataprep/__init__.py index b6840f247f..048f9b8010 100644 --- a/unsloth/dataprep/__init__.py +++ b/unsloth/dataprep/__init__.py @@ -13,4 +13,4 @@ # limitations under the License. from .synthetic import * -from raw_text import * +from .raw_text import * From 3fb6335f0259fa846330f0a25262bca4a6943422 Mon Sep 17 00:00:00 2001 From: vangmay Date: Wed, 10 Dec 2025 10:46:29 +0530 Subject: [PATCH 18/44] =?UTF-8?q?Fix=20Chunking=20loop=20can=20hang=20when?= =?UTF-8?q?=20stride=20=E2=89=A5=20chunk=5Fsize?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- unsloth/dataprep/raw_text.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index b880e97338..f7c1bf7856 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -101,6 +101,13 @@ class RawTextDataLoader: 3. Maintains context with stride overlap 4. Returns tokenized chunks directly (more efficient) or text chunks """ + 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}) to progress the chunking loop" + ) + # 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"] From c1086e3ed63a2519b3181b01985a3b02b5659a32 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 10 Dec 2025 05:17:01 +0000 Subject: [PATCH 19/44] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/__init__.py | 1 + 1 file changed, 1 insertion(+) diff --git a/unsloth/__init__.py b/unsloth/__init__.py index b77af7c242..26e43eec80 100644 --- a/unsloth/__init__.py +++ b/unsloth/__init__.py @@ -247,6 +247,7 @@ 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 ( From 1f85f39e0a85324de87cda091f28300dea853e3d Mon Sep 17 00:00:00 2001 From: danielhanchen Date: Thu, 8 Jan 2026 04:14:53 +0000 Subject: [PATCH 20/44] Fix FBGEMM/CUTLASS errors on SM100 (Blackwell) GPUs This PR fixes the "Arch conditional MMA instruction used without targeting appropriate compute capability. Aborting." errors that occur when using FBGEMM on Blackwell GPUs (B200/B100, SM100). Changes: - Add stderr filters in import_fixes.py for CUTLASS/FBGEMM MMA errors - Add warning filters for various deprecation messages - Update check_fbgemm_gpu_version() to disable FBGEMM instead of raising an error when old versions are detected - Update test_has_fbgemm() in fp8.py to catch broader CUTLASS/CUDA errors and gracefully fall back to Triton kernels - Update loader_utils.py to disable FBGEMM instead of raising ValueError for old fbgemm_gpu versions The key behavior change is that FBGEMM errors no longer crash the script. Instead, FBGEMM is disabled and Triton kernels are used automatically. This allows Unsloth to work on SM100 GPUs where CUTLASS SM90 kernels fail, and also gracefully handles old FBGEMM versions. --- unsloth/import_fixes.py | 34 +++++++++++++++++++++++++++++++--- unsloth/kernels/fp8.py | 22 +++++++++++++++++++--- unsloth/models/loader_utils.py | 10 +++++++--- 3 files changed, 57 insertions(+), 9 deletions(-) diff --git a/unsloth/import_fixes.py b/unsloth/import_fixes.py index 1e05e462e9..5ad341e2ac 100644 --- a/unsloth/import_fixes.py +++ b/unsloth/import_fixes.py @@ -94,16 +94,40 @@ 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 +347,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/models/loader_utils.py b/unsloth/models/loader_utils.py index fe2a89d893..9656cc9d26 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,11 @@ 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 From 41b7fe0c672f518b41519f25f049f87084694f1c Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 8 Jan 2026 04:15:17 +0000 Subject: [PATCH 21/44] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/import_fixes.py | 12 +++--------- unsloth/models/loader_utils.py | 1 + 2 files changed, 4 insertions(+), 9 deletions(-) diff --git a/unsloth/import_fixes.py b/unsloth/import_fixes.py index 5ad341e2ac..27e5342e20 100644 --- a/unsloth/import_fixes.py +++ b/unsloth/import_fixes.py @@ -114,20 +114,14 @@ if os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") != "1": "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" - ) + 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" - ) + 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' diff --git a/unsloth/models/loader_utils.py b/unsloth/models/loader_utils.py index 9656cc9d26..1e5533c25c 100644 --- a/unsloth/models/loader_utils.py +++ b/unsloth/models/loader_utils.py @@ -419,6 +419,7 @@ def _get_fp8_mode_and_check_settings( # 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." From 24bbe8a97a0bcb97b204eca82f680e04c5634b07 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 8 Jan 2026 11:35:00 +0000 Subject: [PATCH 22/44] Fix bugs and add improvements to RawTextDataLoader - Fix test file: use return_tokenized instead of return_tensors - Fix test file: use text_dataset instead of undefined dataset variable - Move parameter validation to constructor (fail fast on invalid params) - Add labels field in tokenized output for causal LM training - Add empty file handling with clear error message - Add tests for constructor validation and labels field --- tests/test_raw_text.py | 23 ++++++++++++++++++++--- unsloth/dataprep/raw_text.py | 19 +++++++++++-------- 2 files changed, 31 insertions(+), 11 deletions(-) diff --git a/tests/test_raw_text.py b/tests/test_raw_text.py index 5306c68fa6..7c7272a551 100644 --- a/tests/test_raw_text.py +++ b/tests/test_raw_text.py @@ -100,12 +100,12 @@ def test_raw_text_loader(): loader = RawTextDataLoader(tokenizer, chunk_size = 5, stride = 2) # Test loading with text output (legacy mode) - text_dataset = loader.load_from_file(test_file, return_tensors = False) + 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_tensors = True) + 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 @@ -124,13 +124,30 @@ def test_raw_text_loader(): 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(dataset) + stats = preprocessor.validate_dataset(text_dataset) assert stats["total_samples"] > 0, "Should count samples" assert "warnings" in stats, "Should include warnings" diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index f7c1bf7856..da64565bbc 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -36,6 +36,12 @@ SUPPORTED_FORMATS = { 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 @@ -52,6 +58,8 @@ class RawTextDataLoader: 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 ) @@ -86,8 +94,10 @@ class RawTextDataLoader: # 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} + {"input_ids": input_ids, "attention_mask": attention_mask, "labels": labels} ) else: # If chunks are text strings (backward compatibility) @@ -101,13 +111,6 @@ class RawTextDataLoader: 3. Maintains context with stride overlap 4. Returns tokenized chunks directly (more efficient) or text chunks """ - 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}) to progress the chunking loop" - ) - # 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"] From 3ac4f3c21318b69b22183a309f8abfe58a8f75a9 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 8 Jan 2026 11:35:21 +0000 Subject: [PATCH 23/44] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/test_raw_text.py | 8 ++++++-- unsloth/dataprep/raw_text.py | 6 +++++- 2 files changed, 11 insertions(+), 3 deletions(-) diff --git a/tests/test_raw_text.py b/tests/test_raw_text.py index 7c7272a551..9f2e8cda4e 100644 --- a/tests/test_raw_text.py +++ b/tests/test_raw_text.py @@ -125,8 +125,12 @@ def test_raw_text_loader(): ), "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" + 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: diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index da64565bbc..ba010edabb 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -97,7 +97,11 @@ class RawTextDataLoader: # 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} + { + "input_ids": input_ids, + "attention_mask": attention_mask, + "labels": labels, + } ) else: # If chunks are text strings (backward compatibility) From 1cdf751f8eaad722be9f248ce097645504d1c3e8 Mon Sep 17 00:00:00 2001 From: Rachel Li Date: Thu, 8 Jan 2026 18:44:22 -0500 Subject: [PATCH 24/44] Fix Kaggle telemetry misclassification when COLAB_ keys exist Problem: Kaggle notebook environments can expose both KAGGLE_* and COLAB_* environment keys. _get_statistics currently checks COLAB_ before KAGGLE_, causing Kaggle sessions to be labeled colab/colabpro. Prefer filesystem markers (e.g. /kaggle/working, /content + /opt/colab) before env-key heuristics, then fall back to the existing env-key checks. This avoids misclassification when providers leak overlapping env vars. Kaggle test notebook: https://www.kaggle.com/code/hnxnq07/kaggle-stats-gathering-test --- unsloth/models/_utils.py | 165 +++++++++++++++++++++------------------ 1 file changed, 90 insertions(+), 75 deletions(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index e6c4a12874..77564c03d5 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -1108,85 +1108,100 @@ def _get_statistics(statistics = None, force_download = True): 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", - ) + # 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" - - pass - try: - statistics = try_vllm_check() - except: - statistics = "other" - if statistics is not None: - import tempfile - from huggingface_hub import snapshot_download - from unsloth_zoo.rl_environments import execute_with_time_limit - - if has_internet(): - - def stats_check(): - with tempfile.TemporaryDirectory(ignore_cleanup_errors = True) as f: - snapshot_download( - f"unslothai/{statistics}", - force_download = True, - cache_dir = f, - local_dir = f, + 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", ) - - time_limited_stats_check = execute_with_time_limit(120)(stats_check) - try: - time_limited_stats_check() - except TimeoutError: - raise TimeoutError( - "Unsloth: HuggingFace seems to be down after trying for 120 seconds :(\n" - "Check https://status.huggingface.co/ for more details.\n" - "As a temporary measure, use modelscope with the same model name ie:\n" - "```\n" - "pip install modelscope\n" - "import os; os.environ['UNSLOTH_USE_MODELSCOPE'] = '1'\n" - "from unsloth import FastLanguageModel\n" - "model = FastLanguageModel.from_pretrained('unsloth/gpt-oss-20b')\n" - "```" - ) - except: - # Try no time limit check - stats_check() + 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" + + pass + try: + statistics = try_vllm_check() + except: + statistics = "other" + if statistics is not None: + import tempfile + from huggingface_hub import snapshot_download + from unsloth_zoo.rl_environments import execute_with_time_limit + + if has_internet(): + + def stats_check(): + with tempfile.TemporaryDirectory(ignore_cleanup_errors = True) as f: + snapshot_download( + f"unslothai/{statistics}", + force_download = True, + cache_dir = f, + local_dir = f, + ) + + time_limited_stats_check = execute_with_time_limit(120)(stats_check) + try: + time_limited_stats_check() + except TimeoutError: + raise TimeoutError( + "Unsloth: HuggingFace seems to be down after trying for 120 seconds :(\n" + "Check https://status.huggingface.co/ for more details.\n" + "As a temporary measure, use modelscope with the same model name ie:\n" + "```\n" + "pip install modelscope\n" + "import os; os.environ['UNSLOTH_USE_MODELSCOPE'] = '1'\n" + "from unsloth import FastLanguageModel\n" + "model = FastLanguageModel.from_pretrained('unsloth/gpt-oss-20b')\n" + "```" + ) + except: + # Try no time limit check + stats_check() def get_statistics(local_files_only = False): From 67bef80b1f2fcfe1f4d6ef0bd96cca2f7955ef83 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 8 Jan 2026 23:49:54 +0000 Subject: [PATCH 25/44] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/models/_utils.py | 16 +++++++++------- 1 file changed, 9 insertions(+), 7 deletions(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 77564c03d5..756d884de1 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -1140,7 +1140,7 @@ def _get_statistics(statistics = None, force_download = True): statistics = "lambda" # else: statistics = "other" else: - + def try_vllm_check(): vendor_files = ( "/sys/class/dmi/id/product_version", @@ -1150,7 +1150,7 @@ def _get_statistics(statistics = None, force_download = True): "/sys/class/dmi/id/sys_vendor", ) from pathlib import Path - + for vendor_file in vendor_files: path = Path(vendor_file) if path.is_file(): @@ -1162,7 +1162,7 @@ def _get_statistics(statistics = None, force_download = True): elif "google" in file_content: return "gcp" return "other" - + pass try: statistics = try_vllm_check() @@ -1172,18 +1172,20 @@ def _get_statistics(statistics = None, force_download = True): import tempfile from huggingface_hub import snapshot_download from unsloth_zoo.rl_environments import execute_with_time_limit - + if has_internet(): - + def stats_check(): - with tempfile.TemporaryDirectory(ignore_cleanup_errors = True) as f: + with tempfile.TemporaryDirectory( + ignore_cleanup_errors = True + ) as f: snapshot_download( f"unslothai/{statistics}", force_download = True, cache_dir = f, local_dir = f, ) - + time_limited_stats_check = execute_with_time_limit(120)(stats_check) try: time_limited_stats_check() From e13160ddc8febd25b569bc31970efce7dd9c57a7 Mon Sep 17 00:00:00 2001 From: Rachel Li Date: Thu, 8 Jan 2026 19:04:30 -0500 Subject: [PATCH 26/44] Update _utils.py fixed indentation --- unsloth/models/_utils.py | 68 +++++++++++++++++++++------------------- 1 file changed, 35 insertions(+), 33 deletions(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 756d884de1..9bf02beb30 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -1121,6 +1121,7 @@ def _get_statistics(statistics = None, force_download = True): statistics = "runpod" except Exception: pass + # Fallback to env-key detection if statistics is None: if "\nKAGGLE_" in keynames: @@ -1168,42 +1169,43 @@ def _get_statistics(statistics = None, force_download = True): statistics = try_vllm_check() except: statistics = "other" - if statistics is not None: - import tempfile - from huggingface_hub import snapshot_download - from unsloth_zoo.rl_environments import execute_with_time_limit + + if statistics is not None: + import tempfile + from huggingface_hub import snapshot_download + from unsloth_zoo.rl_environments import execute_with_time_limit - if has_internet(): + if has_internet(): - def stats_check(): - with tempfile.TemporaryDirectory( - ignore_cleanup_errors = True - ) as f: - snapshot_download( - f"unslothai/{statistics}", - force_download = True, - cache_dir = f, - local_dir = f, - ) - - time_limited_stats_check = execute_with_time_limit(120)(stats_check) - try: - time_limited_stats_check() - except TimeoutError: - raise TimeoutError( - "Unsloth: HuggingFace seems to be down after trying for 120 seconds :(\n" - "Check https://status.huggingface.co/ for more details.\n" - "As a temporary measure, use modelscope with the same model name ie:\n" - "```\n" - "pip install modelscope\n" - "import os; os.environ['UNSLOTH_USE_MODELSCOPE'] = '1'\n" - "from unsloth import FastLanguageModel\n" - "model = FastLanguageModel.from_pretrained('unsloth/gpt-oss-20b')\n" - "```" + def stats_check(): + with tempfile.TemporaryDirectory( + ignore_cleanup_errors = True + ) as f: + snapshot_download( + f"unslothai/{statistics}", + force_download = True, + cache_dir = f, + local_dir = f, ) - except: - # Try no time limit check - stats_check() + + time_limited_stats_check = execute_with_time_limit(120)(stats_check) + try: + time_limited_stats_check() + except TimeoutError: + raise TimeoutError( + "Unsloth: HuggingFace seems to be down after trying for 120 seconds :(\n" + "Check https://status.huggingface.co/ for more details.\n" + "As a temporary measure, use modelscope with the same model name ie:\n" + "```\n" + "pip install modelscope\n" + "import os; os.environ['UNSLOTH_USE_MODELSCOPE'] = '1'\n" + "from unsloth import FastLanguageModel\n" + "model = FastLanguageModel.from_pretrained('unsloth/gpt-oss-20b')\n" + "```" + ) + except: + # Try no time limit check + stats_check() def get_statistics(local_files_only = False): From e7d68f3e5788d21d5afaac182099716c9b5e248e Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 9 Jan 2026 00:04:59 +0000 Subject: [PATCH 27/44] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/models/_utils.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 9bf02beb30..701e6f633c 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -1121,7 +1121,7 @@ def _get_statistics(statistics = None, force_download = True): statistics = "runpod" except Exception: pass - + # Fallback to env-key detection if statistics is None: if "\nKAGGLE_" in keynames: @@ -1169,7 +1169,7 @@ def _get_statistics(statistics = None, force_download = True): statistics = try_vllm_check() except: statistics = "other" - + if statistics is not None: import tempfile from huggingface_hub import snapshot_download @@ -1178,9 +1178,7 @@ def _get_statistics(statistics = None, force_download = True): if has_internet(): def stats_check(): - with tempfile.TemporaryDirectory( - ignore_cleanup_errors = True - ) as f: + with tempfile.TemporaryDirectory(ignore_cleanup_errors = True) as f: snapshot_download( f"unslothai/{statistics}", force_download = True, From 84701c55ffa7437d11d0e32c5d7a6e77681862c4 Mon Sep 17 00:00:00 2001 From: Rachel Li Date: Thu, 8 Jan 2026 19:20:24 -0500 Subject: [PATCH 28/44] Fix telemetry ping regression for explicit statistics Fixed Codex regression: keep snapshot_download pings for explicit statistics values; detection only runs when statistics is None. Also replaced bare except. --- unsloth/models/_utils.py | 62 ++++++++++++++++++++-------------------- 1 file changed, 31 insertions(+), 31 deletions(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 701e6f633c..0cb6c8975e 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -1170,40 +1170,40 @@ def _get_statistics(statistics = None, force_download = True): except: statistics = "other" - if statistics is not None: - import tempfile - from huggingface_hub import snapshot_download - from unsloth_zoo.rl_environments import execute_with_time_limit + if statistics is not None: + import tempfile + from huggingface_hub import snapshot_download + from unsloth_zoo.rl_environments import execute_with_time_limit - if has_internet(): + if has_internet(): - def stats_check(): - with tempfile.TemporaryDirectory(ignore_cleanup_errors = True) as f: - snapshot_download( - f"unslothai/{statistics}", - force_download = True, - cache_dir = f, - local_dir = f, - ) - - time_limited_stats_check = execute_with_time_limit(120)(stats_check) - try: - time_limited_stats_check() - except TimeoutError: - raise TimeoutError( - "Unsloth: HuggingFace seems to be down after trying for 120 seconds :(\n" - "Check https://status.huggingface.co/ for more details.\n" - "As a temporary measure, use modelscope with the same model name ie:\n" - "```\n" - "pip install modelscope\n" - "import os; os.environ['UNSLOTH_USE_MODELSCOPE'] = '1'\n" - "from unsloth import FastLanguageModel\n" - "model = FastLanguageModel.from_pretrained('unsloth/gpt-oss-20b')\n" - "```" + def stats_check(): + with tempfile.TemporaryDirectory(ignore_cleanup_errors = True) as f: + snapshot_download( + f"unslothai/{statistics}", + force_download = True, + cache_dir = f, + local_dir = f, ) - except: - # Try no time limit check - stats_check() + + time_limited_stats_check = execute_with_time_limit(120)(stats_check) + try: + time_limited_stats_check() + except TimeoutError: + raise TimeoutError( + "Unsloth: HuggingFace seems to be down after trying for 120 seconds :(\n" + "Check https://status.huggingface.co/ for more details.\n" + "As a temporary measure, use modelscope with the same model name ie:\n" + "```\n" + "pip install modelscope\n" + "import os; os.environ['UNSLOTH_USE_MODELSCOPE'] = '1'\n" + "from unsloth import FastLanguageModel\n" + "model = FastLanguageModel.from_pretrained('unsloth/gpt-oss-20b')\n" + "```" + ) + except Exception: + # Try no time limit check + stats_check() def get_statistics(local_files_only = False): From 1193c9f5268f33123054a58c8baf699bf2597a80 Mon Sep 17 00:00:00 2001 From: Rachel Li Date: Thu, 8 Jan 2026 19:32:33 -0500 Subject: [PATCH 29/44] Fix Kaggle telemetry detection & address review feedback - Fix Kaggle misclassification by prioritizing filesystem markers over env vars - Preserve telemetry pings when statistics is explicitly provided - Replace bare except with except Exception - Minor cleanup based on automated review feedback --- unsloth/models/_utils.py | 10 +++------- 1 file changed, 3 insertions(+), 7 deletions(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 0cb6c8975e..77c8da2576 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -1106,9 +1106,7 @@ 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 - else: + if statistics is None: # Prefer filesystem markers (harder to misidentify) before env-key matching try: from pathlib import Path @@ -1150,7 +1148,6 @@ def _get_statistics(statistics = None, force_download = True): "/sys/class/dmi/id/chassis_asset_tag", "/sys/class/dmi/id/sys_vendor", ) - from pathlib import Path for vendor_file in vendor_files: path = Path(vendor_file) @@ -1163,11 +1160,10 @@ def _get_statistics(statistics = None, force_download = True): elif "google" in file_content: return "gcp" return "other" - - pass + try: statistics = try_vllm_check() - except: + except Exception: statistics = "other" if statistics is not None: From 3ce1060dd162cc6a1451e57ae129e384343878e0 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 9 Jan 2026 00:33:00 +0000 Subject: [PATCH 30/44] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/models/_utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 77c8da2576..b0ad0a3082 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -1160,7 +1160,7 @@ def _get_statistics(statistics = None, force_download = True): elif "google" in file_content: return "gcp" return "other" - + try: statistics = try_vllm_check() except Exception: From 0bff0ffbe504d35ba407c20baf3250e1d25d81c6 Mon Sep 17 00:00:00 2001 From: Kaitao Yang Date: Wed, 7 Jan 2026 22:54:35 -0800 Subject: [PATCH 31/44] reduce code duplication by _offload_frozen_module_for_training --- unsloth/models/llama.py | 91 ++++++++++++++++++++++++++--------------- 1 file changed, 57 insertions(+), 34 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 92d51b73ad..cfd900c460 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.tuners.tuners_utils 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: From 6465496ab2b7885f1b3d7c98046b14bb1f7f2ac6 Mon Sep 17 00:00:00 2001 From: danielhanchen Date: Fri, 9 Jan 2026 23:24:39 +0000 Subject: [PATCH 32/44] fix: use peft.utils.other for ModulesToSaveWrapper import ModulesToSaveWrapper was removed from peft.tuners.tuners_utils in PEFT 0.16.0. The class has been available in peft.utils.other since at least PEFT 0.7.1, which is the minimum version Unsloth requires. This fixes the ImportError when using PEFT >= 0.16.0. --- unsloth/models/llama.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index cfd900c460..39f2ba1460 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -146,7 +146,7 @@ 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.tuners.tuners_utils import ModulesToSaveWrapper +from peft.utils.other import ModulesToSaveWrapper def _offload_frozen_module_for_training( From 5b422f7a0634c059cd1885b189d83fea7ec148e3 Mon Sep 17 00:00:00 2001 From: Duc-Viet Hoang Date: Mon, 12 Jan 2026 10:03:54 +0700 Subject: [PATCH 33/44] Complete disable `gradient_checkpointing` for vision when `use_gradient_checkpointing=False` --- unsloth/models/rl.py | 6 ++++-- unsloth/models/vision.py | 2 +- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index e945a80354..9ec4d76ed3 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -264,17 +264,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() 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"): From eaf3f932e007f809fe67cef56765886f8f9b1cae Mon Sep 17 00:00:00 2001 From: Francesco Bertolotti Date: Mon, 12 Jan 2026 16:19:43 +0100 Subject: [PATCH 34/44] wrong number of dimensions --- unsloth/kernels/swiglu.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/unsloth/kernels/swiglu.py b/unsloth/kernels/swiglu.py index b321f5179e..9e2680e862 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 n_elements = e.numel() grid = lambda meta: (triton.cdiv(n_elements, meta["BLOCK_SIZE"]),) with torch_gpu_device(e.device): From 2f8c4d962be5e09b54a7b3ea6c9372d684204139 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 12 Jan 2026 19:08:13 +0000 Subject: [PATCH 35/44] [pre-commit.ci] pre-commit autoupdate MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit updates: - [github.com/astral-sh/ruff-pre-commit: v0.14.10 → v0.14.11](https://github.com/astral-sh/ruff-pre-commit/compare/v0.14.10...v0.14.11) --- .pre-commit-config.yaml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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: From 0f0b87078157435861e74ada0892d5269907d544 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 12 Jan 2026 21:32:20 -0800 Subject: [PATCH 36/44] Apply suggestion from @danielhanchen --- unsloth/kernels/swiglu.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/unsloth/kernels/swiglu.py b/unsloth/kernels/swiglu.py index 9e2680e862..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): From 45eeae95c5f7257a4fa39fa9c0f70933541d83f5 Mon Sep 17 00:00:00 2001 From: Michael Han <107991372+shimmyshimmer@users.noreply.github.com> Date: Wed, 14 Jan 2026 03:45:35 -0800 Subject: [PATCH 37/44] Update template.md --- .github/ISSUE_TEMPLATE/bug---issue.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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/ From 7f6dc63dc89aa5706b5650d95b9ff82bb5cb3c42 Mon Sep 17 00:00:00 2001 From: Datta Nimmaturi Date: Thu, 15 Jan 2026 11:11:50 +0000 Subject: [PATCH 38/44] use non lora model as base for RL --- unsloth/models/rl.py | 32 ++++++++++++++++++++++++++++++++ 1 file changed, 32 insertions(+) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index e945a80354..62141b8700 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -1166,6 +1166,38 @@ 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. From f5dde984a16c6e2903c7e36bd39ed7e6aad0f404 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 15 Jan 2026 11:25:10 +0000 Subject: [PATCH 39/44] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/models/rl.py | 21 ++++++++++++--------- 1 file changed, 12 insertions(+), 9 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 62141b8700..d5524ff5f5 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -1173,30 +1173,33 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import # 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\)' + 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') + 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") + 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)] + 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) + 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: From 2164423ea6e7ef0e475dfd0d01068b68e38f6094 Mon Sep 17 00:00:00 2001 From: pluesclues <136766175+pluesclues@users.noreply.github.com> Date: Thu, 15 Jan 2026 08:01:19 -0500 Subject: [PATCH 40/44] Merge pull request #3628 from pluesclues/alternative_compute_chunked_loss Chunk Across Batch and Context length for logprob calculations for grpo --- unsloth/models/rl.py | 28 ++- unsloth/models/rl_replacements.py | 303 +++++++++++++++++++++++------- 2 files changed, 267 insertions(+), 64 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 14a07e193a..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 @@ -300,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} @@ -321,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, ): @@ -332,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 @@ -1029,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 ) @@ -1037,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, @@ -1058,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, ) 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 From 6edbfbc43511162335d8d0e601f8f858783227b1 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 15 Jan 2026 05:09:26 -0800 Subject: [PATCH 41/44] Update _utils.py --- unsloth/models/_utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index b0ad0a3082..6028635fb9 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", From ecd10f2e55e90cdafe4f71c9cbbfa58a4b89dc70 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 15 Jan 2026 07:00:25 -0800 Subject: [PATCH 42/44] Update pyproject.toml --- pyproject.toml | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) 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", From c4a718ca3161158418602c4ee55d09c86cfa3839 Mon Sep 17 00:00:00 2001 From: Michael Han <107991372+shimmyshimmer@users.noreply.github.com> Date: Thu, 15 Jan 2026 08:01:01 -0800 Subject: [PATCH 43/44] Update README.md --- README.md | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) 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) From 20c434cd7774aa3cae1241d9c2c8c7d1f344eeb7 Mon Sep 17 00:00:00 2001 From: electroglyph Date: Thu, 15 Jan 2026 20:02:29 -0800 Subject: [PATCH 44/44] add weight-only int8 QAT scheme and update tests for torchao 0.15.0 (#3859) * add int8 weight-only QAT scheme, add test, fix tests for current torchao version * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * change quantization to PerAxis * lambda =/ * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * add torchao messages, remove group_size from int8 * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * raise exception on missing torchao * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * touch up the torchao imports * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> --- tests/utils/test_qat.py | 63 ++++++++++++++++++++++++++-------------- unsloth/models/_utils.py | 59 ++++++++++++++++++++++++++++--------- 2 files changed, 87 insertions(+), 35 deletions(-) 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/models/_utils.py b/unsloth/models/_utils.py index 6028635fb9..28dbe450ed 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -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): @@ -2211,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. @@ -2230,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 = ( @@ -2243,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() ) @@ -2252,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 = [ @@ -2276,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 = ( @@ -2288,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}"