"""Merge grug-think + grug-think-v3-10k, tokenize, filter >64k, save to disk.""" import argparse import json import os import glob import pandas as pd from datasets import Dataset from transformers import AutoTokenizer def load_jsonl_dir(path): files = glob.glob(os.path.join(path, "**", "*.jsonl"), recursive=True) rows = [] for f in files: with open(f, "r", encoding="utf-8") as fh: for line in fh: line = line.strip() if line: rows.append(json.loads(line)) return rows def normalize(row): msgs = row.get("messages") or row.get("conversations") or [] tools = row.get("tools") or None return {"messages": msgs, "tools": tools, "source": row.get("source", "unknown")} def main(): ap = argparse.ArgumentParser() ap.add_argument("--grug-think-dir", default="/data/grug-think") ap.add_argument("--grug10k-dir", default="/data/grug10k") ap.add_argument("--out", default="/workspace/data_merged") ap.add_argument("--model", default="deepseek-ai/DeepSeek-V4-Flash-0731") ap.add_argument("--max-seq-len", type=int, default=65536) args = ap.parse_args() print(f"[data_prep] loading grug-think from {args.grug_think_dir}", flush=True) rows1 = load_jsonl_dir(args.grug_think_dir) print(f"[data_prep] grug-think rows: {len(rows1)}", flush=True) print(f"[data_prep] loading grug-think-v3-10k from {args.grug10k_dir}", flush=True) rows2 = load_jsonl_dir(args.grug10k_dir) print(f"[data_prep] grug-think-v3-10k rows: {len(rows2)}", flush=True) all_rows = [normalize(r) for r in (rows1 + rows2)] print(f"[data_prep] merged rows: {len(all_rows)}", flush=True) tok = AutoTokenizer.from_pretrained(args.model, trust_remote_code=True) if tok.pad_token is None: tok.pad_token = tok.eos_token def render_text(example): parts = [] if example.get("tools"): parts.append(json.dumps(example["tools"], default=list)) for m in example["messages"]: role = m.get("role", "") content = m.get("content", "") parts.append(f"<|{role}|>{content}") return {"text": "".join(parts)} ds = Dataset.from_list(all_rows) ds = ds.map(render_text, num_proc=8, desc="rendering text") def count_tokens(example): ids = tok(example["text"], add_special_tokens=True, truncation=False)["input_ids"] return {"token_len": len(ids)} ds = ds.map(count_tokens, num_proc=8, desc="counting tokens") before = len(ds) ds = ds.filter(lambda x: x["token_len"] <= args.max_seq_len, num_proc=8, desc="filtering >max_seq_len") after = len(ds) print(f"[data_prep] filtered > {args.max_seq_len} tokens: {before} -> {after} (dropped {before - after})", flush=True) lens = ds["token_len"] lens.sort() n = len(lens) stats = { "n_samples": n, "n_tokens_total": sum(lens), "seq_len_min": lens[0] if n else 0, "seq_len_p50": lens[n // 2] if n else 0, "seq_len_p90": lens[int(n * 0.9)] if n else 0, "seq_len_p99": lens[int(n * 0.99)] if n else 0, "seq_len_max": lens[-1] if n else 0, } print(f"[data_prep] stats: {json.dumps(stats, indent=2)}", flush=True) ds = ds.remove_columns(["token_len"]) os.makedirs(args.out, exist_ok=True) ds.save_to_disk(args.out) with open(os.path.join(args.out, "dataset_stats.json"), "w") as fh: json.dump(stats, fh, indent=2) print(f"[data_prep] saved to {args.out} ({n} samples)", flush=True) if __name__ == "__main__": main()