File size: 2,371 Bytes
fd4b9b1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
"""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