From e57b1c73fde87f382f3123d02a84a164f934d6bb Mon Sep 17 00:00:00 2001 From: vangmay Date: Thu, 20 Nov 2025 21:08:33 +0800 Subject: [PATCH] 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):