| """ |
| Preprocess local code-message JSONL datasets for llm.c GPT-2 training. |
| |
| Expected inputs: |
| - text-only JSONL: {"text": "..."} |
| - messages JSONL: {"messages": [{"role": "...", "content": "..."}, ...]} |
| |
| The output is llm.c's GPT-2 data format: |
| - 256 int32 header values (1024 bytes) |
| - followed by uint16 GPT-2 token ids |
| """ |
|
|
| import argparse |
| import glob |
| import hashlib |
| import json |
| import multiprocessing as mp |
| import os |
| from pathlib import Path |
|
|
| import numpy as np |
| import tiktoken |
| from tqdm import tqdm |
|
|
|
|
| HEADER_SIZE = 256 |
| GPT2_DATA_MAGIC = 20240520 |
| GPT2_DATA_VERSION = 1 |
|
|
| _ENC = None |
| _EOT = None |
|
|
|
|
| def init_worker(): |
| global _ENC, _EOT |
| _ENC = tiktoken.get_encoding("gpt2") |
| _EOT = _ENC._special_tokens["<|endoftext|>"] |
|
|
|
|
| def tokenize_text(text): |
| ids = _ENC.encode_ordinary(text) |
| tokens = np.empty(len(ids) + 1, dtype=np.uint16) |
| tokens[0] = _EOT |
| tokens[1:] = ids |
| return tokens |
|
|
|
|
| def tokenize_record(record): |
| doc_idx, text = record |
| return doc_idx, tokenize_text(text) |
|
|
|
|
| def write_datafile_np(filename, tokens): |
| assert tokens.dtype == np.uint16 |
| assert len(tokens) < 2**31, "token count too large for one shard" |
| header = np.zeros(HEADER_SIZE, dtype=np.int32) |
| header[0] = GPT2_DATA_MAGIC |
| header[1] = GPT2_DATA_VERSION |
| header[2] = len(tokens) |
| num_bytes = HEADER_SIZE * 4 + len(tokens) * tokens.itemsize |
| print(f"writing {len(tokens):,} tokens to {filename} ({num_bytes:,} bytes)") |
| with open(filename, "wb") as f: |
| f.write(header.tobytes()) |
| f.write(tokens.tobytes()) |
|
|
|
|
| class ShardWriter: |
| def __init__(self, output_dir, prefix, shard_size): |
| self.output_dir = Path(output_dir) |
| self.prefix = prefix |
| self.shard_size = shard_size |
| self.shard_index = 0 |
| self.token_count = 0 |
| self.total_tokens = 0 |
| self.buffer = None |
|
|
| def _ensure_buffer(self): |
| if self.buffer is None: |
| self.buffer = np.empty(self.shard_size, dtype=np.uint16) |
|
|
| def _filename(self): |
| return self.output_dir / f"{self.prefix}_{self.shard_index:06d}.bin" |
|
|
| def write(self, tokens): |
| if len(tokens) == 0: |
| return |
| self._ensure_buffer() |
| offset = 0 |
| while offset < len(tokens): |
| space = self.shard_size - self.token_count |
| take = min(space, len(tokens) - offset) |
| self.buffer[self.token_count:self.token_count + take] = tokens[offset:offset + take] |
| self.token_count += take |
| self.total_tokens += take |
| offset += take |
| if self.token_count == self.shard_size: |
| write_datafile_np(self._filename(), self.buffer) |
| self.shard_index += 1 |
| self.token_count = 0 |
|
|
| def close(self): |
| if self.buffer is not None and self.token_count > 0: |
| write_datafile_np(self._filename(), self.buffer[:self.token_count].copy()) |
| self.shard_index += 1 |
| self.token_count = 0 |
|
|
|
|
| def serialize_messages(messages): |
| parts = [] |
| for message in messages: |
| role = str(message.get("role", "unknown")) |
| content = str(message.get("content", "")) |
| if content: |
| parts.append(f"<|{role}|>\n{content}") |
| return "\n".join(parts) |
|
|
|
|
| def iter_texts(path, input_format, text_key, messages_key, limit_docs): |
| skipped = 0 |
| yielded = 0 |
| with open(path, "r", encoding="utf-8") as f: |
| for line_idx, line in enumerate(f): |
| if limit_docs is not None and yielded >= limit_docs: |
| break |
| try: |
| obj = json.loads(line) |
| except json.JSONDecodeError: |
| skipped += 1 |
| continue |
|
|
| if input_format == "text": |
| text = obj.get(text_key, "") |
| elif input_format == "messages": |
| text = serialize_messages(obj.get(messages_key, [])) |
| else: |
| if text_key in obj: |
| text = obj.get(text_key, "") |
| elif messages_key in obj: |
| text = serialize_messages(obj.get(messages_key, [])) |
| else: |
| text = "" |
|
|
| if not isinstance(text, str): |
| text = str(text) |
| if not text.strip(): |
| skipped += 1 |
| continue |
| yield line_idx, text |
| yielded += 1 |
|
|
| if skipped: |
| print(f"skipped {skipped:,} empty or invalid records") |
|
|
|
|
| def goes_to_val(doc_index, val_fraction, seed): |
| if val_fraction <= 0: |
| return False |
| key = f"{seed}:{doc_index}".encode("utf-8") |
| digest = hashlib.blake2b(key, digest_size=8).digest() |
| value = int.from_bytes(digest, "little") / 2**64 |
| return value < val_fraction |
|
|
|
|
| def ensure_no_existing_bins(output_dir, dataset_name, overwrite): |
| patterns = [ |
| os.path.join(output_dir, f"{dataset_name}_train_*.bin"), |
| os.path.join(output_dir, f"{dataset_name}_val_*.bin"), |
| ] |
| existing = [path for pattern in patterns for path in glob.glob(pattern)] |
| if existing and not overwrite: |
| raise SystemExit( |
| f"Refusing to overwrite {len(existing)} existing .bin files in {output_dir}. " |
| "Pass --overwrite to replace them." |
| ) |
| for path in existing: |
| os.remove(path) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="Preprocess code-message JSONL for llm.c GPT-2 training") |
| parser.add_argument("--input", required=True, help="Input JSONL file") |
| parser.add_argument("--output_dir", default=None, help="Output directory for .bin shards") |
| parser.add_argument("--dataset_name", default="code_messages", help="Prefix for output shard files") |
| parser.add_argument("--format", choices=["auto", "text", "messages"], default="auto", help="Input JSONL schema") |
| parser.add_argument("--text_key", default="text", help="Text field name for text JSONL") |
| parser.add_argument("--messages_key", default="messages", help="Messages field name for chat JSONL") |
| parser.add_argument("--shard_size", type=int, default=100_000_000, help="Tokens per output shard") |
| parser.add_argument("--val_fraction", type=float, default=0.001, help="Doc fraction to reserve for validation") |
| parser.add_argument("--seed", type=int, default=1337, help="Seed for deterministic validation split") |
| parser.add_argument("--workers", type=int, default=min(16, max(1, (os.cpu_count() or 2) - 2))) |
| parser.add_argument("--chunksize", type=int, default=16, help="Multiprocessing chunksize") |
| parser.add_argument("--limit_docs", type=int, default=None, help="Only preprocess this many docs, for smoke tests") |
| parser.add_argument("--total_docs", type=int, default=None, help="Optional tqdm total") |
| parser.add_argument("--overwrite", action="store_true", help="Delete existing output shards first") |
| args = parser.parse_args() |
|
|
| if not (0.0 <= args.val_fraction < 1.0): |
| raise SystemExit("--val_fraction must be in [0, 1)") |
| if args.shard_size <= 0: |
| raise SystemExit("--shard_size must be positive") |
|
|
| script_dir = Path(__file__).resolve().parent |
| output_dir = Path(args.output_dir) if args.output_dir else script_dir / args.dataset_name |
| output_dir.mkdir(parents=True, exist_ok=True) |
| ensure_no_existing_bins(str(output_dir), args.dataset_name, args.overwrite) |
|
|
| train_writer = ShardWriter(output_dir, f"{args.dataset_name}_train", args.shard_size) |
| val_writer = ShardWriter(output_dir, f"{args.dataset_name}_val", args.shard_size) |
|
|
| total = args.total_docs |
| if args.limit_docs is not None: |
| total = args.limit_docs if total is None else min(total, args.limit_docs) |
|
|
| with mp.Pool(args.workers, initializer=init_worker) as pool: |
| records = iter_texts(args.input, args.format, args.text_key, args.messages_key, args.limit_docs) |
| token_iter = pool.imap(tokenize_record, records, chunksize=args.chunksize) |
| for doc_idx, tokens in tqdm(token_iter, total=total, unit="docs"): |
| if goes_to_val(doc_idx, args.val_fraction, args.seed): |
| val_writer.write(tokens) |
| else: |
| train_writer.write(tokens) |
|
|
| train_writer.close() |
| val_writer.close() |
| print(f"train tokens: {train_writer.total_tokens:,}") |
| print(f"val tokens: {val_writer.total_tokens:,}") |
| print(f"wrote shards under: {output_dir}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|