| """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() |