File size: 15,067 Bytes
4a3e194
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
79ed4eb
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
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
"""Fast verification for the paper: codec, framing, serialization, interpreter.

Everything here finishes in minutes. The long generation runs are in
paper_runs.py.
"""

from __future__ import annotations

import hashlib
import json
import os
import sys
import time
from typing import Dict, List

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

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)
R: Dict = {}


def sha(b: bytes) -> str:
    return hashlib.sha256(b).hexdigest()


def family_files() -> List[str]:
    files = [os.path.join(REPO, "neural_computer.safetensors")]
    vdir = os.path.join(REPO, "variants")
    files += sorted(os.path.join(vdir, f) for f in os.listdir(vdir)
                    if f.endswith(".safetensors"))
    return files


# ---------------------------------------------------------------------------
def check_codec():
    from selfrep import describe, describe_literal, decode
    print("[1] Recipe language")
    files = family_files()
    total = 0
    bad = 0
    tape_total = 0
    for p in files:
        data = open(p, "rb").read()
        r = describe(data)
        if decode(r) != data:
            bad += 1
        total += len(data)
        tape_total += len(r)
    print(f"    delta(describe(f)) = f on all {len(files)} artifact files "
          f"({total:,} bytes -> {tape_total:,} bytes of recipe): "
          f"{'all exact' if bad == 0 else f'{bad} FAILED'}")

    # the length bound of the completeness lemma, with c1 = 1, c2 = 127, c3 = 1
    worst = 0.0
    lb = 0
    for p in files:
        data = open(p, "rb").read()
        rl = describe_literal(data)
        assert decode(rl) == data
        bound = len(data) + -(-len(data) // 127) + 1
        assert len(rl) == bound, (len(rl), bound)
        lb += 1
    # exhaustive on short strings
    import random
    rng = random.Random(0)
    ebad = 0
    for n in range(0, 600):
        s = bytes(rng.randrange(256) for _ in range(n))
        for enc in (describe, describe_literal):
            r = enc(s)
            if decode(r) != s:
                ebad += 1
        if len(describe_literal(s)) != len(s) + -(-len(s) // 127) + 1:
            ebad += 1
    print(f"    literal encoding meets |r| = |f| + ceil(|f|/127) + 1 on all "
          f"{len(files)} files and on 600 random strings: "
          f"{'exact' if ebad == 0 else 'FAILED'}")
    R["codec"] = {"files": len(files), "bytes": total, "recipe_bytes": tape_total,
                  "roundtrip_failures": bad, "bound_failures": ebad}


# ---------------------------------------------------------------------------
def check_framing():
    from selfrep import ser, inst, M_STAR, M_P, read_host, tau_star, describe
    print("[2] Framing")
    sigma = read_host()
    bad = 0
    cases = [(sigma, M_STAR, tau_star(sigma, M_STAR)),
             (sigma, M_P, describe(b"")),
             (b"", M_P, b"\x00"),
             (b"\x00" * 3, M_STAR, b"")]
    for sg, m, tau in cases:
        s = ser(sg, bytes(m), tau)
        got = inst(s)
        if got != (sg, bytes(m), tau):
            bad += 1
    print(f"    inst(ser(Omega)) = Omega on {len(cases)} instances: "
          f"{'exact' if bad == 0 else 'FAILED'}")
    R["framing"] = {"cases": len(cases), "failures": bad}


# ---------------------------------------------------------------------------
def check_serialization():
    """Under the canonical specification, rebuilding the host from source
    reproduces its distributed encoding byte for byte, and two independent
    builds agree."""
    print("[3] Canonical serialization")
    import struct
    from selfrep import read_host, sigma_host
    from check_sigma import resigma
    original = read_host()
    rebuilt = sigma_host()
    deterministic = sigma_host() == rebuilt
    same = rebuilt == original
    round_trip = resigma(original) == original
    print(f"    rebuild from source is byte-identical: {same} "
          f"({len(original):,} bytes, sha {sha(original)[:16]})")
    print(f"    two independent builds agree: {deterministic}")
    print(f"    the netlist the encoding denotes serializes back to it: "
          f"{round_trip}")
    n = struct.unpack("<Q", original[:8])[0]
    hdr = json.loads(original[8:8 + n])
    meta = hdr.pop("__metadata__", {})
    keys = list(hdr)
    dtypes = sorted({v["dtype"] for v in hdr.values()})
    order = {"F64": 0, "F32": 1, "F16": 2, "I64": 3, "I32": 4, "I16": 5,
             "I8": 6, "U8": 7, "BOOL": 8}
    by_dtype_name = keys == sorted(keys, key=lambda k: (order[hdr[k]["dtype"]], k))
    pad = len(original[8:8 + n]) - len(original[8:8 + n].rstrip(b" "))
    print(f"    header {n:,} bytes, {len(keys)} tensors, dtypes {dtypes}, "
          f"metadata keys {sorted(meta)}")
    print(f"    tensor entries ordered by (element type, name): {by_dtype_name}; "
          f"data begins at {8+n:,}, 8-aligned: {(8 + n) % 8 == 0}, "
          f"{pad} bytes of padding")
    R["serialization"] = {"byte_identical": same, "deterministic": deterministic,
                          "round_trip": round_trip,
                          "bytes": len(original), "sha256": sha(original),
                          "header_bytes": n, "tensors": len(keys),
                          "dtypes": dtypes, "metadata_keys": sorted(meta),
                          "ordered_by_dtype_then_name": by_dtype_name,
                          "data_offset": 8 + n, "header_padding": pad}


# ---------------------------------------------------------------------------
def _batch_vectors(m, p, device):
    """State vectors for a batch of (memory image, counter) pairs."""
    import torch
    v = torch.zeros(len(p), 2057, device=device)
    for k in range(8):
        v[:, k] = ((p >> (7 - k)) & 1).float()
        v[:, 9 + k::8] = ((m >> (7 - k)) & 1).float()
    return v


def _bytes_of(v):
    """The counter and the 256 memory bytes of every state of a batch."""
    import torch
    pw = torch.tensor([1 << (7 - k) for k in range(8)], dtype=torch.long,
                      device=v.device)
    q = v.detach().long()
    pc = (q[:, :8] * pw).sum(1)
    mem = (q[:, 9:].reshape(q.shape[0], 256, 8) * pw).sum(2)
    return pc, mem


def _reference_batch(mem, pc):
    """One step of Definition 3.1 on a batch of memory images and counters."""
    import torch
    take = lambda idx: mem.gather(1, idx.unsqueeze(1)).squeeze(1)
    A = take(pc)
    B = take((pc + 1) & 0xFF)
    C = take((pc + 2) & 0xFF)
    r = (take(B) - take(A)) & 0xFF
    out = mem.clone()
    out.scatter_(1, B.unsqueeze(1), r.unsqueeze(1))
    leq = (r == 0) | (r >= 0x80)
    return out, torch.where(leq, C, (pc + 3) & 0xFF)


def check_host_datapath():
    """One step of the netlist that sigma(N_host) denotes, against
    Definition 3.1, over all 2^16 operand pairs and all 2^16 (pc, C) pairs."""
    print("[4] Host datapath")
    import torch
    from selfrep import NetEvaluator, read_host
    device = "cuda" if torch.cuda.is_available() else "cpu"
    E = NetEvaluator(read_host(), device=device)

    def sweep(mems, pcs, chunk=256):
        bad = 0
        for i in range(0, len(pcs), chunk):
            m = mems[i:i + chunk].to(device)
            p = pcs[i:i + chunk].to(device)
            got_pc, got_mem = _bytes_of(E.step(_batch_vectors(m, p, device)))
            want_mem, want_pc = _reference_batch(m, p)
            bad += int(((got_mem != want_mem).any(1)
                        | (got_pc != want_pc)).sum())
        return bad

    a = torch.arange(65536) % 256
    b = torch.arange(65536) // 256
    mems = torch.zeros(65536, 256, dtype=torch.long)
    mems[:, 0], mems[:, 1], mems[:, 2] = 0x90, 0x91, 0x30
    mems[:, 0x90] = a
    mems[:, 0x91] = b
    bad_ops = sweep(mems, torch.zeros(65536, dtype=torch.long))
    print(f"    one step against Definition 3.1 over all 65,536 operand pairs: "
          f"{'exact' if bad_ops == 0 else f'{bad_ops} FAILED'}")

    pc = torch.arange(65536) // 256
    c = torch.arange(65536) % 256
    mems = torch.zeros(65536, 256, dtype=torch.long)
    mems[:, 0xA0], mems[:, 0xA1] = 4, 5
    rows = torch.arange(65536)
    mems[rows, pc] = 0xA0
    mems[rows, (pc + 1) % 256] = 0xA1
    mems[rows, (pc + 2) % 256] = c
    bad_ctr = sweep(mems, pc)
    print(f"    one step against Definition 3.1 over all 65,536 (pc, C) pairs, "
          f"the counter wrapping across the top of memory: "
          f"{'exact' if bad_ctr == 0 else f'{bad_ctr} FAILED'}")
    R["datapath"] = {"operand_pairs": 65536, "operand_failures": bad_ops,
                     "counter_pairs": 65536, "counter_failures": bad_ctr}


# ---------------------------------------------------------------------------
def check_interpreter():
    """One-record semantics exhaustively over record fields, and the dependence
    hypothesis of the size proposition."""
    print("[5] Interpreter U")
    import torch
    from reflect import (Cfg, build_net, Leveled, encode_gate, encode_netlist,
                         pad, ref_step, state_to_vec, vec_to_state)
    cfg = Cfg()
    net, inputs, outputs = build_net(cfg)
    lev = Leveled(net, inputs, outputs, device="cpu")

    def U(sig, gp, halt=0):
        v = state_to_vec(cfg, {"sig": sig, "gp": gp, "halt": halt}).unsqueeze(0)
        out = lev.step(v[:, :len(inputs)])[0]
        return vec_to_state(cfg, out)

    ra, rb, ro = cfg.WORK_BASE, cfg.WORK_BASE + 1, cfg.WORK_BASE + 2
    bad = total = 0
    for w0 in (-1, 0, 1):
        for w1 in (-1, 0, 1):
            for bias in range(-(1 << (cfg.BB - 1)), 1 << (cfg.BB - 1)):
                nl = encode_netlist(cfg, [([(ra, w0, 0), (rb, w1, 0)], bias, (ro, 0))])
                for va in (0, 1):
                    for vb in (0, 1):
                        sig = [0] * cfg.S
                        sig[cfg.NET0:cfg.NET0 + len(nl)] = nl
                        sig[ra], sig[rb] = va, vb
                        g = U(sig, 0)
                        exp = 1 if w0 * va + w1 * vb + bias >= 0 else 0
                        total += 1
                        if g["sig"][ro] != exp or g["gp"] != 1:
                            bad += 1
    print(f"    one-record semantics over all {total} (w0,w1,bias,x0,x1) "
          f"configurations: {'exact' if bad == 0 else f'{bad} FAILED'}")

    # record counter advances modulo G
    cbad = 0
    nl = encode_netlist(cfg, [])
    for gp in range(cfg.G):
        sig = [0] * cfg.S
        sig[cfg.NET0:cfg.NET0 + len(nl)] = nl
        g = U(sig, gp)
        if g["gp"] != (gp + 1) % cfg.G:
            cbad += 1
    print(f"    record counter advances mod G on all {cfg.G} values: "
          f"{'exact' if cbad == 0 else 'FAILED'}")

    # dependence hypothesis of the size proposition: for every record-region
    # coordinate q, exhibit a state x and a data coordinate s with
    # U(x)_s != U(x xor e_q)_s. A random state rarely witnesses this, because a
    # random output address lands in the 106-bit data region about one time in
    # ten; the targeted pass of check_dependence forces the record under test to
    # write into the data region and makes the tested field decisive.
    import random
    from check_dependence import targeted

    def U1(sig, gp):
        return U(sig, gp)["sig"]

    rng = random.Random(20260910)
    D0, D1 = cfg.WORK_BASE, cfg.NET0
    Q0 = cfg.NET0
    Qbits = list(range(Q0, Q0 + cfg.G * cfg.R))
    influenced = set()
    for q in Qbits:
        g, r = divmod(q - Q0, cfg.R)
        found = 0
        for attempt in range(160):
            slots = [(rng.randrange(D0, D1 - 1), rng.choice((-1, 0, 1)),
                      rng.randint(0, 1)) for _ in range(2)]
            bias = rng.randrange(-(1 << (cfg.BB - 1)), 1 << (cfg.BB - 1))
            out = (rng.randrange(D0, D1 - 1), rng.randint(0, 1))
            rec = encode_gate(cfg, pad(cfg, slots), bias, out)
            sig = [0] * cfg.S
            for d in range(D0, D1):
                sig[d] = rng.randint(0, 1)
            ptr = rng.choice((0, 1, 2, 3))
            for k in range(cfg.A):
                sig[cfg.PTR_BASE + k] = (ptr >> (cfg.A - 1 - k)) & 1
            base = cfg.bank_base(0) + g * cfg.R
            sig[base:base + cfg.R] = rec
            s0, s1 = list(sig), list(sig)
            s1[q] ^= 1
            a0, a1 = U1(s0, g), U1(s1, g)
            if any(a0[d] != a1[d] for d in range(D0, D1)):
                found = attempt + 1
                break
        if not found:
            found = targeted(cfg, U1, rng, q, g, r, D0, D1)
        if found:
            influenced.add(q)
    print(f"    every one of the {len(Qbits)} record-region bits influences some "
          f"data-region bit of U(x): {len(influenced)}/{len(Qbits)}"
          f"{'  (hypothesis holds)' if len(influenced) == len(Qbits) else '  FAILED'}")
    R["interpreter"] = {
        "onerecord_cases": total, "onerecord_failures": bad,
        "counter_failures": cbad,
        "Q_bits": len(Qbits), "influencing": len(influenced),
        "S": cfg.S, "G": cfg.G, "R": cfg.R, "b": cfg.R,
        "Q_size": cfg.G * cfg.R, "records_max": (cfg.G * cfg.R) // cfg.R,
        "units_lower_bound": (cfg.G * cfg.R + 1) // 2,
    }
    print(f"    |Q| = {cfg.G*cfg.R}, b = {cfg.R}: at most {(cfg.G*cfg.R)//cfg.R} "
          f"records stored; any fan-in-2 netlist realising U has at least "
          f"{(cfg.G*cfg.R+1)//2} units")


# ---------------------------------------------------------------------------
def check_lev_sparse_dense():
    """The sparse and dense evaluations of Lev(N_host) agree."""
    print("[6] Lev(N_host): sparse vs dense")
    import torch
    from selfrep import LevEvaluator, NetEvaluator, read_host
    sigma = read_host()
    dense = LevEvaluator(sigma, device="cpu", dense=True)
    sparse = LevEvaluator(sigma, device="cpu", dense=False)
    netl = NetEvaluator(sigma, device="cpu")
    gen = torch.Generator().manual_seed(4242)
    V = (torch.rand(64, LevEvaluator.N, generator=gen) < 0.5).float()
    a = dense.step(V)
    b = sparse.step(V)
    c = netl.step(V)
    same = bool((a == b).all()) and bool((a == c).all())
    print(f"    identical next state on 64 uniformly random states: {same}")
    print(f"    dense: {dense.info['layers']} layers, max width "
          f"{dense.info['max_width']}, {dense.info['total_weights']:,} entries")
    R["lev"] = {"sparse_dense_agree": same,
                "layers": dense.info["layers"],
                "max_width": dense.info["max_width"],
                "total_weights": dense.info["total_weights"]}


def main() -> int:
    t0 = time.perf_counter()
    check_codec()
    check_framing()
    check_serialization()
    check_host_datapath()
    check_interpreter()
    check_lev_sparse_dense()
    R["seconds"] = time.perf_counter() - t0
    path = runs_path("paper_verify.json")
    json.dump(R, open(path, "w"), indent=1)
    print(f"\nwrote {path}  ({R['seconds']:.0f}s)")
    return 0


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