Files
abap-llm/train/prepare.py

114 lines
5.5 KiB
Python

"""Stage 1 data preparation: real token counts, length filter, split by document, report.
train/.venv/bin/python train/prepare.py [--corpus PATH] [--max-len 16384] [--seed 20261003]
Output: train/data/train.jsonl, valid.jsonl, test.jsonl (copy of valid: `mlx_lm.lora --test` needs a test file),
train/data/report.md, train/data/stats.json. Only the tokenizer is loaded (no model).
"""
import argparse
import collections
import json
import os
import random
import re
from transformers import AutoTokenizer
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
OUT = os.path.join(ROOT, "train", "data")
def pct(v, p):
return v[min(len(v) - 1, int(len(v) * p))]
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--corpus", default=os.path.expanduser("~/projects/abap-llm/corpus/corpus.jsonl"))
ap.add_argument("--model", default=os.path.expanduser("~/models/Qwen3.8-27B-4bit"))
ap.add_argument("--max-len", type=int, default=16384)
ap.add_argument("--seed", type=int, default=20261003)
ap.add_argument("--valid-frac", type=float, default=0.05)
a = ap.parse_args()
tok = AutoTokenizer.from_pretrained(a.model)
docs = [json.loads(l) for l in open(a.corpus)]
for i, d in enumerate(docs):
d["_id"] = i
d["real_tokens"] = len(tok(d["text"], add_special_tokens=False)["input_ids"])
est = sum(d["tokens"] for d in docs)
real = sum(d["real_tokens"] for d in docs)
kept = [d for d in docs if d["real_tokens"] <= a.max_len]
removed = [d for d in docs if d["real_tokens"] > a.max_len]
# Family = same demo in several release versions (__V755 ...) or parts of one text (__P01 ...). A family
# is never split between train and valid (near-duplicates would make the valid loss too good).
fam = collections.defaultdict(list)
for d in kept:
fam[re.sub(r"__(V\d+|P\d+)", "", d["objects"][0] if d["objects"] else str(d["_id"]))].append(d)
keys = sorted(fam)
random.Random(a.seed).shuffle(keys)
n_valid, valid, train = max(1, round(len(kept) * a.valid_frac)), [], []
for k in keys:
(valid if len(valid) < n_valid else train).extend(fam[k])
os.makedirs(OUT, exist_ok=True)
for name, part in (("train", train), ("valid", valid), ("test", valid)):
with open(os.path.join(OUT, name + ".jsonl"), "w") as f:
for d in part:
f.write(json.dumps({"text": d["text"]}, ensure_ascii=False) + "\n")
def summ(ds):
v = sorted(d["real_tokens"] for d in ds)
return {"docs": len(ds), "tokens": sum(v), "min": v[0], "p50": pct(v, .5), "p90": pct(v, .9),
"p99": pct(v, .99), "max": v[-1]} if v else {"docs": 0, "tokens": 0}
by_type = collections.defaultdict(lambda: [0, 0])
for d in kept:
k = "+".join(d["types"]) if len(d["types"]) <= 2 else "mixed(%d)" % len(d["types"])
by_type[k][0] += 1
by_type[k][1] += d["real_tokens"]
bins = [512, 1024, 2048, 4096, 8192, 16384, 10 ** 9]
hist = collections.Counter()
for d in docs:
for b in bins:
if d["real_tokens"] <= b:
hist[b] += 1
break
stats = {"corpus": a.corpus, "seed": a.seed, "max_len": a.max_len, "estimate_tokens": est,
"real_tokens": real, "ratio_real_to_estimate": round(real / est, 3), "all": summ(docs),
"kept": summ(kept), "train": summ(train), "valid": summ(valid),
"removed": [{"objects": d["objects"][:3], "tokens": d["real_tokens"]} for d in removed],
"by_type": {k: {"docs": v[0], "tokens": v[1]} for k, v in by_type.items()},
"source": sorted({d.get("source", "?") for d in docs}),
"iterations_2_epochs": len(train) * 2}
json.dump(stats, open(os.path.join(OUT, "stats.json"), "w"), indent=1)
L = ["# Stage 1 data report", "",
f"Corpus: `{a.corpus}` (source {', '.join(stats['source'])}, Apache-2.0, see corpus/out/ATTRIBUTION.md).",
f"Tokenizer: `{a.model}`. Seed {a.seed}. max_seq_length {a.max_len}.",
f"Split by family (release versions and parts of one document stay together): {len(fam)} families.", "",
"| | docs | tokens |", "|---|---|---|",
f"| corpus | {len(docs)} | {real} (estimate chars/4: {est}, ratio {stats['ratio_real_to_estimate']}) |",
f"| removed (> {a.max_len}) | {len(removed)} | {sum(d['real_tokens'] for d in removed)} |",
f"| train | {len(train)} | {stats['train']['tokens']} |",
f"| valid (= test file) | {len(valid)} | {stats['valid']['tokens']} |", "",
f"Iterations for 2 epochs (batch 1): {len(train) * 2}.", "",
"## Token distribution per document (real tokenizer)", "", "| set | min | p50 | p90 | p99 | max |", "|---|---|---|---|---|---|"]
for n in ("all", "kept", "train", "valid"):
s = stats[n]
L.append(f"| {n} | {s['min']} | {s['p50']} | {s['p90']} | {s['p99']} | {s['max']} |")
L += ["", "| tokens per document up to | docs |", "|---|---|"]
L += [f"| {b if b < 10 ** 9 else 'more'} | {hist[b]} |" for b in bins]
L += ["", "## By object type (kept documents)", "", "| types | docs | tokens |", "|---|---|---|"]
for k, v in sorted(by_type.items(), key=lambda kv: -kv[1][1]):
L.append(f"| {k} | {v[0]} | {v[1]} |")
L += ["", "## Removed documents", ""]
L += [f"- {', '.join(d['objects'][:2])}: {d['real_tokens']} tokens" for d in removed] or ["None."]
open(os.path.join(OUT, "report.md"), "w").write("\n".join(L) + "\n")
print("\n".join(L))
if __name__ == "__main__":
main()