| |
| """ |
| prepare_data.py: Download HuggingFaceH4/no_robots and pre-tokenize it as ChatML |
| for MetaDiffusion chat SFT. |
| |
| Output (in --data-dir): |
| train.pt list of {"input_ids": LongTensor, "assistant_start": int, |
| "assistant_end": int} |
| val.pt same, held-out |
| tokenizer/ Supra tokenizer with ChatML + rainbow tokens added |
| stats.json counts and length stats |
| |
| The assistant region (content + trailing <|im_end|>) is the only part that |
| will be masked during training; everything before it is visible prompt. |
| |
| Usage: |
| python3 prepare_data.py --model-path ../hf_release --data-dir data/no_robots_chatml |
| """ |
|
|
| import argparse |
| import hashlib |
| import json |
| import random |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| from datasets import load_dataset |
| from transformers import AutoTokenizer |
|
|
| IM_START = "<|im_start|>" |
| IM_END = "<|im_end|>" |
| RESERVED = "<|reserved|>" |
| CHAT_TOKENS = [IM_START, IM_END] + [f"<|r{i}|>" for i in range(1, 8)] |
| BASE_VOCAB = 32000 |
|
|
|
|
| def add_chat_tokens(tokenizer): |
| """Add chat + rainbow tokens at ids 32001..32009. |
| |
| The base tokenizer has 32000 entries, so the first added token would take |
| id 32000, which is the diffusion [MASK] id. A reserved filler token takes |
| that slot first; mask id 32000 must never exist in the tokenizer. |
| """ |
| if len(tokenizer) == BASE_VOCAB: |
| tokenizer.add_special_tokens({"additional_special_tokens": [RESERVED]}) |
| n = tokenizer.add_special_tokens( |
| {"additional_special_tokens": CHAT_TOKENS} |
| ) |
| im_start = tokenizer.convert_tokens_to_ids(IM_START) |
| im_end = tokenizer.convert_tokens_to_ids(IM_END) |
| assert im_start == 32001, f"im_start id {im_start} != 32001 (collides with [MASK])" |
| assert im_end == 32002, f"im_end id {im_end} != 32002" |
| print(f"[*] Added {n} special tokens, vocab now {len(tokenizer)}") |
| print(f"[*] im_start={im_start} im_end={im_end} " |
| f"rainbow={[tokenizer.convert_tokens_to_ids(f'<|r{i}|>') for i in range(1, 8)]}") |
| return n |
|
|
|
|
| def format_segment(role: str, content: str) -> str: |
| return f"{IM_START}{role}\n{content}{IM_END}" |
|
|
|
|
| def load_messages(dataset_name: str): |
| """Return a list of message lists for one dataset name.""" |
| if dataset_name == "no_robots": |
| ds = load_dataset("HuggingFaceH4/no_robots", split="train") |
| return [list(row["messages"]) for row in ds] |
| if dataset_name == "alpaca": |
| ds = load_dataset("yahma/alpaca-cleaned", split="train") |
| out = [] |
| for row in ds: |
| user = row["instruction"] |
| if row.get("input"): |
| user += f"\n\n{row['input']}" |
| out.append([{"role": "user", "content": user}, |
| {"role": "assistant", "content": row["output"]}]) |
| return out |
| if dataset_name == "dolly": |
| ds = load_dataset("databricks/databricks-dolly-15k", split="train") |
| out = [] |
| for row in ds: |
| user = row["instruction"] |
| if row.get("context"): |
| user += f"\n\n{row['context']}" |
| out.append([{"role": "user", "content": user}, |
| {"role": "assistant", "content": row["response"]}]) |
| return out |
| if dataset_name == "smol-smoltalk": |
| |
| |
| |
| ds = load_dataset("HuggingFaceTB/smol-smoltalk", split="train") |
| return [list(row["messages"]) for row in ds] |
| if dataset_name == "math": |
| return load_math_data() |
| raise ValueError(f"Unknown dataset: {dataset_name} " |
| f"(choose from: no_robots, alpaca, dolly, smol-smoltalk, math)") |
|
|
| MATH_QUESTION_TEMPLATES = { |
| "add": ["What is {a} + {b}?", "What is {a} plus {b}?", "Add {a} and {b}.", |
| "What does {a} + {b} equal?"], |
| "sub": ["What is {a} - {b}?", "What is {a} minus {b}?", "Subtract {b} from {a}.", |
| "What does {a} - {b} equal?"], |
| "mul": ["What is {a} × {b}?", "What is {a} times {b}?", "Multiply {a} by {b}.", |
| "What does {a} × {b} equal?"], |
| "div": ["What is {a} ÷ {b}?", "What is {a} divided by {b}?", "Divide {a} by {b}.", |
| "What does {a} ÷ {b} equal?"], |
| } |
| MATH_ANSWER_TEMPLATES = ["The answer is {r}.", "It is {r}.", "{r}"] |
|
|
|
|
| def load_math_data(seed: int = 42): |
| """Exhaustive basic-arithmetic QA pairs (add/sub/mul/div), ChatML messages. |
| |
| A 150M model learns arithmetic by memorization, so cover EVERY pair in a |
| small range rather than sampling: add a<=b in 1..99, sub b<a in 1..99, |
| times tables 1..12, exact divisions 1..12. |
| """ |
| rng = random.Random(seed) |
| pairs = [] |
| for a in range(1, 100): |
| for b in range(a, 100): |
| pairs.append(("add", a, b, a + b)) |
| for a in range(2, 100): |
| for b in range(1, a): |
| pairs.append(("sub", a, b, a - b)) |
| for a in range(1, 13): |
| for b in range(1, 13): |
| pairs.append(("mul", a, b, a * b)) |
| for b in range(1, 13): |
| for q in range(1, 13): |
| pairs.append(("div", b * q, b, q)) |
| rng.shuffle(pairs) |
|
|
| out = [] |
| for op, a, b, r in pairs: |
| question = rng.choice(MATH_QUESTION_TEMPLATES[op]).format(a=a, b=b) |
| answer = rng.choice(MATH_ANSWER_TEMPLATES).format(r=r) |
| out.append([{"role": "user", "content": question}, |
| {"role": "assistant", "content": answer}]) |
| print(f" (synthetic arithmetic: {len(out)} pairs, " |
| f"add {sum(1 for p in pairs if p[0]=='add')}, " |
| f"sub {sum(1 for p in pairs if p[0]=='sub')}, " |
| f"mul {sum(1 for p in pairs if p[0]=='mul')}, " |
| f"div {sum(1 for p in pairs if p[0]=='div')})") |
| return out |
|
|
|
|
| def example_key(input_ids_tensor): |
| """Compact dedup key: SHA-256 of token ids. |
| |
| Tuple keys for 484K x 600-token conversations cost ~5 GB of set memory; |
| digests cost ~50 MB. (Hashing was never the bottleneck; per-segment |
| tokenizer.encode() calls were.) |
| """ |
| return hashlib.sha256( |
| input_ids_tensor.numpy().astype(np.uint32).tobytes() |
| ).digest() |
|
|
|
|
| def build_conv_segments(messages): |
| """Split a conversation into ChatML segment strings + roles.""" |
| segs, roles = [], [] |
| for m in messages: |
| role = m.get("role", "user") |
| content = m.get("content", "") |
| if not content.strip(): |
| continue |
| segs.append(format_segment(role, content)) |
| roles.append(role) |
| return segs, roles |
|
|
|
|
| def reconstruct_example(seg_ids, roles, tokenizer, max_len, max_resp): |
| """Rebuild a conversation from pre-encoded segments (tokenize_example |
| logic, but with tokenization already done in batch).""" |
| ids = [] |
| assistant_ranges = [] |
| for role, seg in zip(roles, seg_ids): |
| start = len(ids) |
| ids.extend(seg) |
| if role == "assistant": |
| assistant_ranges.append((start, len(ids))) |
|
|
| if not ids or not assistant_ranges: |
| return None |
| a0, a1 = assistant_ranges[-1] |
| if a1 <= a0: |
| return None |
|
|
| |
| if a1 - a0 > max_resp: |
| im_end_id = tokenizer.convert_tokens_to_ids(IM_END) |
| ids = ids[:a0] + ids[a0:a0 + max_resp - 1] + [im_end_id] |
| a1 = a0 + max_resp |
|
|
| |
| resp = ids[a0:a1] |
| if len(resp) > max_len: |
| resp = resp[:max_len] |
| a1 = a0 + len(resp) |
| hist = ids[:a0] |
| room = max_len - len(resp) |
| if len(hist) > room: |
| hist = hist[len(hist) - room:] if room > 0 else [] |
| ids = hist + resp |
| return { |
| "input_ids": torch.tensor(ids, dtype=torch.long), |
| "assistant_start": len(hist), |
| "assistant_end": len(ids), |
| } |
|
|
|
|
| def tokenize_and_split(convs, math_indices, tokenizer, val_size, max_len, |
| max_resp, seed=42): |
| """Batched tokenize + reconstruct + dedup + val split. |
| |
| convs: list of (roles, segs) from build_conv_segments. |
| math_indices: conversation indices exempt from dedup (math repeats are |
| intentional). Returns (train_examples, val_examples, skipped, dupes). |
| """ |
| all_seg_strs = [s for _, segs in convs for s in segs] |
| print(f"[*] Encoding {len(all_seg_strs)} segments (batched)...") |
| encoded = [] |
| CHUNK = 200_000 |
| for i in range(0, len(all_seg_strs), CHUNK): |
| chunk = all_seg_strs[i:i + CHUNK] |
| encoded.extend(tokenizer(chunk, add_special_tokens=False)["input_ids"]) |
|
|
| rng = random.Random(seed) |
| idxs = list(range(len(convs))) |
| rng.shuffle(idxs) |
| val_idxs = set(idxs[:val_size]) |
|
|
| train_examples, val_examples = [], [] |
| skipped = dupes = 0 |
| seen = set() |
| ptr = 0 |
| for i, (roles, segs) in enumerate(convs): |
| seg_ids = encoded[ptr:ptr + len(segs)] |
| ptr += len(segs) |
| ex = reconstruct_example(seg_ids, roles, tokenizer, max_len, max_resp) |
| if ex is None: |
| skipped += 1 |
| continue |
| if i not in math_indices: |
| key = example_key(ex["input_ids"]) |
| if key in seen: |
| dupes += 1 |
| continue |
| seen.add(key) |
| (val_examples if i in val_idxs else train_examples).append(ex) |
| return train_examples, val_examples, skipped, dupes |
|
|
|
|
| def save_dataset(out_dir, train_examples, val_examples, tokenizer, |
| skipped, dupes): |
| """Save train.pt / val.pt / stats.json and print the summary.""" |
| out = Path(out_dir) |
| torch.save({"examples": train_examples}, out / "train.pt") |
| torch.save({"examples": val_examples}, out / "val.pt") |
|
|
| lens = [e["input_ids"].numel() for e in train_examples] |
| stats = { |
| "train": len(train_examples), |
| "val": len(val_examples), |
| "skipped": skipped, |
| "dupes": dupes, |
| "avg_tokens": sum(lens) / len(lens) if lens else 0, |
| "max_tokens": max(lens) if lens else 0, |
| "vocab_size": len(tokenizer), |
| "im_start_id": tokenizer.convert_tokens_to_ids(IM_START), |
| "im_end_id": tokenizer.convert_tokens_to_ids(IM_END), |
| "rainbow_ids": [tokenizer.convert_tokens_to_ids(f"<|r{i}|>") for i in range(1, 8)], |
| "mask_token_id": 32000, |
| } |
| with open(out / "stats.json", "w") as f: |
| json.dump(stats, f, indent=2) |
|
|
| print(f"[*] Done: {stats['train']} train / {stats['val']} val " |
| f"(skipped {skipped}, dupes {dupes})") |
| print(f"[*] Avg tokens per example: {stats['avg_tokens']:.0f} (max {stats['max_tokens']})") |
| print(f"[*] im_start={stats['im_start_id']} im_end={stats['im_end_id']} " |
| f"rainbow={stats['rainbow_ids']}") |
| return stats |
|
|
|
|
| def tokenize_example(tokenizer, messages, max_len: int, max_resp: int = 256): |
| """Tokenize a conversation segment-wise; return ids + assistant bounds. |
| |
| The target (last assistant response) is capped at `max_resp` tokens |
| including its trailing <|im_end|>: long web-text targets teach rambling, |
| and the terminator must stay in the target so the model learns to emit it. |
| """ |
| ids = [] |
| assistant_ranges = [] |
| for m in messages: |
| role = m.get("role", "user") |
| content = m.get("content", "") |
| if not content.strip(): |
| continue |
| seg = format_segment(role, content) |
| seg_ids = tokenizer.encode(seg, add_special_tokens=False) |
| start = len(ids) |
| ids.extend(seg_ids) |
| if role == "assistant": |
| assistant_ranges.append((start, len(ids))) |
|
|
| if not ids or not assistant_ranges: |
| return None |
|
|
| |
| a0, a1 = assistant_ranges[-1] |
| if a1 <= a0: |
| return None |
|
|
| |
| if a1 - a0 > max_resp: |
| im_end_id = tokenizer.convert_tokens_to_ids(IM_END) |
| ids = ids[:a0] + ids[a0:a0 + max_resp - 1] + [im_end_id] |
| a1 = a0 + max_resp |
|
|
| |
| resp = ids[a0:a1] |
| if len(resp) > max_len: |
| resp = resp[:max_len] |
| a1 = a0 + len(resp) |
| hist = ids[:a0] |
| room = max_len - len(resp) |
| if len(hist) > room: |
| hist = hist[len(hist) - room:] if room > 0 else [] |
| ids = hist + resp |
| return { |
| "input_ids": torch.tensor(ids, dtype=torch.long), |
| "assistant_start": len(hist), |
| "assistant_end": len(ids), |
| } |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="Prepare no_robots as ChatML for MetaDiffusion SFT") |
| parser.add_argument("--model-path", default="../hf_release", help="Dir with tokenizer.json (base model)") |
| parser.add_argument("--data-dir", default="data/no_robots_chatml", help="Output dir") |
| parser.add_argument("--datasets", default="no_robots", |
| help="Comma list: no_robots, alpaca, dolly, smol-smoltalk, " |
| "math (e.g. no_robots,alpaca,dolly,smol-smoltalk,math)") |
| parser.add_argument("--max-examples", type=int, default=0, |
| help="Cap per downloaded dataset (0 = no cap). smol-smoltalk " |
| "is 484K; use ~60000 for a 150M model. Does not apply " |
| "to synthetic math (exhaustive coverage is the point).") |
| parser.add_argument("--val-size", type=int, default=500, help="Held-out examples") |
| parser.add_argument("--max-len", type=int, default=1024, help="Max tokens per example") |
| parser.add_argument("--max-resp-tokens", type=int, default=256, |
| help="Cap on target response tokens (keeps <|im_end|>)") |
| parser.add_argument("--math-repeat", type=int, default=1, |
| help="Upsample synthetic math: N passes over all pairs " |
| "(different phrasings per pass; 3-4 helps a 150M " |
| "model memorize the mapping)") |
| parser.add_argument("--seed", type=int, default=42) |
| args = parser.parse_args() |
|
|
| out = Path(args.data_dir) |
| out.mkdir(parents=True, exist_ok=True) |
|
|
| print(f"[*] Loading tokenizer from {args.model_path}") |
| tokenizer = AutoTokenizer.from_pretrained(args.model_path) |
| add_chat_tokens(tokenizer) |
| tokenizer.save_pretrained(out / "tokenizer") |
|
|
| convs = [] |
| math_indices = set() |
| for name in args.datasets.split(","): |
| name = name.strip() |
| print(f"[*] Loading dataset: {name}") |
| if name == "math": |
| |
| |
| |
| |
| start = len(convs) |
| for rep in range(args.math_repeat): |
| msgs = load_math_data(seed=args.seed + rep) |
| for messages in msgs: |
| segs, roles = build_conv_segments(messages) |
| convs.append((roles, segs)) |
| print(f" pass {rep + 1}/{args.math_repeat}: {len(msgs)}") |
| math_indices.update(range(start, len(convs))) |
| else: |
| msgs = load_messages(name) |
| if args.max_examples > 0 and len(msgs) > args.max_examples: |
| rng_cap = random.Random(args.seed) |
| rng_cap.shuffle(msgs) |
| msgs = msgs[: args.max_examples] |
| print(f" capped to {len(msgs)}") |
| for messages in msgs: |
| segs, roles = build_conv_segments(messages) |
| convs.append((roles, segs)) |
| print(f" {len(msgs)} conversations") |
| print(f"[*] Total: {len(convs)} conversations") |
|
|
| train_examples, val_examples, skipped, dupes = tokenize_and_split( |
| convs, math_indices, tokenizer, args.val_size, args.max_len, |
| args.max_resp_tokens, args.seed) |
|
|
| save_dataset(out, train_examples, val_examples, tokenizer, skipped, dupes) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|