#!/usr/bin/env python3 """ 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|>" # filler so id 32000 stays free for [MASK] CHAT_TOKENS = [IM_START, IM_END] + [f"<|r{i}|>" for i in range(1, 8)] # rainbow pads BASE_VOCAB = 32000 # Supra tokenizer entry count (ids 0..31999) 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": # The actual SFT set for SmolLM2-135M-Instruct: 484K short, high-quality # conversations designed for <1B models (no function calling, no # advanced math). Cap with --max-examples (e.g. 60000) for a 150M model. 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 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 # Truncate from the front (history), keep the target response intact 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 = [] # (start, end) per message, in token space 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 # Target turn = the LAST assistant response (LLaDA-MoE style) a0, a1 = assistant_ranges[-1] if a1 <= a0: return None # Cap the response length, always keeping the trailing <|im_end|> 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 # Truncate from the front (history), keep the target response intact 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 = [] # (roles, segs) per conversation math_indices = set() for name in args.datasets.split(","): name = name.strip() print(f"[*] Loading dataset: {name}") if name == "math": # Upsample: several passes over the same pairs, each pass re-rolling # phrasings (fresh seed). Math repeats are NOT deduped: identical # copies are intentional extra gradient steps (memorization needs # repetition, not variety). 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()