Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
92 changes: 92 additions & 0 deletions tests/test_raw_text.py
Original file line number Diff line number Diff line change
Expand Up @@ -312,8 +312,100 @@ def test_load_from_file_skips_non_object_json_lines():
return True


def test_smart_chunk_text_empty_input_returns_no_chunks():
"""Empty/whitespace text must yield no chunks. This tokenizer keeps one token
per char (like BPE/SentencePiece keeping spaces), so a len(tokens)==0 check
would miss whitespace; the fix guards on text.strip() before tokenizing."""

class WhitespacePreservingTokenizer:
def __init__(self, eos_token_id):
self.eos_token = "</s>" if eos_token_id is not None else None
self.eos_token_id = eos_token_id

def __call__(
self,
text,
return_tensors = None,
add_special_tokens = False,
):
token_ids = [ord(c) % 100 for c in text] # whitespace -> real tokens
if return_tensors == "pt":
return {"input_ids": [token_ids]}
return {"input_ids": token_ids}

def decode(
self,
token_ids,
skip_special_tokens = False,
):
return "".join(chr(32 + (t % 90)) for t in token_ids)

for eos_token_id in (2, None):
loader = RawTextDataLoader(
WhitespacePreservingTokenizer(eos_token_id), chunk_size = 2048, stride = 512
)
# Whitespace tokenizes to >0 tokens, so [] proves the pre-tokenize guard.
assert len(loader.tokenizer(" \n\t ")["input_ids"]) > 0
for text in ("", " \n\t "):
for return_tokenized in (True, False):
assert (
loader.smart_chunk_text(
text, chunk_size = 2048, stride = 512, return_tokenized = return_tokenized
)
== []
), f"no chunks for empty input (eos={eos_token_id}, text={text!r}, tokenized={return_tokenized})"
assert loader.chunk_text(text, return_tokenized = return_tokenized) == [], (
f"chunk_text: no chunks for empty input "
f"(eos={eos_token_id}, text={text!r}, tokenized={return_tokenized})"
)
print("test_smart_chunk_text_empty_input_returns_no_chunks passed")
return True


def test_load_from_files_all_empty_raises():
"""All-empty file list must raise (like load_from_file) instead of returning
a 0-row text-column dataset in return_tokenized mode."""

class WhitespacePreservingTokenizer:
eos_token = "</s>"
eos_token_id = 2

def __call__(
self,
text,
return_tensors = None,
add_special_tokens = False,
):
token_ids = [ord(c) % 100 for c in text]
if return_tensors == "pt":
return {"input_ids": [token_ids]}
return {"input_ids": token_ids}

loader = RawTextDataLoader(WhitespacePreservingTokenizer(), chunk_size = 2048, stride = 512)
paths = []
try:
for content in ("", " \n\t "):
with tempfile.NamedTemporaryFile("w", suffix = ".txt", delete = False) as f:
f.write(content)
paths.append(f.name)
raised = False
try:
loader.load_from_files(paths, return_tokenized = True)
except ValueError as e:
raised = True
assert "empty" in str(e).lower() or "whitespace" in str(e).lower(), str(e)
assert raised, "load_from_files must raise when all files are empty/whitespace"
finally:
for p in paths:
os.unlink(p)
print("test_load_from_files_all_empty_raises passed")
return True


if __name__ == "__main__":
success = test_raw_text_loader()
success = test_smart_chunk_text_single_chunk_no_eos_returns_plain_list() and success
success = test_load_from_file_skips_non_object_json_lines() and success
success = test_smart_chunk_text_empty_input_returns_no_chunks() and success
success = test_load_from_files_all_empty_raises() and success
sys.exit(0 if success else 1)
10 changes: 10 additions & 0 deletions unsloth/dataprep/raw_text.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,10 @@ def load_from_files(
text_content, self.chunk_size, self.stride, return_tokenized
)
all_chunks.extend(chunks)
if not all_chunks:
# All files empty/whitespace: raise like load_from_file instead of
# create_causal_dataset([]) returning a 0-row text-column dataset.
raise ValueError("All files are empty or contain only whitespace")
return self.create_causal_dataset(all_chunks)

def chunk_text(
Expand Down Expand Up @@ -139,6 +143,12 @@ def smart_chunk_text(
f"stride ({stride}) must be smaller than chunk_size ({chunk_size}) to progress the chunking loop"
)

# Skip empty/whitespace text before tokenizing: BPE/SentencePiece emit
# real tokens for spaces/newlines, so a len(tokens)==0 check misses it
# and would yield a degenerate lone-EOS sample. Mirrors load_from_file.
if not text or not text.strip():
return []

# Tokenize the whole text once for accurate token counts
tokenized = self.tokenizer(text, return_tensors = "pt", add_special_tokens = False)
tokens = tokenized["input_ids"]
Expand Down
Loading