Make the chunk function efficient
This commit is contained in:
parent
12474b14d9
commit
e57b1c73fd
3 changed files with 100 additions and 31 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue