Skip to content
Open
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
2 changes: 1 addition & 1 deletion python/dolma/cli/tokenizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -131,7 +131,7 @@ class TokenizationConfig:
help="Number of parallel processes to use.",
)
files_per_process: Optional[int] = field(
default=None,
default=1,
help="Number of files to process per process.",
)
batch_size: int = field(
Expand Down
13 changes: 9 additions & 4 deletions python/dolma/tokenizer/executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,8 @@ def process_single(cls, source_path: str, destination_path: str, queue: QueueTyp

max_size: int = kwargs.pop("max_size", None) or 1024 * 1024 * 1024
dtype: np.dtype = np.dtype(kwargs.pop("dtype", None) or "uint16")
# max_size is configured in bytes; MemmapWriter expects a token count.
max_tokens = max_size // dtype.itemsize
local_shuffle: int = kwargs.pop("local_shuffle", None) or 10_000
ring_size: int = kwargs.pop("ring_size", None) or 8
sample_ring_prop: bool = kwargs.pop("sample_ring_prop", None) or False
Expand Down Expand Up @@ -132,7 +134,7 @@ def process_single(cls, source_path: str, destination_path: str, queue: QueueTyp

with ExitStack() as stack:
memwriter = stack.enter_context(
MemmapWriter(path=destination_path + f"-{mm_cnt:05d}", dtype=dtype, max_tokens=max_size)
MemmapWriter(path=destination_path + f"-{mm_cnt:05d}", dtype=dtype, max_tokens=max_tokens)
)
cls.increment_progressbar(queue, memmaps=1)

Expand Down Expand Up @@ -232,16 +234,19 @@ def process_single(cls, source_path: str, destination_path: str, queue: QueueTyp
MemmapWriter(
path=destination_path + f"-{mm_cnt:05d}",
dtype=dtype,
max_tokens=max_size,
max_tokens=max_tokens,
)
)
cls.increment_progressbar(queue, memmaps=1)

# shuffle the remaining sequences
random.shuffle(remaining)

# finally, write the remaining sequences
remaining = memwriter.write_many(outputs=remaining, flush=True)
if remaining and mm_cnt > 10_000:
raise RuntimeError(
"Exceeded memmap rollover limit while writing tokenized output. "
"A single tokenized sequence may exceed max_size."
)

# done writing, flush (triggers a write to disk)
memwriter.flush()
Expand Down
18 changes: 14 additions & 4 deletions python/dolma/tokenizer/memmap_writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,17 @@ def write(self, output: TokenizerOutput, flush: bool = False) -> bool:
if self._metadata_file is None:
raise RuntimeError("Metadata file is not open")

if (len(output.tokens) + self._written_tokens) >= self.max_tokens:
token_slice = output.tokens[output.start : output.end]
token_len = len(token_slice)
if token_len == 0:
return True

if token_len > self.max_tokens:
raise ValueError(
f"Tokenized sequence length ({token_len}) exceeds memmap max_tokens ({self.max_tokens})."
)

if (token_len + self._written_tokens) >= self.max_tokens:
# return false if the memmap file is full
return False

Expand All @@ -83,10 +93,10 @@ def write(self, output: TokenizerOutput, flush: bool = False) -> bool:
src=output.src,
loc=output.loc,
start=self._written_tokens,
end=self._written_tokens + output.end,
end=self._written_tokens + token_len,
)
self._memmap_file[self._written_tokens : self._written_tokens + output.end] = output.tokens
self._written_tokens += output.end
self._memmap_file[self._written_tokens : self._written_tokens + token_len] = token_slice
self._written_tokens += token_len

# self._metadata_file.write(msgspec.json.encode(metadata) + b"\n")
self.metadata_writer.writerow(metadata)
Expand Down