File size: 3,624 Bytes
6eb0505 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 | """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() |