Stage 1 data: version dedup, splitting of long documents, family split; 748 iterations
Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_014aUaQeLnwbb1zTpN7kHeat
This commit is contained in:
315
train/prepare.py
315
train/prepare.py
@@ -1,12 +1,18 @@
|
||||
"""Stage 1 data preparation: real token counts, length filter, split by document, report.
|
||||
"""Stage 1 data preparation (Kral decisions 2026-10-03).
|
||||
|
||||
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).
|
||||
1. Real token counts (tokenizer of the base model).
|
||||
2. Version dedup: in each family keep the newest version; keep an older version only if its source differs
|
||||
from the newest by more than 5 % of lines (difflib).
|
||||
3. Documents longer than max_len are split with the real tokenizer: classes at ENDMETHOD boundaries, markdown
|
||||
at "##" headings (fallbacks: "###", then lines). Header lines are repeated in each piece.
|
||||
4. Split 95/5 by family (all versions and pieces of one family stay in one split).
|
||||
Output: train/data/{train,valid,test}.jsonl (test = copy of valid), report.md, stats.json. Only the tokenizer is loaded.
|
||||
"""
|
||||
import argparse
|
||||
import collections
|
||||
import difflib
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
@@ -16,12 +22,155 @@ from transformers import AutoTokenizer
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
OUT = os.path.join(ROOT, "train", "data")
|
||||
VER = re.compile(r"__V\d+")
|
||||
FAM = re.compile(r"__(V\d+|P\d+)")
|
||||
RANK = {"main": 100, "v816": 90, "v758": 80, "v757": 70, "v756": 60, "v755": 50}
|
||||
MARK = re.compile(r"^\* ---- .* ----$")
|
||||
|
||||
|
||||
def pct(v, p):
|
||||
return v[min(len(v) - 1, int(len(v) * p))]
|
||||
|
||||
|
||||
def version_of(d):
|
||||
parts = d["package_path"].split("/")
|
||||
return parts[1] if len(parts) > 1 else parts[0]
|
||||
|
||||
|
||||
def norm_lines(text):
|
||||
out = []
|
||||
for i, ln in enumerate(text.split("\n")):
|
||||
if i == 0 or MARK.match(ln):
|
||||
continue
|
||||
ln = VER.sub("", ln).rstrip()
|
||||
if ln.strip():
|
||||
out.append(ln)
|
||||
return out
|
||||
|
||||
|
||||
def diff_fraction(a, b):
|
||||
n = max(len(a), len(b))
|
||||
if n == 0:
|
||||
return 0.0
|
||||
m = sum(x.size for x in difflib.SequenceMatcher(None, a, b, autojunk=False).get_matching_blocks())
|
||||
return (n - m) / n
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- splitting
|
||||
|
||||
def split_blocks(text, kind):
|
||||
"""Cut the body into atomic blocks: after each ENDMETHOD. (code), before each '## ' heading (markdown)."""
|
||||
lines = text.split("\n")
|
||||
blocks, cur = [], []
|
||||
for ln in lines:
|
||||
if kind == "md" and ln.startswith("## ") and cur:
|
||||
blocks.append("\n".join(cur))
|
||||
cur = []
|
||||
cur.append(ln)
|
||||
if kind == "code" and re.match(r"^\s*ENDMETHOD\.", ln, re.I):
|
||||
blocks.append("\n".join(cur))
|
||||
cur = []
|
||||
if cur:
|
||||
blocks.append("\n".join(cur))
|
||||
return blocks
|
||||
|
||||
|
||||
def split_document(d, tok, max_len):
|
||||
"""Return a list of piece texts (each <= max_len tokens, as far as possible) and the number of hard splits."""
|
||||
text = d["text"]
|
||||
ntok = lambda s: len(tok(s, add_special_tokens=False)["input_ids"])
|
||||
kind = "md" if d["types"] == ["DOC"] else "code"
|
||||
lines = text.split("\n")
|
||||
head = lines[0]
|
||||
# sections: marker line starts a section; body until next marker
|
||||
secs, cur_marker, cur = [], None, []
|
||||
for ln in lines[1:]:
|
||||
if MARK.match(ln):
|
||||
secs.append((cur_marker, cur))
|
||||
cur_marker, cur = ln, []
|
||||
else:
|
||||
cur.append(ln)
|
||||
secs.append((cur_marker, cur))
|
||||
units = [] # (marker, block_text)
|
||||
for marker, body in secs:
|
||||
body_text = "\n".join(body).strip("\n")
|
||||
if not body_text:
|
||||
continue
|
||||
sub = split_blocks(body_text, kind)
|
||||
for b in sub:
|
||||
units.append((marker, b))
|
||||
hard = 0
|
||||
# oversize unit: fall back (### for markdown), then line cut
|
||||
fixed = []
|
||||
for marker, b in units:
|
||||
if ntok(b) <= max_len - 400:
|
||||
fixed.append((marker, b))
|
||||
continue
|
||||
pieces = []
|
||||
if kind == "md":
|
||||
cur_p = []
|
||||
for ln in b.split("\n"):
|
||||
if ln.startswith("### ") and cur_p:
|
||||
pieces.append("\n".join(cur_p))
|
||||
cur_p = []
|
||||
cur_p.append(ln)
|
||||
pieces.append("\n".join(cur_p))
|
||||
else:
|
||||
pieces = [b]
|
||||
for p in pieces:
|
||||
if ntok(p) <= max_len - 400:
|
||||
fixed.append((marker, p))
|
||||
continue
|
||||
hard += 1
|
||||
ls, buf = p.split("\n"), []
|
||||
for ln in ls:
|
||||
buf.append(ln)
|
||||
if ntok("\n".join(buf)) > max_len - 600:
|
||||
fixed.append((marker, "\n".join(buf[:-1])))
|
||||
buf = [ln]
|
||||
fixed.append((marker, "\n".join(buf)))
|
||||
# greedy packing; header repeated; state of the open implementation class tracked for code
|
||||
title = next((ln for ln in lines if ln.startswith("# ")), "") if kind == "md" else ""
|
||||
groups, cur, cur_tok, last_marker = [], [], 0, None
|
||||
for marker, b in fixed:
|
||||
t = ntok(b) + 20
|
||||
if cur and cur_tok + t > max_len - 300:
|
||||
groups.append(cur)
|
||||
cur, cur_tok = [], 0
|
||||
cur.append((marker, b))
|
||||
cur_tok += t
|
||||
groups.append(cur)
|
||||
n = len(groups)
|
||||
out = []
|
||||
impl = None # last 'CLASS x IMPLEMENTATION.' still open at the start of the next piece
|
||||
for gi, g in enumerate(groups):
|
||||
h = head if n == 1 else f"{head} [part {gi + 1}/{n}]"
|
||||
txt = [h, ""]
|
||||
if gi > 0 and title and not g[0][1].startswith("# "):
|
||||
pass
|
||||
prev_marker = None
|
||||
for marker, b in g:
|
||||
if marker and marker != prev_marker:
|
||||
txt.append(marker)
|
||||
prev_marker = marker
|
||||
if gi > 0 and b is g[0][1] and impl and kind == "code" and not re.match(r"^\s*CLASS\b", b, re.I):
|
||||
txt.append(impl)
|
||||
elif gi > 0 and b is g[0][1] and kind == "md" and title and b.splitlines()[0] != title:
|
||||
txt.append(title)
|
||||
txt.append("")
|
||||
txt.append(b)
|
||||
for ln in b.split("\n"):
|
||||
m = re.match(r"^\s*CLASS\s+\S+\s+IMPLEMENTATION\.", ln, re.I)
|
||||
if m:
|
||||
impl = ln.strip()
|
||||
if re.match(r"^\s*ENDCLASS\.", ln, re.I):
|
||||
impl = None
|
||||
if kind == "code" and impl and gi < n - 1:
|
||||
txt.append("ENDCLASS.")
|
||||
out.append("\n".join(txt).rstrip("\n") + "\n")
|
||||
return out, hard
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--corpus", default=os.path.expanduser("~/projects/abap-llm/corpus/corpus.jsonl"))
|
||||
@@ -29,23 +178,78 @@ def main():
|
||||
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)
|
||||
ap.add_argument("--diff", type=float, default=0.05)
|
||||
a = ap.parse_args()
|
||||
|
||||
tok = AutoTokenizer.from_pretrained(a.model)
|
||||
ntok = lambda s: len(tok(s, add_special_tokens=False)["input_ids"])
|
||||
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]
|
||||
d["real_tokens"] = ntok(d["text"])
|
||||
d["fam"] = VER.sub("", d["objects"][0]) if d["objects"] else str(i)
|
||||
d["ver"] = version_of(d)
|
||||
est, real0 = sum(d["tokens"] for d in docs), sum(d["real_tokens"] for d in docs)
|
||||
|
||||
# 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).
|
||||
# ---- version dedup (an object version = all records of one object in one branch)
|
||||
objv = collections.defaultdict(list)
|
||||
for d in docs:
|
||||
objv[(d["fam"], d["ver"])].append(d)
|
||||
byfam = collections.defaultdict(dict)
|
||||
for (f, v), ds in objv.items():
|
||||
byfam[f][v] = ds
|
||||
removed_versions, kept_versions, strict_extra = [], 0, 0
|
||||
keep_ids = set()
|
||||
n_multi = 0
|
||||
for f, vers in byfam.items():
|
||||
order = sorted(vers, key=lambda v: -RANK.get(v, 40))
|
||||
newest = order[0]
|
||||
for d in vers[newest]:
|
||||
keep_ids.add(d["_id"])
|
||||
kept_versions += 1
|
||||
if len(order) > 1:
|
||||
n_multi += 1
|
||||
base = [x for d in vers[newest] for x in norm_lines(d["text"])]
|
||||
kept_lines = [base]
|
||||
for v in order[1:]:
|
||||
lines = [x for d in vers[v] for x in norm_lines(d["text"])]
|
||||
frac = diff_fraction(base, lines)
|
||||
if frac > a.diff:
|
||||
kept_versions += 1
|
||||
for d in vers[v]:
|
||||
keep_ids.add(d["_id"])
|
||||
# info only: would a stricter rule (compare with all kept versions) remove it?
|
||||
if min(diff_fraction(k, lines) for k in kept_lines) <= a.diff:
|
||||
strict_extra += 1
|
||||
kept_lines.append(lines)
|
||||
else:
|
||||
removed_versions.append({"object": f, "version": v, "newest": newest, "diff": round(frac, 3),
|
||||
"tokens": sum(d["real_tokens"] for d in vers[v])})
|
||||
after_dedup = [d for d in docs if d["_id"] in keep_ids]
|
||||
real1 = sum(d["real_tokens"] for d in after_dedup)
|
||||
|
||||
# ---- split long documents
|
||||
items, split_info, hard_total = [], [], 0
|
||||
for d in after_dedup:
|
||||
if d["real_tokens"] <= a.max_len:
|
||||
items.append({"text": d["text"], "types": d["types"], "fam": FAM.sub("", d["objects"][0]) if d["objects"] else str(d["_id"]),
|
||||
"src": d["objects"][:1], "tokens": d["real_tokens"], "pieces": 1})
|
||||
continue
|
||||
pieces, hard = split_document(d, tok, a.max_len)
|
||||
hard_total += hard
|
||||
for p in pieces:
|
||||
items.append({"text": p, "types": d["types"], "fam": FAM.sub("", d["objects"][0]) if d["objects"] else str(d["_id"]),
|
||||
"src": d["objects"][:1], "tokens": ntok(p), "pieces": len(pieces)})
|
||||
split_info.append({"object": d["objects"][:1], "tokens": d["real_tokens"], "pieces": len(pieces),
|
||||
"piece_tokens": [x["tokens"] for x in items[-len(pieces):]]})
|
||||
over = [it for it in items if it["tokens"] > a.max_len]
|
||||
kept = [it for it in items if it["tokens"] <= a.max_len]
|
||||
real2 = sum(it["tokens"] for it in items)
|
||||
|
||||
# ---- split by family
|
||||
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)
|
||||
for it in kept:
|
||||
fam[it["fam"]].append(it)
|
||||
keys = sorted(fam)
|
||||
random.Random(a.seed).shuffle(keys)
|
||||
n_valid, valid, train = max(1, round(len(kept) * a.valid_frac)), [], []
|
||||
@@ -55,58 +259,71 @@ def main():
|
||||
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")
|
||||
for it in part:
|
||||
f.write(json.dumps({"text": it["text"]}, ensure_ascii=False) + "\n")
|
||||
|
||||
def summ(ds):
|
||||
v = sorted(d["real_tokens"] for d in ds)
|
||||
v = sorted(x["tokens"] for x 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}
|
||||
"p99": pct(v, .99), "max": v[-1]}
|
||||
|
||||
def share(ds):
|
||||
c, tot = collections.Counter(), sum(x["tokens"] for x in ds)
|
||||
for x in ds:
|
||||
c["DOC" if x["types"] == ["DOC"] else "CLAS" if x["types"] == ["CLAS"] else "other"] += x["tokens"]
|
||||
return {k: [v, round(100 * v / tot, 1)] for k, v in c.items()}
|
||||
|
||||
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"])
|
||||
for it in kept:
|
||||
k = "+".join(it["types"]) if len(it["types"]) <= 2 else "mixed(%d)" % len(it["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
|
||||
by_type[k][1] += it["tokens"]
|
||||
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}
|
||||
"records": len(docs), "tokens_before_dedup": real0, "tokens_after_dedup": real1,
|
||||
"records_after_dedup": len(after_dedup), "tokens_after_split": real2,
|
||||
"objects_with_versions": n_multi, "families": len(byfam), "versions_kept": kept_versions,
|
||||
"versions_removed": len(removed_versions), "strict_rule_would_remove_more": strict_extra,
|
||||
"removed_versions": removed_versions, "split_docs": len(split_info), "hard_line_splits": hard_total,
|
||||
"pieces_over_limit": len(over), "train": summ(train), "valid": summ(valid), "all": summ(kept),
|
||||
"share_before_dedup": share([{"types": d["types"], "tokens": d["real_tokens"]} for d in docs]),
|
||||
"share_final": share(kept), "by_type": {k: {"docs": v[0], "tokens": v[1]} for k, v in by_type.items()},
|
||||
"split_info": split_info, "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"Corpus: `{a.corpus}` (SAP-samples/abap-cheat-sheets, Apache-2.0, see corpus/out/ATTRIBUTION.md).",
|
||||
f"Tokenizer: `{a.model}`. Seed {a.seed}. max_seq_length {a.max_len}.", "",
|
||||
"## Version dedup (keep the newest; keep an older version only if it differs by more than "
|
||||
f"{int(a.diff * 100)} % of lines)", "",
|
||||
f"- Object families: {len(byfam)}; with more than one version: {n_multi}.",
|
||||
f"- Versions kept: {kept_versions}. Versions removed: {len(removed_versions)}.",
|
||||
f"- Records: {len(docs)} before, {len(after_dedup)} after.",
|
||||
f"- Tokens before: {real0} (chars/4 estimate {est}). After dedup: {real1}. After splitting: {real2}.",
|
||||
f"- Info: a stricter rule (compare also with the other kept versions) would remove {strict_extra} more kept versions.",
|
||||
"", "## Splitting of documents over the limit", "",
|
||||
f"- Documents split: {len(split_info)}; pieces: {sum(s['pieces'] for s in split_info)}; hard line splits "
|
||||
f"(a block was too big): {hard_total}; pieces still over the limit: {len(over)}.", "",
|
||||
"## Result", "", "| | docs | tokens |", "|---|---|---|",
|
||||
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 |", "|---|---|---|"]
|
||||
f"Split by family: {len(fam)} families. **Iterations for 2 epochs (batch 1): {len(train) * 2}.**", "",
|
||||
"## Token share by type", "", "| | before dedup | final (kept pieces) |", "|---|---|---|"]
|
||||
for k in ("DOC", "CLAS", "other"):
|
||||
b, f = stats["share_before_dedup"].get(k, [0, 0]), stats["share_final"].get(k, [0, 0])
|
||||
L.append(f"| {k} | {b[0]} ({b[1]} %) | {f[0]} ({f[1]} %) |")
|
||||
L += ["", "## Token distribution per document (final)", "", "| set | min | p50 | p90 | p99 | max |", "|---|---|---|---|---|---|"]
|
||||
for n_ in ("all", "train", "valid"):
|
||||
s = stats[n_]
|
||||
L.append(f"| {n_} | {s['min']} | {s['p50']} | {s['p90']} | {s['p99']} | {s['max']} |")
|
||||
L += ["", "## By object type (final)", "", "| 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."]
|
||||
L += ["", "## Removed versions", "", "| object | version | newest | diff | tokens |", "|---|---|---|---|---|"]
|
||||
L += [f"| {r['object']} | {r['version']} | {r['newest']} | {r['diff']} | {r['tokens']} |" for r in removed_versions]
|
||||
L += ["", "## Split documents", ""]
|
||||
L += [f"- {s['object']}: {s['tokens']} tokens -> {s['pieces']} pieces {s['piece_tokens']}" for s in split_info]
|
||||
open(os.path.join(OUT, "report.md"), "w").write("\n".join(L) + "\n")
|
||||
print("\n".join(L))
|
||||
print("\n".join(L[:60]))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user