forked from unslothai/unsloth
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_raw_text.py
More file actions
411 lines (340 loc) · 15.9 KB
/
Copy pathtest_raw_text.py
File metadata and controls
411 lines (340 loc) · 15.9 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
#!/usr/bin/env python3
"""Minimal test for raw text training, without heavy dependencies."""
import sys
import os
import tempfile
from pathlib import Path
import importlib.util
# Mock the datasets module (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):
# Column access, e.g. dataset['text'].
return self.data[idx]
elif isinstance(idx, int):
# Row access 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)
# __spec__ must be set so importlib.util.find_spec doesn't raise ValueError when
# transformers' import_utils later probes for the real `datasets` package.
datasets_mock = type(sys)("datasets")
datasets_mock.__spec__ = importlib.util.spec_from_loader("datasets", loader = None)
datasets_mock.Dataset = MockDataset
sys.modules["datasets"] = datasets_mock
# Import raw_text 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."""
class MockTokenizer:
def __init__(self):
self.eos_token = "</s>"
self.eos_token_id = 2
def __call__(
self,
text,
return_tensors = None,
add_special_tokens = False,
):
words = text.split()
token_ids = list(range(len(words)))
if return_tensors == "pt":
class MockTensor:
def __init__(self, data):
self.data = data
def __getitem__(self, idx):
return self.data
def __len__(self):
return len(self.data)
def tolist(self):
return 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])
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:
tokenizer = MockTokenizer()
loader = RawTextDataLoader(tokenizer, chunk_size = 5, stride = 2)
# Text output (legacy mode).
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"
# Tokenized output (new efficient mode).
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
), "Dataset should have 'input_ids' column"
assert (
"attention_mask" in tokenized_dataset.column_names
), "Dataset should have 'attention_mask' column"
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"
# labels field (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"
# 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)
# smart_chunk_text validation: called directly, chunk_size/stride are its own
# arguments and bypass the constructor guard, so it must guard itself or an
# invalid stride makes `start_idx += chunk_size - stride` non-positive and the
# chunking loop never terminates (hangs).
long_text = "This is a test file for raw text training. " * 10
valid_chunks = loader.smart_chunk_text(long_text, chunk_size = 5, stride = 2)
assert len(valid_chunks) > 0, "Valid stride should produce chunks"
try:
loader.smart_chunk_text(long_text, chunk_size = 5, stride = 5)
assert False, "Should raise ValueError for stride == chunk_size"
except ValueError as e:
assert "stride" in str(e) and "chunk_size" in str(e)
try:
loader.smart_chunk_text(long_text, 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)
# Preprocessor.
preprocessor = TextPreprocessor()
clean_text = preprocessor.clean_text(" messy text \n\n\n ")
assert "messy text" in clean_text, "Should clean text properly"
paragraph_text = preprocessor.clean_text("Line 1\r\n\r\n\r\nLine 2")
assert (
paragraph_text == "Line 1\n\nLine 2"
), "Should preserve paragraph breaks while normalizing newlines"
# Non-ASCII horizontal whitespace (NBSP, thin/em/ideographic space, VT, FF) must
# normalize to one ASCII space, not be deleted, or adjacent words fuse on HTML/PDF/OCR input.
unicode_whitespace_cases = [
("hello\u00a0world", "hello world"),
("hello\u202fworld", "hello world"),
("hello\u2009world", "hello world"),
("hello\u3000world", "hello world"),
("hello\u2002world", "hello world"),
("hello\x0bworld", "hello world"),
("hello\x0cworld", "hello world"),
]
for raw, expected in unicode_whitespace_cases:
assert preprocessor.clean_text(raw) == expected, (
f"Should normalize Unicode/control whitespace to a single space " f"for {raw!r}"
)
# Mixed paragraph + Unicode whitespace.
mixed = preprocessor.clean_text("Section\u00a01\r\n\r\nBody\ftext\u202fhere")
assert (
mixed == "Section 1\n\nBody text here"
), "Should preserve paragraph breaks and normalize Unicode whitespace simultaneously"
# Tabs collapse to a single space.
assert preprocessor.clean_text("a\tb") == "a b"
assert preprocessor.clean_text("a\t\tb") == "a b"
# Spaces around newlines trimmed on both sides, even across multiple newlines.
assert preprocessor.clean_text("foo \n\n bar") == "foo\n\nbar"
# Stripping a non-ASCII char between spaces must not leave a double space
# (also guards idempotence: otherwise "word1 (c) word2" needs a second pass).
assert preprocessor.clean_text("word1 \u00a9 word2") == "word1 word2"
assert preprocessor.clean_text("a \u00e9 b") == "a b"
assert preprocessor.clean_text("prefix \U0001f600 suffix") == "prefix suffix"
# Stripping a non-ASCII char adjacent to a newline must not leave a stray space.
assert preprocessor.clean_text("foo \u00e9\nbar") == "foo\nbar"
assert preprocessor.clean_text("foo\n\u00e9 bar") == "foo\nbar"
# The double-space collapse must not swallow a paragraph break near a non-ASCII char.
assert preprocessor.clean_text("a \u00a9\n\nb") == "a\n\nb"
# Idempotence: clean_text twice == once.
idempotent_inputs = [
" messy text \n\n\n ",
"Line 1\r\n\r\n\r\nLine 2",
"hello\u00a0world",
"Section\u00a01\r\n\r\nBody\ftext\u202fhere",
"word1 \u00a9 word2",
"a \u00e9 b",
]
for raw in idempotent_inputs:
once = preprocessor.clean_text(raw)
twice = preprocessor.clean_text(once)
assert once == twice, f"clean_text should be idempotent for {raw!r}"
# Validation.
stats = preprocessor.validate_dataset(text_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:
os.unlink(test_file)
def test_smart_chunk_text_single_chunk_no_eos_returns_plain_list():
"""smart_chunk_text's single-chunk branch must return a plain list for
input_ids even when the tokenizer has no eos_token_id, matching the
multi-chunk branch's unconditional tolist()/list() conversion."""
class MockTensor:
def __init__(self, data):
self.data = data
def __getitem__(self, idx):
return self.data
def __len__(self):
return len(self.data)
def tolist(self):
return self.data
class MockTokenizerNoEos:
def __init__(self):
self.eos_token = None
self.eos_token_id = None
def __call__(
self,
text,
return_tensors = None,
add_special_tokens = False,
):
token_ids = list(range(len(text.split())))
if return_tensors == "pt":
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)
loader = RawTextDataLoader(MockTokenizerNoEos(), chunk_size = 2048, stride = 512)
result = loader.smart_chunk_text(
"hello world short text", chunk_size = 2048, stride = 512, return_tokenized = True
)
input_ids = result[0]["input_ids"]
assert isinstance(
input_ids, list
), f"input_ids should be a plain list even without an eos_token_id, got {type(input_ids)}"
assert input_ids == [0, 1, 2, 3], f"unexpected input_ids: {input_ids}"
print("✅ test_smart_chunk_text_single_chunk_no_eos_returns_plain_list passed!")
return True
def test_load_from_file_skips_non_object_json_lines():
"""Non-object .jsonl lines (valid JSON, not dicts) are skipped, not fatal."""
# "context" contains "text", ["text"] holds it, 42 isn't iterable -- each
# would reach data[field] and raise TypeError without the isinstance guard.
with tempfile.NamedTemporaryFile("w", suffix = ".jsonl", delete = False) as f:
f.write('"context"\n["text", "x"]\n42\n{"text": "keep this"}\n')
path = f.name
try:
text = RawTextDataLoader(None)._read_file_by_format(path, "json_lines")
assert text == "keep this", text
finally:
os.unlink(path)
print("test_load_from_file_skips_non_object_json_lines passed")
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)