File size: 2,015 Bytes
3579fb4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Order-independent settling: a stored acyclic netlist evaluates correctly
after d passes for every order of its records."""
import os, random, sys, itertools
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import torch
from reflect import (Cfg, build_net, Leveled, encode_netlist, compile_to_reflect,
                     adder_net, eval_family_net, state_to_vec, vec_to_state, Runner)


def main() -> int:
    cfg = Cfg()
    net, inputs, outputs = build_net(cfg)
    R = Runner(cfg, Leveled(net, inputs, outputs, device="cpu"))
    W = 1
    anet, an, bn, sums, cout = adder_net(W)
    ia = {an[0]: cfg.WORK_BASE, bn[0]: cfg.WORK_BASE + 1}
    rg, addr, _ = compile_to_reflect(cfg, anet, ia, sums + [cout], cfg.WORK_BASE + 2)
    n = len(rg)
    print(f"stored netlist: {n} records (a full adder cell), G = {cfg.G}")
    orders = list(itertools.permutations(range(n))) if n <= 7 else None
    if orders is None:
        rng = random.Random(7)
        orders = [tuple(rng.sample(range(n), n)) for _ in range(500)]
        print(f"  sampling {len(orders)} random record orders")
    else:
        print(f"  all {len(orders)} record orders")
    bad = 0
    for perm in orders:
        nl = encode_netlist(cfg, [rg[p] for p in perm])
        for x in (0, 1):
            for y in (0, 1):
                sig = [0] * cfg.S
                sig[cfg.NET0:cfg.NET0 + len(nl)] = nl
                sig[ia[an[0]]] = x
                sig[ia[bn[0]]] = y
                st = R.run({"sig": sig, "gp": 0, "halt": 0}, cfg.G * cfg.G)
                got = st["sig"][addr[sums[0]]] | (st["sig"][addr[cout]] << 1)
                fam = eval_family_net(anet, {an[0]: x, bn[0]: y})
                want = fam[sums[0]] | (fam[cout] << 1)
                if got != want or got != x + y:
                    bad += 1
    print(f"  correct for every order and every input: "
          f"{'yes' if bad == 0 else f'NO ({bad} failures)'}")
    return 0 if bad == 0 else 1


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