serialize the host netlist unit by unit and instantiate each generation from the emitted bytes
0654a8d | """Exact self-reproduction for the SUBLEQ threshold host. | |
| This module supplies the parts the universal constructor needs in order to | |
| emit its own complete instance and not only its weight file: a tape device | |
| with a rewind request and an end-of-tape status, the three-phase program P*, | |
| the framing maps ser / inst, and runners for the three evaluators. | |
| Device (memory-mapped, all logic in the runtime, none of it in the netlist): | |
| 0xF9 C_WR write request writing 1 emits R_OUT, then R_OUT <- 0 | |
| 0xFA C_EOT end-of-tape status maintained by the device | |
| 0xFB C_RW rewind request writing 1 sets the head to 0 | |
| 0xFC C_RD read request writing 1 loads R_IN and advances the head | |
| 0xFD R_IN input register | |
| 0xFE R_OUT output register | |
| 0xFF halt program counter | |
| One step is: execute the SUBLEQ instruction, then apply the device to the | |
| resulting state in the order (write, read, rewind, clear requests). A request | |
| fires on the value 1, not on the fact that the cell was addressed, so the | |
| device reads the machine state and nothing else. | |
| Program variables: | |
| 0xF0 Z constant 0 (restored by every instruction that uses it) | |
| 0xF1 ONE constant 1 | |
| 0xF2 T1 scratch: repeat counter | |
| 0xF3 T2 scratch: literal counter | |
| 0xF4 NEG1 constant 0xFF | |
| 0xF5 EOK constant 1, the end-of-tape test target | |
| """ | |
| from __future__ import annotations | |
| import hashlib | |
| import os | |
| import struct | |
| import sys | |
| from typing import Dict, List, Optional, Tuple | |
| sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) | |
| REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| HOST_PATH = os.path.join(REPO, "variants", "neural_subleq8io_netlist.safetensors") | |
| # device cells | |
| C_WR, C_EOT, C_RW, C_RD, R_IN, R_OUT = 0xF9, 0xFA, 0xFB, 0xFC, 0xFD, 0xFE | |
| HALT_PC = 0xFF | |
| # program variables | |
| Z, ONE, T1, T2, NEG1, EOK = 0xF0, 0xF1, 0xF2, 0xF3, 0xF4, 0xF5 | |
| # ============================================================================= | |
| # Recipe language (grammar and decoder), reused from constructor8 | |
| # ============================================================================= | |
| def describe(data: bytes) -> bytes: | |
| """Compile bytes into a recipe. Literal tokens carry up to 127 bytes; a run | |
| of at least 4 equal bytes becomes a repeat token. Every byte is stored | |
| negated mod 256 so the machine recovers it with one subtraction.""" | |
| tape = bytearray() | |
| i, n = 0, len(data) | |
| while i < n: | |
| j = i | |
| while j < n and data[j] == data[i] and j - i < 127: | |
| j += 1 | |
| if j - i >= 4: | |
| tape.append(256 - (j - i)) | |
| tape.append((256 - data[i]) % 256) | |
| i = j | |
| continue | |
| k = i | |
| while k < n and k - i < 127: | |
| m = k | |
| while m < n and data[m] == data[k] and m - k < 4: | |
| m += 1 | |
| if m - k >= 4: | |
| break | |
| k += 1 | |
| k = max(k, i + 1) | |
| tape.append(k - i) | |
| tape.extend((256 - x) % 256 for x in data[i:k]) | |
| i = k | |
| tape.append(0) | |
| return bytes(tape) | |
| def describe_literal(data: bytes) -> bytes: | |
| """The all-literal encoding used in the proof of the length bound.""" | |
| tape = bytearray() | |
| for i in range(0, len(data), 127): | |
| block = data[i:i + 127] | |
| tape.append(len(block)) | |
| tape.extend((256 - x) % 256 for x in block) | |
| tape.append(0) | |
| return bytes(tape) | |
| def decode(tape: bytes) -> bytes: | |
| """delta: the decoding map on well-formed recipes.""" | |
| out = bytearray() | |
| i = 0 | |
| while True: | |
| t = tape[i] | |
| i += 1 | |
| if t == 0: | |
| return bytes(out) | |
| if t <= 127: | |
| for _ in range(t): | |
| out.append((256 - tape[i]) % 256) | |
| i += 1 | |
| elif t == 128: | |
| raise ValueError("tag 128 is reserved") | |
| else: | |
| out.extend([(256 - tape[i]) % 256] * (256 - t)) | |
| i += 1 | |
| # ============================================================================= | |
| # Framing: ser and inst | |
| # ============================================================================= | |
| def lam(k: int) -> bytes: | |
| """Eight-byte little-endian length field.""" | |
| return struct.pack("<Q", k) | |
| def field(u: bytes) -> bytes: | |
| return lam(len(u)) + u | |
| def ser(sigma: bytes, m: bytes, tau: bytes) -> bytes: | |
| assert len(m) == 256 | |
| return field(sigma) + field(m) + tau | |
| def inst(s: bytes) -> Tuple[bytes, bytes, bytes]: | |
| """Partial inverse of ser: recover (sigma, m, tau).""" | |
| if len(s) < 8: | |
| raise ValueError("truncated") | |
| a = struct.unpack("<Q", s[:8])[0] | |
| if len(s) < 8 + a + 8: | |
| raise ValueError("truncated") | |
| sigma = s[8:8 + a] | |
| b = struct.unpack("<Q", s[8 + a:16 + a])[0] | |
| if b != 256: | |
| raise ValueError("memory image is not 256 bytes") | |
| m = s[16 + a:16 + a + b] | |
| if len(m) != b: # declared 256 bytes, fewer present | |
| raise ValueError("truncated memory image") | |
| tau = s[16 + a + b:] | |
| return sigma, m, tau | |
| # ============================================================================= | |
| # Programs | |
| # ============================================================================= | |
| def _emit(prog: List[Tuple[int, int, int]]) -> Dict[int, int]: | |
| mem = {} | |
| for idx, (a, b, c) in enumerate(prog): | |
| mem[idx * 3] = a | |
| mem[idx * 3 + 1] = b | |
| mem[idx * 3 + 2] = c | |
| return mem | |
| # The decoding loop. Addresses are 3k; every branch target below is written as | |
| # an instruction index and resolved to 3*index by _emit. | |
| # | |
| # k0 T2 <- 0 | |
| # k1 T1 <- 0 | |
| # k2 request read of the tag | |
| # k3 T1 <- -T | |
| # k4 T2 <- T | |
| # k5 branch to k7 when T = 0 or T >= 128 | |
| # k6 goto LITERAL | |
| # k7 branch to END when 256-T <= 0, that is T in {0,128} | |
| # k8 goto REPEAT with T1 = 256-T the run length | |
| # k9 LITERAL: request read of the next byte | |
| # k10 R_OUT <- b | |
| # k11 emit | |
| # k12 T2 <- T2-1; branch to k14 when the count is exhausted | |
| # k13 goto k9 | |
| # k14 goto k0 | |
| # k15 REPEAT: request read of the value byte | |
| # k16 R_OUT <- b | |
| # k17 emit | |
| # k18 T1 <- T1-1; branch to k14 when the count is exhausted | |
| # k19 goto k16 | |
| # k20 END | |
| DECODE_LOOP = [ | |
| (T2, T2, 3 * 1), # k0 | |
| (T1, T1, 3 * 2), # k1 | |
| (NEG1, C_RD, 3 * 3), # k2 | |
| (R_IN, T1, 3 * 4), # k3 | |
| (T1, T2, 3 * 5), # k4 | |
| (Z, T2, 3 * 7), # k5 | |
| (Z, Z, 3 * 9), # k6 | |
| (Z, T1, 3 * 20), # k7 | |
| (Z, Z, 3 * 15), # k8 | |
| (NEG1, C_RD, 3 * 10), # k9 | |
| (R_IN, R_OUT, 3 * 11), # k10 | |
| (NEG1, C_WR, 3 * 12), # k11 | |
| (ONE, T2, 3 * 14), # k12 | |
| (Z, Z, 3 * 9), # k13 | |
| (Z, Z, 3 * 0), # k14 | |
| (NEG1, C_RD, 3 * 16), # k15 | |
| (R_IN, R_OUT, 3 * 17), # k16 | |
| (NEG1, C_WR, 3 * 18), # k17 | |
| (ONE, T1, 3 * 14), # k18 | |
| (Z, Z, 3 * 16), # k19 | |
| ] | |
| # P: the constructor. The end token halts. | |
| P = DECODE_LOOP + [(Z, Z, HALT_PC)] | |
| # P*: the end token enters the rewind phase, then the copy phase. | |
| # k20 R: request rewind | |
| # k21 B0: request read | |
| # k22 B1: EOK <- 1-EOT; halt when the end of tape is reached | |
| # k23 B2: Z <- -b | |
| # k24 B3: R_OUT <- b | |
| # k25 B4: emit | |
| # k26 B5: Z <- 0 and go to B0 | |
| P_STAR = DECODE_LOOP + [ | |
| (NEG1, C_RW, 3 * 21), # k20 | |
| (NEG1, C_RD, 3 * 22), # k21 | |
| (C_EOT, EOK, HALT_PC), # k22 | |
| (R_IN, Z, 3 * 24), # k23 | |
| (Z, R_OUT, 3 * 25), # k24 | |
| (NEG1, C_WR, 3 * 26), # k25 | |
| (Z, Z, 3 * 21), # k26 | |
| ] | |
| # P_e: P with an end-of-tape guard after the tag read, so that the machine | |
| # halts on every tape, recipe or not. Instruction j3 assigns | |
| # M[EOK] <- M[EOK] - M[C_EOT]; a successful tag read has cleared C_EOT and | |
| # leaves EOK at 1, while a read at the end of the tape sets it and the result | |
| # 0 transfers control to the halt cell. | |
| P_TOTAL = [ | |
| (T2, T2, 3 * 1), # j0 | |
| (T1, T1, 3 * 2), # j1 | |
| (NEG1, C_RD, 3 * 3), # j2 read the tag | |
| (C_EOT, EOK, HALT_PC), # j3 halt if that read was past the end of the tape | |
| (R_IN, T1, 3 * 5), # j4 | |
| (T1, T2, 3 * 6), # j5 | |
| (Z, T2, 3 * 8), # j6 | |
| (Z, Z, 3 * 10), # j7 | |
| (Z, T1, 3 * 21), # j8 | |
| (Z, Z, 3 * 16), # j9 | |
| (NEG1, C_RD, 3 * 11), # j10 LITERAL | |
| (R_IN, R_OUT, 3 * 12), # j11 | |
| (NEG1, C_WR, 3 * 13), # j12 | |
| (ONE, T2, 3 * 15), # j13 | |
| (Z, Z, 3 * 10), # j14 | |
| (Z, Z, 3 * 0), # j15 | |
| (NEG1, C_RD, 3 * 17), # j16 REPEAT | |
| (R_IN, R_OUT, 3 * 18), # j17 | |
| (NEG1, C_WR, 3 * 19), # j18 | |
| (ONE, T1, 3 * 15), # j19 | |
| (Z, Z, 3 * 17), # j20 | |
| (Z, Z, HALT_PC), # j21 END | |
| ] | |
| def decode_any(tape: bytes, guard: bool) -> Tuple[bytes, bool]: | |
| """The output of the decoding loop on an arbitrary tape. | |
| `guard` selects P_e over P. A read at the end of the tape sets the | |
| end-of-tape status and leaves the input register holding the last byte it | |
| received, so a token whose payload overruns the tape is completed with | |
| copies of that byte. Without the guard the loop diverges exactly when the | |
| tape is exhausted at a tag read and the stale input register holds neither | |
| 0 nor 128; the second component of the result records whether the machine | |
| halts. | |
| """ | |
| out = bytearray() | |
| h, rin = 0, 0 | |
| while True: | |
| if h < len(tape): | |
| rin = tape[h] | |
| h += 1 | |
| eot = 0 | |
| else: | |
| eot = 1 | |
| if guard and eot: | |
| return bytes(out), True | |
| t = rin | |
| if not guard and eot: | |
| return bytes(out), t in (0, 128) | |
| if t == 0 or t == 128: | |
| return bytes(out), True | |
| reps = 1 if t > 128 else t | |
| count = (256 - t) if t > 128 else 1 | |
| for _ in range(reps): | |
| if h < len(tape): | |
| rin = tape[h] | |
| h += 1 | |
| out.extend([(256 - rin) % 256] * count) | |
| def memory_image(prog: List[Tuple[int, int, int]]) -> List[int]: | |
| """The 256-byte initial memory image holding a program and its constants.""" | |
| mem = [0] * 256 | |
| for addr, val in _emit(prog).items(): | |
| mem[addr] = val & 0xFF | |
| mem[Z] = 0 | |
| mem[ONE] = 1 | |
| mem[T1] = 0 | |
| mem[T2] = 0 | |
| mem[NEG1] = 0xFF | |
| mem[EOK] = 1 | |
| return mem | |
| M_P = memory_image(P) | |
| M_STAR = memory_image(P_STAR) | |
| # ============================================================================= | |
| # Device | |
| # ============================================================================= | |
| class Tape: | |
| """Environment state (tau, h, omega) with the operations of the definition.""" | |
| def __init__(self, tau: bytes): | |
| self.tau = tau | |
| self.h = 0 | |
| self.out = bytearray() | |
| def apply(self, mem: List[int]) -> None: | |
| """One device application to the post-instruction memory image.""" | |
| if mem[C_WR] == 1: | |
| self.out.append(mem[R_OUT]) | |
| mem[R_OUT] = 0 | |
| if mem[C_RD] == 1: | |
| if self.h < len(self.tau): | |
| mem[R_IN] = self.tau[self.h] | |
| mem[C_EOT] = 0 | |
| self.h += 1 | |
| else: | |
| mem[C_EOT] = 1 | |
| if mem[C_RW] == 1: | |
| self.h = 0 | |
| mem[C_WR] = 0 | |
| mem[C_RD] = 0 | |
| mem[C_RW] = 0 | |
| # ============================================================================= | |
| # Evaluator 1: integer reference | |
| # ============================================================================= | |
| def run_reference(mem0: List[int], tau: bytes, max_steps: int = 1 << 34, | |
| expect: Optional[bytes] = None) -> Tuple[bytes, int]: | |
| mem = list(mem0) | |
| dev = Tape(tau) | |
| pc = 0 | |
| steps = 0 | |
| while pc != HALT_PC and steps < max_steps: | |
| A = mem[pc] | |
| B = mem[(pc + 1) & 0xFF] | |
| C = mem[(pc + 2) & 0xFF] | |
| r = (mem[B] - mem[A]) & 0xFF | |
| mem[B] = r | |
| pc = C if (r == 0 or r >= 0x80) else (pc + 3) & 0xFF | |
| n_before = len(dev.out) | |
| dev.apply(mem) | |
| if expect is not None and len(dev.out) > n_before: | |
| k = len(dev.out) - 1 | |
| if k >= len(expect) or dev.out[k] != expect[k]: | |
| raise AssertionError(f"stream diverged at byte {k}") | |
| steps += 1 | |
| return bytes(dev.out), steps | |
| # ============================================================================= | |
| # The host netlist and its canonical serialization | |
| # ============================================================================= | |
| STATE_LAYOUT = {"pc": [0, 8], "halt": [8, 1], "mem": [9, 256, 8]} | |
| IO_CELLS = {"c_wr": C_WR, "c_eot": C_EOT, "c_rw": C_RW, "c_rd": C_RD, | |
| "r_in": R_IN, "r_out": R_OUT, "halt_pc": HALT_PC} | |
| STATE_BITS = 8 + 1 + 2048 | |
| def host_netlist(): | |
| """The clocked netlist of the host, assembled from its source description.""" | |
| import host_netlist as H | |
| return H.build_subleq_step_net() | |
| def sigma_host() -> bytes: | |
| """sigma(N_host): the canonical serialization of that netlist.""" | |
| from netlist_io import sigma_of_net | |
| net, inputs, outputs = host_netlist() | |
| return sigma_of_net(net, inputs, outputs, "subleq8io", STATE_LAYOUT, | |
| IO_CELLS) | |
| def read_host() -> bytes: | |
| """The distributed serialization of the host.""" | |
| return open(HOST_PATH, "rb").read() | |
| def sha(b: bytes) -> str: | |
| return hashlib.sha256(b).hexdigest() | |
| def tau_star(sigma: bytes, m: List[int]) -> bytes: | |
| return describe(field(sigma) + field(bytes(m))) | |
| # ============================================================================= | |
| # Evaluators of the step map, each built from sigma alone | |
| # ============================================================================= | |
| class _Transducer: | |
| """State marshalling shared by the threshold evaluators. | |
| A subclass supplies `step`, which maps a batch of state vectors to the next | |
| ones, and `n_in`, the width of the state. One step of the transducer of | |
| Definition 2.8 is that map followed by the device, so the run loop counts | |
| the step on which the halt bit is set and applies the device to it. | |
| """ | |
| N = STATE_BITS | |
| _DEVCELLS = (C_WR, C_EOT, C_RW, C_RD, R_IN, R_OUT) | |
| def _vec(self, pc: int, mem: List[int]): | |
| torch = self.torch | |
| v = torch.zeros(self.N) | |
| for k in range(8): | |
| v[k] = (pc >> (7 - k)) & 1 | |
| for j in range(256): | |
| for k in range(8): | |
| v[9 + j * 8 + k] = (mem[j] >> (7 - k)) & 1 | |
| return v | |
| def _byte(vc, j: int) -> int: | |
| x = 0 | |
| for k in range(8): | |
| x = (x << 1) | int(vc[9 + j * 8 + k]) | |
| return x | |
| def _set_byte(self, v, j: int, val: int) -> None: | |
| for k in range(8): | |
| v[0, 9 + j * 8 + k] = float((val >> (7 - k)) & 1) | |
| def capture(self): | |
| """Capture one application of the map as a CUDA graph. | |
| The map is a fixed sequence of operations on fixed shapes, so the whole | |
| step replays as one graph launch in place of some hundreds. The captured | |
| map is compared against the eager one before it is used. | |
| """ | |
| torch = self.torch | |
| assert self.device.startswith("cuda") | |
| self.gin = torch.zeros(1, self.N, device=self.device) | |
| # every buffer the body writes must stay alive for the life of the | |
| # graph, or the allocator will hand its memory to something else and | |
| # the replay will overwrite that instead | |
| self.gbuf = self._buffers() | |
| for _ in range(5): | |
| self._body(self.gin) | |
| torch.cuda.synchronize() | |
| self.graph = torch.cuda.CUDAGraph() | |
| with torch.cuda.graph(self.graph): | |
| self.gout = self._body(self.gin) | |
| torch.cuda.synchronize() | |
| gen = torch.Generator(device=self.device).manual_seed(3) | |
| probe = (torch.rand(1, self.N, generator=gen, | |
| device=self.device) < 0.5).float() | |
| self.gin.copy_(probe) | |
| self.graph.replay() | |
| torch.cuda.synchronize() | |
| assert bool((self.gout == self._eager(probe)).all()), \ | |
| "the captured graph differs from the eager step" | |
| return self | |
| def _replay(self, v): | |
| # the state lives in gin, which is ordinary memory; gout belongs to the | |
| # graph's private pool and is only ever read as a whole | |
| if v.data_ptr() != self.gin.data_ptr(): | |
| self.gin.copy_(v) | |
| self.graph.replay() | |
| self.gin.copy_(self.gout) | |
| return self.gin | |
| def run(self, mem0: List[int], tau: bytes, max_steps: int, | |
| expect: Optional[bytes] = None, progress: int = 0, | |
| margin: bool = False) -> Tuple[bytes, int]: | |
| """Iterate the map with the device applied after each step. | |
| Only the halt bit and the six device cells cross to the host each step; | |
| the rest of the state stays on the accelerator. The six cells are read | |
| out in one operation and written back in one, so the cost of a step is | |
| the map and not the marshalling. With margin=True the minimum distance | |
| of any pre-activation from -1/2 along the whole trajectory is | |
| accumulated (dense mode only).""" | |
| import time | |
| torch = self.torch | |
| v = self._vec(0, mem0).unsqueeze(0).to(self.device) | |
| dev = Tape(tau) | |
| cells = list(self._DEVCELLS) | |
| cell_t = torch.tensor([9 + j * 8 + k for j in cells for k in range(8)], | |
| device=self.device) | |
| pow2_t = torch.tensor([1 << (7 - k) for k in range(8)], | |
| device=self.device, dtype=torch.float32) | |
| shift_t = torch.tensor([7 - k for k in range(8)]) | |
| shadow = [0] * 256 | |
| self.min_margin = float("inf") | |
| n = 0 | |
| t0 = time.perf_counter() | |
| while n < max_steps: | |
| if margin: | |
| self._accumulate_margin(v) | |
| v = self.step(v) | |
| n += 1 | |
| vals = torch.cat([v[0, 8:9], | |
| (v[0, cell_t].reshape(len(cells), 8) * pow2_t) | |
| .sum(-1)]).to("cpu").to(torch.int64).tolist() | |
| for c, j in enumerate(cells): | |
| shadow[j] = vals[1 + c] | |
| before = len(dev.out) | |
| dev.apply(shadow) | |
| new = torch.tensor([shadow[j] for j in cells], dtype=torch.int64) | |
| v[0, cell_t] = (((new.unsqueeze(-1) >> shift_t) & 1) | |
| .reshape(-1).float().to(self.device)) | |
| if expect is not None and len(dev.out) > before: | |
| k = len(dev.out) - 1 | |
| if k >= len(expect) or dev.out[k] != expect[k]: | |
| raise AssertionError(f"stream diverged at byte {k}") | |
| if vals[0] >= 1: | |
| break | |
| if progress and n % progress == 0: | |
| rate = n / (time.perf_counter() - t0) | |
| print(f" {self.tag} {n:,} steps, {len(dev.out):,} bytes " | |
| f"({rate:,.0f} steps/s)", flush=True) | |
| self.seconds = time.perf_counter() - t0 | |
| return bytes(dev.out), n | |
| class NetEvaluator(_Transducer): | |
| """The netlist of sigma, evaluated unit by unit. | |
| The units are grouped by depth and each reads its predecessors out of one | |
| signal vector, so no unit is evaluated before its predecessors and none is | |
| padded to a common width: the evaluation carries the netlist's own | |
| 50,250 predecessor entries and nothing else. | |
| """ | |
| tag = "net" | |
| def __init__(self, sigma: bytes, device: str = "cpu", graph: bool = False): | |
| import time | |
| import torch | |
| from netlist_io import net_of_sigma | |
| from reflect import Leveled | |
| t0 = time.perf_counter() | |
| self.torch = torch | |
| self.device = device | |
| net, inputs, outputs, meta = net_of_sigma(sigma) | |
| self.net, self.inputs, self.outputs, self.meta = net, inputs, outputs, meta | |
| assert len(inputs) == self.N, "state width disagrees with the layout" | |
| self.lev = Leveled(net, inputs, outputs, device=device) | |
| self.info = {"levels": len(self.lev.plan), "units": len(net.gates), | |
| "entries": sum(len(i) for i, _ in net.gates.values())} | |
| self.graph = None | |
| if graph: | |
| self.capture() | |
| self.build_seconds = time.perf_counter() - t0 | |
| def _buffers(self): | |
| return self.torch.zeros(self.lev.n_sig, 1, device=self.device) | |
| def _body(self, inp): | |
| V = self.gbuf | |
| lev = self.lev | |
| V.zero_() | |
| V[1] = 1.0 | |
| V[lev.in_slots] = inp.T | |
| for idx, w, b, out in lev.plan: | |
| g = V[idx] | |
| V[out] = ((g * w[:, :, None]).sum(1) + b[:, None] >= 0).float() | |
| return V[lev.out_slots].T | |
| def _eager(self, v): | |
| return self.lev.step(v) | |
| def step(self, v): | |
| if self.graph is not None: | |
| return self._replay(v) | |
| return self.lev.step(v) | |
| class LevEvaluator(_Transducer): | |
| """Lev(N) for the netlist N of sigma. | |
| `dense=True` materialises the matrices of Lemma 2.6 and iterates | |
| matrix-vector products; `dense=False` evaluates the same map from its | |
| nonzero entries, the identity rows included, which is the same function | |
| layer by layer. The two agree on every state tested by check_lev. | |
| """ | |
| tag = "lev" | |
| def __init__(self, sigma: bytes, device: str = "cpu", dense: bool = True, | |
| graph: bool = False): | |
| import time | |
| import torch | |
| from netlist_io import net_of_sigma | |
| from matrix8 import compile_net | |
| t0 = time.perf_counter() | |
| self.torch = torch | |
| self.device = device | |
| self.dense = dense | |
| net, inputs, outputs, meta = net_of_sigma(sigma) | |
| self.net, self.inputs, self.outputs, self.meta = net, inputs, outputs, meta | |
| assert len(inputs) == self.N, "state width disagrees with the layout" | |
| layers, info = compile_net(net, inputs, outputs) | |
| for W, _ in layers: | |
| assert set(torch.unique(W).tolist()) <= {-1, 0, 1} | |
| self.info = dict(info) | |
| self.info["size"] = sum(int(W.shape[0]) for W, _ in layers) | |
| self.info["nonzero"] = sum(int((W != 0).sum()) for W, _ in layers) | |
| if dense: | |
| self.W = [W.to(device=device, dtype=torch.float32) | |
| for W, _ in layers] | |
| self.B = [b.to(device=device, dtype=torch.float32) | |
| for _, b in layers] | |
| else: | |
| pad = device.startswith("cuda") | |
| self.plan = [self._sparsify(W, b, device, pad) for W, b in layers] | |
| self.widths = [int(W.shape[0]) for W, _ in layers] | |
| self.graph = None | |
| if graph: | |
| self.capture() | |
| self.build_seconds = time.perf_counter() - t0 | |
| def _sparsify(W, b, device, pad): | |
| """One layer as the nonzero entries of its rows. | |
| With `pad`, every row of the layer is padded to the largest number of | |
| nonzero entries in it by an entry of weight zero, which contributes | |
| nothing to any pre-activation and makes the layer one gather and one | |
| reduction; that is the faster arrangement on the accelerator, where the | |
| cost is the number of operations issued. Without it the rows are grouped | |
| by their number of nonzero entries and no padding is read, which is the | |
| faster arrangement on a processor core. | |
| """ | |
| import torch | |
| nz = (W != 0) | |
| counts = nz.sum(1) | |
| n = int(W.shape[0]) | |
| sizes = [int(counts.max())] if pad else sorted(set(counts.tolist())) | |
| groups = [] | |
| for k in sizes: | |
| rows = (torch.arange(n) if pad | |
| else torch.nonzero(counts == k, as_tuple=False).flatten()) | |
| idx = torch.zeros(len(rows), max(k, 1), dtype=torch.long) | |
| w = torch.zeros(len(rows), max(k, 1)) | |
| for r, row in enumerate(rows.tolist()): | |
| cols = torch.nonzero(nz[row], as_tuple=False).flatten() | |
| idx[r, :len(cols)] = cols | |
| w[r, :len(cols)] = W[row, cols] | |
| groups.append((rows.to(device), idx.to(device), | |
| w.to(device=device, dtype=torch.float32), | |
| b[rows].to(device=device, dtype=torch.float32))) | |
| return groups | |
| def _buffers(self): | |
| import torch | |
| return [torch.zeros(1, n, device=self.device) for n in self.widths] | |
| def _sparse_step(self, v, buf): | |
| x = v | |
| for groups, y in zip(self.plan, buf): | |
| for rows, idx, w, b in groups: | |
| y[:, rows] = ((x[:, idx] * w).sum(-1) + b >= 0).float() | |
| x = y | |
| return x | |
| def _body(self, inp): | |
| return self._sparse_step(inp, self.gbuf) | |
| def _eager(self, v): | |
| return self._sparse_step(v, [self.torch.zeros(v.shape[0], n, | |
| device=self.device) | |
| for n in self.widths]) | |
| def step(self, v): | |
| if self.dense: | |
| for W, b in zip(self.W, self.B): | |
| v = ((v @ W.T + b) >= 0).float() | |
| return v | |
| if self.graph is not None: | |
| return self._replay(v) | |
| return self._eager(v) | |
| def _accumulate_margin(self, v): | |
| y = v | |
| for W, b in zip(self.W, self.B): | |
| pre = y @ W.T + b | |
| m = float((pre + 0.5).abs().min()) | |
| if m < self.min_margin: | |
| self.min_margin = m | |
| y = (pre >= 0).float() | |
| def step_noisy(self, v, sigma: float, gen): | |
| """One step with additive Gaussian read noise per pre-activation and the | |
| comparator at -1/2. | |
| The two forms compute the same pre-activations, so the noise is drawn | |
| for the same quantities whichever is used; the sparse form draws it | |
| layer by layer over that layer's rows. | |
| """ | |
| torch = self.torch | |
| if self.dense: | |
| for W, b in zip(self.W, self.B): | |
| pre = v @ W.T + b | |
| pre = pre + torch.randn(pre.shape, generator=gen, | |
| device=pre.device) * sigma | |
| v = (pre >= -0.5).float() | |
| return v | |
| x = v | |
| for groups, n in zip(self.plan, self.widths): | |
| y = torch.empty(x.shape[0], n, device=x.device) | |
| for rows, idx, w, b in groups: | |
| pre = (x[:, idx] * w).sum(-1) + b | |
| pre = pre + torch.randn(pre.shape, generator=gen, | |
| device=pre.device) * sigma | |
| y[:, rows] = (pre >= -0.5).float() | |
| x = y | |
| return x | |