Make the chunk function efficient

This commit is contained in:
vangmay 2025-11-20 21:08:33 +08:00
commit e57b1c73fd
3 changed files with 100 additions and 31 deletions

View file

@ -61,6 +61,7 @@ def test_raw_text_loader():
class MockTokenizer:
def __init__(self):
self.eos_token = "</s>"
self.eos_token_id = 2 # Mock EOS token ID
def __call__(self, text, return_tensors = None, add_special_tokens = False):
words = text.split()
@ -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()

View file

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

View file

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