Zap / training /train /build_data.py
ProCreations's picture
Zap v1: 20.7M-param web form & field classifier for autofill (bf16, MIT)
e4805d0 verified
Raw History Blame Contribute Delete
3.69 kB
"""Join teacher labels onto forms and write train/val/test sets + an MLM corpus.
Usage: python build_data.py REPS_DIR OUT_DIR PATTERN [PATTERN ...] --labels DIR [--labels DIR ...]
REPS_DIR/reps.jsonl.gz (from select.py over the same PATTERNs) fixes each signature's split/representative;
teacher labels are matched by signature across all --labels dirs.
Every harvested copy of a labeled signature inherits its labels (max 3 copies per signature).
"""
import glob
import gzip
import json
import os
import random
import sys
from collections import Counter, defaultdict
sys.path.insert(0, "/job/work")
from label.common import load_forms # noqa: E402
from label.select import labeled_sigs # noqa: E402
from label.teacher import FIELD_LABELS, FORM_LABELS # noqa: E402
MAX_COPIES = 3
def main(label_dir, out_dir, patterns, label_dirs):
os.makedirs(out_dir, exist_ok=True)
reps = {}
with gzip.open(os.path.join(label_dir, "reps.jsonl.gz"), "rt") as f:
for line in f:
r = json.loads(line)
reps[r["sig"]] = {"split": r["split"], "uid": r["uid"], "prio": r["prio"]}
by_sig = {sig: lab for sig, lab in labeled_sigs(label_dirs).items() if sig in reps}
print(f"{len(reps)} signatures, {len(by_sig)} labeled signatures", flush=True)
outs = {s: open(os.path.join(out_dir, f"{s}.jsonl"), "w") for s in ("train", "val", "test")}
mlm = gzip.open(os.path.join(out_dir, "mlm.jsonl.gz"), "wt")
n_copies = defaultdict(int)
seen_uid = set()
stats = Counter()
field_c, form_c = Counter(), Counter()
for pat in patterns:
for r in load_forms(pat):
if r["uid"] in seen_uid:
continue
seen_uid.add(r["uid"])
rep = reps.get(r["sig"])
if rep is None:
continue
split = rep["split"]
# held-out splits keep only the representative (exact copies would inflate the test set)
if split != "train" and r["uid"] != rep["uid"]:
continue
if n_copies[r["sig"]] >= MAX_COPIES:
continue
n_copies[r["sig"]] += 1
if split == "train" and n_copies[r["sig"]] == 1:
mlm.write(json.dumps({"ctx": r["ctx"], "fields": r["fields"]}, ensure_ascii=False) + "\n")
stats["mlm"] += 1
lab = by_sig.get(r["sig"])
if lab is None or len(lab["fields"]) != len(r["fields"]):
continue
row = {"uid": r["uid"], "sig": r["sig"], "url": r["url"], "host": r["host"], "lang": r.get("lang", ""),
"ctx": r["ctx"], "fields": r["fields"], "form_label": lab["form"], "field_labels": lab["fields"],
"prio": rep["prio"], "rep": r["uid"] == rep["uid"]}
outs[split].write(json.dumps(row, ensure_ascii=False) + "\n")
stats[split] += 1
if split == "train":
form_c[lab["form"]] += 1
field_c.update(lab["fields"])
for f in outs.values():
f.close()
mlm.close()
summary = {"stats": stats, "train_form_labels": form_c, "train_field_labels": field_c,
"field_labels": FIELD_LABELS, "form_labels": FORM_LABELS}
json.dump(summary, open(os.path.join(out_dir, "data_stats.json"), "w"), indent=1)
print(json.dumps(summary, indent=1), flush=True)
if __name__ == "__main__":
import argparse
ap = argparse.ArgumentParser()
ap.add_argument("reps_dir")
ap.add_argument("out")
ap.add_argument("patterns", nargs="+")
ap.add_argument("--labels", action="append", default=[])
a = ap.parse_args()
main(a.reps_dir, a.out, a.patterns, a.labels)