Spaces:
Sleeping
Sleeping
Download mips_units.py from build-small-hackathon/neural-n64: direct link, hf CLI and curl.
- Browser
- Download file 17.1 kB
-
https://huggingface.co/spaces/build-small-hackathon/neural-n64/resolve/main/mips_units.py
- Command line
-
hf download hf://spaces/build-small-hackathon/neural-n64/mips_units.py
-
curl -L -o mips_units.py https://huggingface.co/spaces/build-small-hackathon/neural-n64/resolve/main/mips_units.py
17.1 kB
| """ | |
| mips_units.py -- the N64's R4300i (MIPS III) integer datapath as verified-exact | |
| neural units. Brick 1 of the neural N64. | |
| MIPS is kind to the methodology: NO flags register (the x86 unit's hardest | |
| outputs simply don't exist here), fixed 32-bit instruction words, and a decode | |
| space that is three tiny enumerable tables (opcode[6 bits], funct[6 bits], | |
| regimm[5 bits]). Wide arithmetic is composed exactly as before: | |
| 32/64-bit add/sub = ADD8C slices, carry rippled (8 slices for 64-bit) | |
| logic = AND8/OR8/XOR8/NOR8 slices | |
| comparisons (SLT*) = subtract ripple + sign/borrow wiring | |
| shifts by k = k applications of 1-bit slice units | |
| MULT/DIV = MASK8 partial products / restoring division | |
| Every unit is trained on its COMPLETE domain and passes only at N/N. | |
| """ | |
| import os, time | |
| import numpy as np, torch, torch.nn as nn | |
| torch.manual_seed(0); np.random.seed(0); torch.set_num_threads(1) | |
| CACHE = os.path.join(os.path.dirname(os.path.abspath(__file__)), "models") | |
| os.makedirs(CACHE, exist_ok=True) | |
| def bits(x, n): | |
| x = np.asarray(x) | |
| return ((x[:, None] >> np.arange(n)[None]) & 1).astype(np.float32) | |
| class BitMLP(nn.Module): | |
| def __init__(self, nin, nout, h=256): | |
| super().__init__() | |
| self.net = nn.Sequential(nn.Linear(nin, h), nn.ReLU(), | |
| nn.Linear(h, h), nn.ReLU(), nn.Linear(h, nout)) | |
| def forward(self, x): return self.net(x) | |
| def _ok(net, X, Y): | |
| with torch.no_grad(): | |
| for i in range(0, X.shape[0], 65536): | |
| if not bool((((net(X[i:i+65536]) > 0).float() == Y[i:i+65536])).all()): | |
| return False | |
| return True | |
| def train_verify(name, X, Y, lr=3e-3, bs=4096, max_steps=60000, cap=600): | |
| X = torch.tensor(X); Y = torch.tensor(Y) | |
| net = BitMLP(X.shape[1], Y.shape[1]) | |
| path = os.path.join(CACHE, f"{name}.pt") | |
| if os.path.exists(path): | |
| net.load_state_dict(torch.load(path, weights_only=True)) | |
| if _ok(net, X, Y): | |
| print(f" {name:8s} {X.shape[0]:6d}/{X.shape[0]:<6d} EXACT (cached)", flush=True) | |
| return net | |
| opt = torch.optim.Adam(net.parameters(), lr) | |
| lossf = nn.BCEWithLogitsLoss(); N = X.shape[0]; t0 = time.time() | |
| for s in range(1, max_steps + 1): | |
| idx = torch.randint(0, N, (min(bs, N),)) | |
| opt.zero_grad(); lossf(net(X[idx]), Y[idx]).backward(); opt.step() | |
| if s % 500 == 0 and (_ok(net, X, Y) or time.time() - t0 > cap): | |
| break | |
| ok = _ok(net, X, Y) | |
| print(f" {name:8s} {X.shape[0]:6d}/{X.shape[0]:<6d} " | |
| f"{'EXACT' if ok else 'NOT EXACT'} ({time.time()-t0:.0f}s)", flush=True) | |
| if not ok: | |
| raise RuntimeError(f"unit {name} failed") | |
| torch.save(net.state_dict(), path) | |
| return net | |
| # ===================== golden slice semantics ===================== | |
| def u_ADD8C(): # a + b + cin -> res(8) + cout (no flags: this is MIPS) | |
| i = np.arange(131072); a, b, c = i >> 9, (i >> 1) & 0xFF, i & 1 | |
| t = a + b + c | |
| X = np.concatenate([bits(a, 8), bits(b, 8), c[:, None].astype(np.float32)], 1) | |
| return X, np.concatenate([bits(t & 0xFF, 8), | |
| (t > 0xFF)[:, None].astype(np.float32)], 1) | |
| def u_SUB8B(): # a - b - bin -> res(8) + borrow | |
| i = np.arange(131072); a, b, c = i >> 9, (i >> 1) & 0xFF, i & 1 | |
| t = a - b - c | |
| X = np.concatenate([bits(a, 8), bits(b, 8), c[:, None].astype(np.float32)], 1) | |
| return X, np.concatenate([bits(t & 0xFF, 8), | |
| (t < 0)[:, None].astype(np.float32)], 1) | |
| def u_logic(f): | |
| a = np.repeat(np.arange(256), 256); b = np.tile(np.arange(256), 256) | |
| return (np.concatenate([bits(a, 8), bits(b, 8)], 1), | |
| bits(f(a, b) & 0xFF, 8)) | |
| def u_SHL1(): | |
| v = np.repeat(np.arange(256), 2); c = np.tile(np.arange(2), 256) | |
| X = np.concatenate([bits(v, 8), c[:, None].astype(np.float32)], 1) | |
| return X, np.concatenate([bits(((v << 1) | c) & 0xFF, 8), | |
| (v >> 7)[:, None].astype(np.float32)], 1) | |
| def u_SHR1(): | |
| v = np.repeat(np.arange(256), 2); c = np.tile(np.arange(2), 256) | |
| X = np.concatenate([bits(v, 8), c[:, None].astype(np.float32)], 1) | |
| return X, np.concatenate([bits((v >> 1) | (c << 7), 8), | |
| (v & 1)[:, None].astype(np.float32)], 1) | |
| def u_MASK8(): | |
| v = np.repeat(np.arange(256), 2); c = np.tile(np.arange(2), 256) | |
| X = np.concatenate([bits(v, 8), c[:, None].astype(np.float32)], 1) | |
| return X, bits(v * c, 8) | |
| def u_LZC8(): # byte -> leading-zero count (0..8), for FP normalization | |
| v = np.arange(256) | |
| lzc = np.array([8 if x == 0 else 8 - x.bit_length() for x in range(256)]) | |
| return bits(v, 8), bits(lzc, 4) | |
| def u_RND(): # IEEE-754 round-up decision: (mode2, sign, lsb, g, r, s) -> up | |
| rows, outs = [], [] | |
| for mode in range(4): # 0 RN, 1 RZ, 2 RP(+inf), 3 RM(-inf) | |
| for sign in range(2): | |
| for lsb in range(2): | |
| for g in range(2): | |
| for r in range(2): | |
| for s in range(2): | |
| if mode == 0: up = g & (r | s | lsb) | |
| elif mode == 1: up = 0 | |
| elif mode == 2: up = (1 - sign) & (g | r | s) | |
| else: up = sign & (g | r | s) | |
| rows.append(list(bits([mode], 2)[0]) + [sign, lsb, g, r, s]) | |
| outs.append([up]) | |
| return np.array(rows, np.float32), np.array(outs, np.float32) | |
| # ===================== decode tables (MIPS: pure enumerable) ===================== | |
| # major opcode -> class id (table below); SPECIAL/REGIMM dispatch to sub-tables | |
| OPC = { | |
| 0x00: "SPECIAL", 0x01: "REGIMM", 0x02: "J", 0x03: "JAL", | |
| 0x04: "BEQ", 0x05: "BNE", 0x06: "BLEZ", 0x07: "BGTZ", | |
| 0x08: "ADDI", 0x09: "ADDIU", 0x0A: "SLTI", 0x0B: "SLTIU", | |
| 0x0C: "ANDI", 0x0D: "ORI", 0x0E: "XORI", 0x0F: "LUI", | |
| 0x10: "COP0", 0x11: "COP1", 0x14: "BEQL", 0x15: "BNEL", | |
| 0x16: "BLEZL", 0x17: "BGTZL", 0x18: "DADDI", 0x19: "DADDIU", | |
| 0x1A: "LDL", 0x1B: "LDR", | |
| 0x20: "LB", 0x21: "LH", 0x22: "LWL", 0x23: "LW", | |
| 0x24: "LBU", 0x25: "LHU", 0x26: "LWR", 0x27: "LWU", | |
| 0x28: "SB", 0x29: "SH", 0x2A: "SWL", 0x2B: "SW", | |
| 0x2C: "SDL", 0x2D: "SDR", 0x2E: "SWR", 0x2F: "CACHE", | |
| 0x30: "LL", 0x31: "LWC1", 0x34: "LLD", 0x35: "LDC1", 0x37: "LD", | |
| 0x38: "SC", 0x39: "SWC1", 0x3C: "SCD", 0x3D: "SDC1", 0x3F: "SD", | |
| 0x12: "COP2", 0x32: "LWC2", 0x3A: "SWC2", | |
| } | |
| FUNCT = { | |
| 0x00: "SLL", 0x02: "SRL", 0x03: "SRA", 0x04: "SLLV", 0x06: "SRLV", | |
| 0x07: "SRAV", 0x08: "JR", 0x09: "JALR", 0x0C: "SYSCALL", 0x0D: "BREAK", | |
| 0x0F: "SYNC", 0x10: "MFHI", 0x11: "MTHI", 0x12: "MFLO", 0x13: "MTLO", | |
| 0x14: "DSLLV", 0x16: "DSRLV", 0x17: "DSRAV", | |
| 0x18: "MULT", 0x19: "MULTU", 0x1A: "DIV", 0x1B: "DIVU", | |
| 0x1C: "DMULT", 0x1D: "DMULTU", 0x1E: "DDIV", 0x1F: "DDIVU", | |
| 0x20: "ADD", 0x21: "ADDU", 0x22: "SUB", 0x23: "SUBU", | |
| 0x24: "AND", 0x25: "OR", 0x26: "XOR", 0x27: "NOR", | |
| 0x2A: "SLT", 0x2B: "SLTU", 0x2C: "DADD", 0x2D: "DADDU", | |
| 0x2E: "DSUB", 0x2F: "DSUBU", | |
| 0x34: "TEQ", | |
| 0x38: "DSLL", 0x3A: "DSRL", 0x3B: "DSRA", | |
| 0x3C: "DSLL32", 0x3E: "DSRL32", 0x3F: "DSRA32", | |
| } | |
| REGIMM = {0x00: "BLTZ", 0x01: "BGEZ", 0x02: "BLTZL", 0x03: "BGEZL", | |
| 0x10: "BLTZAL", 0x11: "BGEZAL"} | |
| OPC_CLASSES = ["ILL"] + sorted(set(OPC.values())) | |
| FUNCT_CLASSES = ["ILL"] + sorted(set(FUNCT.values())) | |
| REGIMM_CLASSES = ["ILL"] + sorted(set(REGIMM.values())) | |
| # precomputed value->class-index tables (decode is a 6/5-bit lookup; avoid the | |
| # per-instruction linear list.index scan). Bit-identical to the .index() result. | |
| OPC_IDX = [OPC_CLASSES.index(OPC.get(v, "ILL")) for v in range(64)] | |
| FUNCT_IDX = [FUNCT_CLASSES.index(FUNCT.get(v, "ILL")) for v in range(64)] | |
| REGIMM_IDX = [REGIMM_CLASSES.index(REGIMM.get(v, "ILL")) for v in range(32)] | |
| def _table_unit(name, table, classes, nbits): | |
| rows = np.arange(1 << nbits) | |
| cls = np.array([classes.index(table.get(int(v), "ILL")) for v in rows]) | |
| ncls = len(classes) | |
| X = torch.tensor(bits(rows, nbits)) | |
| Yc = torch.tensor(cls) | |
| net = nn.Sequential(nn.Linear(nbits, 128), nn.ReLU(), nn.Linear(128, 128), | |
| nn.ReLU(), nn.Linear(128, ncls)) | |
| path = os.path.join(CACHE, f"{name}.pt") | |
| def check(): | |
| with torch.no_grad(): | |
| return bool((net(X).argmax(1) == Yc).all()) | |
| if os.path.exists(path): | |
| net.load_state_dict(torch.load(path, weights_only=True)) | |
| if check(): | |
| print(f" {name:8s} {len(rows):6d}/{len(rows):<6d} EXACT (cached)", flush=True) | |
| return net | |
| opt = torch.optim.Adam(net.parameters(), 3e-3); ce = nn.CrossEntropyLoss() | |
| t0 = time.time() | |
| for s in range(1, 20001): | |
| l = ce(net(X), Yc) | |
| opt.zero_grad(); l.backward(); opt.step() | |
| if s % 200 == 0 and (check() or time.time() - t0 > 60): | |
| break | |
| ok = check() | |
| print(f" {name:8s} {len(rows):6d}/{len(rows):<6d} " | |
| f"{'EXACT' if ok else 'NOT EXACT'} ({time.time()-t0:.0f}s)", flush=True) | |
| if not ok: | |
| raise RuntimeError(name) | |
| torch.save(net.state_dict(), path) | |
| return net | |
| def build_all(): | |
| print("Training/loading + exhaustively verifying MIPS R4300i units:", flush=True) | |
| u = {} | |
| u["ADD8C"] = train_verify("ADD8C", *u_ADD8C()) | |
| u["SUB8B"] = train_verify("SUB8B", *u_SUB8B()) | |
| u["AND8"] = train_verify("AND8", *u_logic(lambda a, b: a & b)) | |
| u["OR8"] = train_verify("OR8", *u_logic(lambda a, b: a | b)) | |
| u["XOR8"] = train_verify("XOR8", *u_logic(lambda a, b: a ^ b)) | |
| u["NOR8"] = train_verify("NOR8", *u_logic(lambda a, b: ~(a | b))) | |
| u["SHL1"] = train_verify("SHL1", *u_SHL1(), bs=512, cap=60) | |
| u["SHR1"] = train_verify("SHR1", *u_SHR1(), bs=512, cap=60) | |
| u["MASK8"] = train_verify("MASK8", *u_MASK8(), bs=512, cap=60) | |
| u["LZC8"] = train_verify("LZC8", *u_LZC8(), bs=256, cap=40) | |
| u["RND"] = train_verify("RND", *u_RND(), bs=64, cap=40) | |
| u["OPC"] = _table_unit("OPC", OPC, OPC_CLASSES, 6) | |
| u["FUNCT"] = _table_unit("FUNCT", FUNCT, FUNCT_CLASSES, 6) | |
| u["REGIMM"] = _table_unit("REGIMM", REGIMM, REGIMM_CLASSES, 5) | |
| return u | |
| # ===================== unit APIs ===================== | |
| class GoldenUnits: | |
| def add8c(self, a, b, c): | |
| t = a + b + c; return t & 0xFF, int(t > 0xFF) | |
| def sub8b(self, a, b, c): | |
| t = a - b - c; return t & 0xFF, int(t < 0) | |
| def logic8(self, kind, a, b): | |
| return {"AND8": a & b, "OR8": a | b, "XOR8": a ^ b, | |
| "NOR8": ~(a | b) & 0xFF}[kind] & 0xFF | |
| def shl1(self, v, c): return ((v << 1) | c) & 0xFF, v >> 7 | |
| def shr1(self, v, c): return (v >> 1) | (c << 7), v & 1 | |
| def mask8(self, v, bit): return v * bit | |
| def lzc8(self, v): return 8 if v == 0 else 8 - v.bit_length() | |
| def rnd(self, mode, sign, lsb, g, r, s): | |
| if mode == 0: return g & (r | s | lsb) | |
| if mode == 1: return 0 | |
| if mode == 2: return (1 - sign) & (g | r | s) | |
| return sign & (g | r | s) | |
| def opc(self, v): return OPC_IDX[v] | |
| def funct(self, v): return FUNCT_IDX[v] | |
| def regimm(self, v): return REGIMM_IDX[v] | |
| class NeuralUnits: | |
| def __init__(self, nets): self.n = nets | |
| def _run(self, name, ib): | |
| with torch.no_grad(): | |
| o = self.n[name](torch.tensor(np.array([ib], np.float32)))[0] | |
| return (o > 0).long().numpy() | |
| def _cls(self, name, v, nbits): | |
| with torch.no_grad(): | |
| o = self.n[name](torch.tensor(bits([v], nbits)))[0] | |
| return int(o.argmax()) | |
| def _i(o, a, n): return int(sum(int(o[a + i]) << i for i in range(n))) | |
| def add8c(self, a, b, c): | |
| o = self._run("ADD8C", list(bits([a], 8)[0]) + list(bits([b], 8)[0]) + [c]) | |
| return self._i(o, 0, 8), int(o[8]) | |
| def sub8b(self, a, b, c): | |
| o = self._run("SUB8B", list(bits([a], 8)[0]) + list(bits([b], 8)[0]) + [c]) | |
| return self._i(o, 0, 8), int(o[8]) | |
| def logic8(self, kind, a, b): | |
| o = self._run(kind, list(bits([a], 8)[0]) + list(bits([b], 8)[0])) | |
| return self._i(o, 0, 8) | |
| def shl1(self, v, c): | |
| o = self._run("SHL1", list(bits([v], 8)[0]) + [c]) | |
| return self._i(o, 0, 8), int(o[8]) | |
| def shr1(self, v, c): | |
| o = self._run("SHR1", list(bits([v], 8)[0]) + [c]) | |
| return self._i(o, 0, 8), int(o[8]) | |
| def mask8(self, v, bit): | |
| return self._i(self._run("MASK8", list(bits([v], 8)[0]) + [bit]), 0, 8) | |
| def lzc8(self, v): | |
| return self._i(self._run("LZC8", list(bits([v], 8)[0])), 0, 4) | |
| def rnd(self, mode, sign, lsb, g, r, s): | |
| return int(self._run("RND", list(bits([mode], 2)[0]) + [sign, lsb, g, r, s])[0]) | |
| def opc(self, v): return self._cls("OPC", v, 6) | |
| def funct(self, v): return self._cls("FUNCT", v, 6) | |
| def regimm(self, v): return self._cls("REGIMM", v, 5) | |
| # ===================== composed wide ops (wiring) ===================== | |
| class ALU: | |
| """Composed 32/64-bit MIPS datapath over the slices. | |
| With the golden reference units the bit-sliced loops are just a slow exact | |
| mirror of native integer arithmetic, so we take native fast paths (proven | |
| bit-identical to the slices). With neural units (audit runs) we stay on the | |
| slices so the verified nets are exercised. nbits is always a multiple of 8. | |
| """ | |
| def __init__(self, u): | |
| self.u = u | |
| self.fast = isinstance(u, GoldenUnits) | |
| def add(self, a, b, nbits, cin=0): | |
| M = ((nbits + 7) // 8) * 8 | |
| if self.fast: | |
| t = a + b + cin | |
| return t & ((1 << M) - 1), (t >> M) & 1 | |
| c = cin; out = 0 | |
| for i in range((nbits + 7) // 8): | |
| r, c = self.u.add8c((a >> (8*i)) & 0xFF, (b >> (8*i)) & 0xFF, c) | |
| out |= r << (8 * i) | |
| return out, c | |
| def sub(self, a, b, nbits): | |
| M = ((nbits + 7) // 8) * 8 | |
| if self.fast: | |
| t = a - b | |
| return t & ((1 << M) - 1), int(t < 0) | |
| c = 0; out = 0 | |
| for i in range((nbits + 7) // 8): | |
| r, c = self.u.sub8b((a >> (8*i)) & 0xFF, (b >> (8*i)) & 0xFF, c) | |
| out |= r << (8 * i) | |
| return out, c # c = borrow | |
| def logic(self, kind, a, b, nbits): | |
| if self.fast: | |
| M = ((nbits + 7) // 8) * 8 | |
| r = {"AND8": a & b, "OR8": a | b, "XOR8": a ^ b, | |
| "NOR8": ~(a | b)}[kind] | |
| return r & ((1 << M) - 1) | |
| out = 0 | |
| for i in range((nbits + 7) // 8): | |
| out |= self.u.logic8(kind, (a >> (8*i)) & 0xFF, (b >> (8*i)) & 0xFF) << (8*i) | |
| return out | |
| def shl(self, v, k, nbits): | |
| if self.fast: | |
| M = ((nbits + 7) // 8) * 8 | |
| return (v << k) & ((1 << M) - 1) | |
| for _ in range(k): | |
| c = 0; out = 0 | |
| for i in range((nbits + 7) // 8): | |
| r, c = self.u.shl1((v >> (8*i)) & 0xFF, c) | |
| out |= r << (8 * i) | |
| v = out | |
| return v | |
| def shr(self, v, k, nbits, arith=False): | |
| if self.fast: | |
| M = ((nbits + 7) // 8) * 8 | |
| if arith and (v >> (nbits - 1)) & 1: | |
| sv = v - (1 << nbits) # treat as signed nbits | |
| return (sv >> k) & ((1 << M) - 1) | |
| return (v >> k) & ((1 << M) - 1) | |
| for _ in range(k): | |
| c = 1 if (arith and (v >> (nbits - 1)) & 1) else 0 | |
| out = 0 | |
| for i in reversed(range((nbits + 7) // 8)): | |
| r, c = self.u.shr1((v >> (8*i)) & 0xFF, c) | |
| out |= r << (8 * i) | |
| v = out | |
| return v | |
| def mul(self, a, b, nbits): | |
| """unsigned a*b -> 2*nbits, masked shifted adds.""" | |
| if self.fast: # a,b < 2^nbits -> exact | |
| return (a * b) & ((1 << (2 * nbits)) - 1) | |
| acc = 0 | |
| for j in range(nbits): | |
| bit = (b >> j) & 1 | |
| masked = 0 | |
| for i in range((nbits + 7) // 8): | |
| masked |= self.u.mask8((a >> (8*i)) & 0xFF, bit) << (8*i) | |
| acc, _ = self.add(acc, masked << j, 2 * nbits) | |
| return acc | |
| def divmod_(self, num, den, nbits): | |
| if den == 0: | |
| return None, None # MIPS DIV by zero: UNPREDICTABLE; core handles | |
| if self.fast: # unsigned, exact | |
| return num // den, num % den | |
| q = 0; rem = 0 | |
| for i in reversed(range(nbits)): | |
| rem = (rem << 1) | ((num >> i) & 1) | |
| diff, borrow = self.sub(rem, den, nbits + 8) | |
| if not borrow: | |
| rem = diff; q |= 1 << i | |
| return q, rem | |
| if __name__ == "__main__": | |
| build_all() | |
| print("\nALL MIPS R4300i UNITS VERIFIED EXACT.") | |