"""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)