Files
abap-llm/train/prepare.py

331 lines
15 KiB
Python

"""Stage 1 data preparation (Kral decisions 2026-10-03).
train/.venv/bin/python train/prepare.py [--corpus PATH] [--max-len 16384] [--seed 20261003]
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
import re
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"))
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)
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"] = 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)
# ---- 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 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)), [], []
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 it in part:
f.write(json.dumps({"text": it["text"]}, ensure_ascii=False) + "\n")
def summ(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]}
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 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] += it["tokens"]
stats = {"corpus": a.corpus, "seed": a.seed, "max_len": a.max_len, "estimate_tokens": est,
"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}` (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"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 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[:60]))
if __name__ == "__main__":
main()