neural-n64 / mips_units.py
Quazim0t0's picture
N64 engine: 2D SP-DMA fix, libdragon boot, 7 CPU/RSP fixes, SDR fix, CPU fast paths; C cores rebuilt. Portal intact.
c167138 verified
Raw History Blame Contribute Delete
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())
@staticmethod
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.")