promotion's picture
Upload code/polyedit/rule_ood.py with huggingface_hub
fd4b9b1 verified
Raw
History Blame Contribute Delete
2.37 kB
"""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