JEV / code /scripts /prepare_data.py
cloudyu's picture
Upload JEV v0.7.0 (Qwen3.5-9B distilled from Jev 1.13)
b2f3bf4 verified
Raw
History Blame Contribute Delete
5.65 kB
#!/usr/bin/env python3
"""M1: jsonl -> parquet (+ n_tokens, is_uniform, n_options) and the data audit report (DESIGN §3.3).
python3 scripts/prepare_data.py --raw data/raw --out data --tokenizer /root/models/Qwen3.5-9B \
--max-seq-len 1024 --report reports/data_audit.md
"""
from __future__ import annotations
import argparse
import json
import os
import sys
import time
import numpy as np
import pandas as pd
sys.path.insert(0, os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "src"))
from jev_judge.data import SPLITS, is_uniform, validate_row # noqa: E402
from jev_judge.template import n_options_for, render # noqa: E402
def read_jsonl(path: str) -> list[dict]:
rows = []
with open(path) as f:
for line in f:
line = line.strip()
if line:
rows.append(json.loads(line))
return rows
def token_lengths(tok, texts: list[str], bs: int = 4096) -> np.ndarray:
out = np.empty(len(texts), dtype=np.int32)
for s in range(0, len(texts), bs):
enc = tok(texts[s : s + bs], add_special_tokens=False, return_length=True)
out[s : s + bs] = np.asarray(enc["length"], dtype=np.int32)
return out
def pct(a: np.ndarray, p: float) -> float:
return float(np.percentile(a, p)) if len(a) else float("nan")
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--raw", default="data/raw")
ap.add_argument("--out", default="data")
ap.add_argument("--tokenizer", default="/root/models/Qwen3.5-9B")
ap.add_argument("--max-seq-len", type=int, default=1024)
ap.add_argument("--report", default="reports/data_audit.md")
args = ap.parse_args()
from transformers import AutoTokenizer
tok = AutoTokenizer.from_pretrained(args.tokenizer)
os.makedirs(args.out, exist_ok=True)
os.makedirs(os.path.dirname(args.report) or ".", exist_ok=True)
lines = [f"# Data audit — `jev-distill-corpus-v3`", "",
f"tokenizer: `{args.tokenizer}` · max_seq_len: {args.max_seq_len} · template: bare-v1 (whole-string tokenization)", ""]
summary_rows = []
d1_rows = []
all_train_tokens = 0
for split in SPLITS:
t0 = time.time()
rows = read_jsonl(os.path.join(args.raw, f"{split}.jsonl"))
violations = []
for i, r in enumerate(rows):
errs = validate_row(r)
if errs:
violations.append((i, r.get("id"), errs))
good = [r for i, r in enumerate(rows) if not validate_row(r)] if violations else rows
df = pd.DataFrame(good)
df["target"] = df["target"].apply(lambda t: np.asarray(t, dtype=np.float32).tolist())
df["options"] = df["options"].apply(lambda o: [str(x) for x in o])
df["is_uniform"] = df["target"].apply(is_uniform)
df["n_options"] = [n_options_for(k, o) for k, o in zip(df["kind"], df["options"])]
texts = [render(k, s, q, o) for k, s, q, o in zip(df["kind"], df["state"], df["question"], df["options"])]
df["n_tokens"] = token_lengths(tok, texts)
df.to_parquet(os.path.join(args.out, f"{split}.parquet"), index=False)
n_tok = df["n_tokens"].to_numpy()
trunc = int((n_tok > args.max_seq_len).sum())
if split == "train":
all_train_tokens = int(n_tok.sum())
summary_rows.append(
f"| {split} | {len(rows):,} | {len(violations)} | {n_tok.mean():.1f} | {pct(n_tok,50):.0f} | {pct(n_tok,95):.0f} | {pct(n_tok,99):.0f} | {n_tok.max()} | {trunc} | {time.time()-t0:.0f}s |"
)
# D1 table per source×kind
for (src, kind), g in df.groupby(["source", "kind"]):
u = int(g["is_uniform"].sum())
d1_rows.append(f"| {split} | {src} | {kind} | {len(g):,} | {u:,} | {100*u/len(g):.1f}% |")
if violations:
lines.append(f"### schema violations in `{split}` ({len(violations)}) — first 10")
for i, rid, errs in violations[:10]:
lines.append(f"- row {i} id={rid}: {errs}")
lines.append("")
print(f"[{split}] rows={len(rows):,} violations={len(violations)} mean_tok={n_tok.mean():.1f} p99={pct(n_tok,99):.0f} max={n_tok.max()} trunc={trunc}", flush=True)
lines += ["## Splits", "",
"| split | rows | schema violations | mean tok | p50 | p95 | p99 | max | > max_seq_len | time |",
"|---|---|---|---|---|---|---|---|---|---|", *summary_rows, "",
f"train tokens / epoch ≈ **{all_train_tokens/1e6:.1f}M**; 2 epochs ≈ **{2*all_train_tokens/1e6:.0f}M**", ""]
lines += ["## D1 — exactly-uniform teacher labels by split × source × kind", "",
"| split | source | kind | rows | uniform | share |", "|---|---|---|---|---|---|", *d1_rows, ""]
# kind/source mix in train
tr = pd.read_parquet(os.path.join(args.out, "train.parquet"), columns=["kind", "source", "family", "n_options"])
lines += ["## Train mix", "", "| kind | rows | share |", "|---|---|---|"]
for k, n in tr["kind"].value_counts().items():
lines.append(f"| {k} | {n:,} | {100*n/len(tr):.1f}% |")
lines += ["", "| source | rows | share |", "|---|---|---|"]
for k, n in tr["source"].value_counts().items():
lines.append(f"| {k} | {n:,} | {100*n/len(tr):.1f}% |")
ch = tr[tr["kind"] == "choice"]["n_options"].value_counts().sort_index()
lines += ["", "choice n_options histogram (train): " + ", ".join(f"{int(k)}:{v:,}" for k, v in ch.items()), ""]
with open(args.report, "w") as f:
f.write("\n".join(lines))
print("report ->", args.report)
if __name__ == "__main__":
main()