"""Leakage-free transformation-family split for the real matched-pair graph.""" from __future__ import annotations import hashlib from collections import defaultdict from itertools import combinations from .realedits import build_mmp_index def family_id(a, b): """Direction-invariant ID, so a held-out replacement cannot leak through its reverse.""" pair = sorted((a, b)) return hashlib.sha1((pair[0] + ">>" + pair[1]).encode()).hexdigest() def build_rule_graphs(records, props, *, holdout_mod=5, holdout_bucket=0, min_support=2): have = {r["canon"]: r["props"] for r in records.values() if all(prop in r["props"] for prop in props)} subset = {pid: {"canon": r["canon"], "props": r["props"]} for pid, r in records.items() if r["canon"] in have} index, _ = build_mmp_index(subset, props[0]) support, pair_families = defaultdict(set), defaultdict(set) for context, variants in index.items(): for (a, pa), (b, pb) in combinations(sorted(variants.items()), 2): if pa == pb: continue fid = family_id(a, b) support[fid].add(context) pair_families[tuple(sorted((pa, pb)))].add(fid) eligible = {fid for fid, contexts in support.items() if len(contexts) >= min_support} edge_family = {pair: min(families & eligible) for pair, families in pair_families.items() if families & eligible} test_families = {fid for fid in eligible if int(fid, 16) % holdout_mod == holdout_bucket} train_families = eligible - test_families graphs = {"train": defaultdict(set), "test": defaultdict(set)} for (a, b), fid in edge_family.items(): split = "test" if fid in test_families else "train" graphs[split][a].add(b) graphs[split][b].add(a) graphs = {split: {p: sorted(ns) for p, ns in graph.items()} for split, graph in graphs.items()} values = {p: {prop: have[p][prop] for prop in props} for p in have} meta = {"min_support": min_support, "holdout_mod": holdout_mod, "holdout_bucket": holdout_bucket, "train_families": sorted(train_families), "test_families": sorted(test_families), "train_edges": sum(map(len, graphs["train"].values())) // 2, "test_edges": sum(map(len, graphs["test"].values())) // 2} return graphs, values, meta