File size: 5,582 Bytes
4a3e194
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
"""Lemma 2.6 on the host: the levelized map equals the netlist it came from.

Lev(N) is asserted to compute, for every state, the next state that N computes.
The netlist here is the one the serialization of the host denotes, evaluated
gate by gate in topological order with no reference to the levelization, and the
two are compared on uniformly random states, which lie off every trajectory of
the constructor.

The same states carry the structural facts the lemma asserts: every weight of
every layer lies in {-1, 0, 1}, and every pre-activation is an integer, so the
comparator has a margin of one half.
"""
import json
import os
import sys

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import torch

REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))


def runs_path(name: str) -> str:
    d = os.path.join(REPO, "paper", "runs")
    os.makedirs(d, exist_ok=True)
    return os.path.join(d, name)


def gate_by_gate(net, inputs, outputs, V):
    """Evaluate the netlist directly: topological order, one gate at a time."""
    gates = net.gates
    indeg, cons = {}, {}
    for g, (ins, _) in gates.items():
        indeg[g] = len([s for s, _ in ins if s in gates])
        for s, _ in ins:
            cons.setdefault(s, []).append(g)
    order = [g for g, d in indeg.items() if d == 0]
    i = 0
    while i < len(order):
        for c in cons.get(order[i], []):
            indeg[c] -= 1
            if indeg[c] == 0:
                order.append(c)
        i += 1
    assert len(order) == len(gates)

    B = V.shape[0]
    val = {"#0": torch.zeros(B), "#1": torch.ones(B)}
    for k, name in enumerate(inputs):
        val[name] = V[:, k]
    for g in order:
        ins, bias = gates[g]
        acc = torch.full((B,), float(bias))
        for s, w in ins:
            acc = acc + w * val[s]
        val[g] = (acc >= 0).float()
    return torch.stack([val[o] for o in outputs], dim=1)


def main() -> int:
    from netlist_io import net_of_sigma
    from selfrep import LevEvaluator, NetEvaluator, read_host

    sigma = read_host()
    net, inputs, outputs, _ = net_of_sigma(sigma)
    L = LevEvaluator(sigma, device="cpu", dense=False)

    gen = torch.Generator().manual_seed(20260910)
    B = 96
    V = (torch.rand(B, len(inputs), generator=gen) < 0.5).float()

    direct = gate_by_gate(net, inputs, outputs, V)
    levelled = L.step(V)
    agree = bool((direct == levelled).all())
    print(f"  netlist has {len(net.gates):,} gates on {len(inputs):,} state bits")
    print(f"  gate-by-gate evaluation equals Lev(N_host) on {B} uniformly random "
          f"states: {'yes' if agree else 'NO'}")
    if not agree:
        d = (direct != levelled).nonzero()
        print("   first disagreements:", d[:5].tolist())

    grouped = NetEvaluator(sigma, device="cpu").step(V)
    agree_net = bool((direct == grouped).all())
    print(f"  the evaluator that carries only the netlist's own predecessor "
          f"entries agrees on the same states: {'yes' if agree_net else 'NO'}")

    # The dense form: ternary weights, integer pre-activations, margin one half.
    D = LevEvaluator(sigma, device="cpu", dense=True)
    nonternary = [i for i, W in enumerate(D.W)
                  if not set(torch.unique(W).tolist()) <= {-1.0, 0.0, 1.0}]
    nonint = [i for i, b in enumerate(D.B) if not bool((b == b.round()).all())]
    print(f"  every weight of all {len(D.W)} layers lies in {{-1,0,1}}: "
          f"{'yes' if not nonternary else f'NO at {nonternary[:3]}'}")
    print(f"  every bias is an integer: {'yes' if not nonint else f'NO at {nonint[:3]}'}")

    y, worst, frac = V, float("inf"), 0.0
    for W, b in zip(D.W, D.B):
        pre = y @ W.T + b
        worst = min(worst, float((pre + 0.5).abs().min()))
        frac = max(frac, float((pre - pre.round()).abs().max()))
        y = (pre >= 0).float()
    same_dense = bool((y == levelled).all())
    print(f"  dense stack agrees with the sparse evaluation on the same states: "
          f"{'yes' if same_dense else 'NO'}")
    print(f"  minimum distance of a pre-activation from -1/2 over all "
          f"{len(D.W)} layers: {worst:.6f}")
    print(f"  maximum distance of a pre-activation from an integer: {frac:.3e}")

    # The same check for the interpreter of Theorem 7.2, whose levelization is
    # the object that Section 7 executes.
    from reflect import Cfg, build_net, Leveled
    cfg = Cfg()
    unet, uin, uout = build_net(cfg)
    UL = Leveled(unet, uin, uout, device="cpu")
    BU = 64
    VU = (torch.rand(BU, len(uin), generator=gen) < 0.5).float()
    udirect = gate_by_gate(unet, uin, uout, VU)
    ulev = UL.step(VU)
    uagree = bool((udirect == ulev).all())
    print(f"  interpreter netlist has {len(unet.gates):,} gates on "
          f"{len(uin):,} state bits")
    print(f"  gate-by-gate evaluation equals Lev(N_U) on {BU} uniformly random "
          f"states: {'yes' if uagree else 'NO'}")

    ok = (agree and agree_net and uagree and same_dense and not nonternary
          and not nonint and abs(worst - 0.5) < 1e-9)
    json.dump({"states": B, "gates": len(net.gates), "agree": agree,
               "agree_grouped": agree_net,
               "dense_agree": same_dense, "nonternary_layers": len(nonternary),
               "noninteger_bias_layers": len(nonint),
               "min_margin": worst, "max_noninteger": frac,
               "u_states": BU, "u_gates": len(unet.gates), "u_agree": uagree},
              open(runs_path("paper_lev.json"), "w"), indent=1)
    return 0 if ok else 1


if __name__ == "__main__":
    sys.exit(main())