neural-n64 / fpu.py
Quazim0t0's picture
Speed up: native MIPS ALU + O(1) decode tables + native FPU shift (bit-identical)
d611ddc verified
Raw
History Blame Contribute Delete
13.1 kB
"""
fpu.py -- IEEE-754 floating point COMPOSED from verified-exact neural units.
Brick 2 of the neural N64: the R4300i's COP1 datapath.
No unit ever sees a domain it wasn't exhaustively verified on:
mantissa add/sub ADD8C / SUB8B slices, carry-rippled
alignment shifts SHR1 slices, sticky collected from the shift-outs
normalization LZC8 per byte + SHL1 slices
multiplication MASK8 partial products + ADC tree
division / sqrt restoring recurrences over SUB8B slices
THE ROUNDING RULE the RND unit: (mode, sign, lsb, guard, round, sticky)
-> round-up bit; 128 verified cases that are the heart
of IEEE-754
Everything else is wiring: bit-field extraction, case selection, exponent
bookkeeping -- the orchestrator stance of every brick before it.
Semantics: the R4300i as modeled by QEMU softfloat (probed empirically):
MIPS legacy NaN: mantissa-MSB SET = signaling; sNaN input -> default NaN
(f32 0x7FBFFFFF / f64 0x7FF7FFFFFFFFFFFF); quiet NaNs propagate unchanged
(a first, then b). Gradual underflow (FS=0). Internal value convention:
unpacked (sign, e, sig) means |x| = sig * 2^(e - M); packing takes sig with
3 GRS bits below: |x| = sig * 2^(e - M - 3).
"""
F32 = dict(E=8, M=23, BIAS=127, DNAN=0x7FBFFFFF, W=32)
F64 = dict(E=11, M=52, BIAS=1023, DNAN=0x7FF7FFFFFFFFFFFF, W=64)
ZERO, NORM, INF, QNAN, SNAN = range(5)
class SoftFP:
def __init__(self, alu, units, rm=0):
self.alu = alu
self.u = units
self.rm = rm # 0 RN, 1 RZ, 2 RP, 3 RM
# mirror the ALU's golden fast-path gate (native == slices, bit-exact)
self.fast = getattr(alu, "fast", False)
# ---------- unit-built helpers ----------
def _nz(self, v, nbits):
if self.fast:
return int((v & ((1 << nbits) - 1)) != 0)
acc = 0
for i in range((nbits + 7) // 8):
acc = self.u.logic8("OR8", acc, (v >> (8 * i)) & 0xFF)
return int(acc != 0)
def _shr_sticky(self, v, k, nbits):
if k <= 0:
return v, 0
if k >= nbits:
return 0, self._nz(v, nbits)
if self.fast: # native, bit-identical
return v >> k, int((v & ((1 << k) - 1)) != 0)
sticky = 0
nb = (nbits + 7) // 8
for _ in range(k):
c = 0; out = 0
for i in reversed(range(nb)):
r, c = self.u.shr1((v >> (8 * i)) & 0xFF, c)
out |= r << (8 * i)
sticky |= c
v = out
return v, sticky
def _lzc(self, v, nbits):
nb = (nbits + 7) // 8
total = 0
for i in reversed(range(nb)):
l = self.u.lzc8((v >> (8 * i)) & 0xFF)
total += l
if l < 8:
break
return total - (nb * 8 - nbits)
# ---------- unpack / classify ----------
def unpack(self, x, f):
E, M = f["E"], f["M"]
sign = (x >> (E + M)) & 1
e = (x >> M) & ((1 << E) - 1)
man = x & ((1 << M) - 1)
if e == (1 << E) - 1:
if not self._nz(man, M):
return INF, sign, 0, 0
return (SNAN if (man >> (M - 1)) & 1 else QNAN), sign, 0, man
if e == 0:
if not self._nz(man, M):
return ZERO, sign, 0, 0
l = self._lzc(man, M + 1) # zeros above implicit slot
sig = self.alu.shl(man, l, M + 1)
return NORM, sign, 1 - l, sig # e biased-adjusted below
return NORM, sign, e, man | (1 << M)
def _unb(self, cls, e, f):
"""unpack returns biased-ish e for NORM; convert to unbiased."""
return e - f["BIAS"]
def _nan2(self, ca, a, cb, b, f):
if ca == SNAN or cb == SNAN:
return f["DNAN"]
return a if ca == QNAN else b
# ---------- round + pack ----------
def _pack_round(self, sign, e_unb, sig, f):
"""sig: significand with 3 GRS bits below target; |x| = sig*2^(e-M-3)."""
E, M = f["E"], f["M"]
maxe = (1 << E) - 1
width = M + 4 # implicit at bit M+3
if sig >> width:
extra = sig.bit_length() - width
sig, st = self._shr_sticky(sig, extra, width + extra)
sig |= st
e_unb += extra
elif not (sig >> (width - 1)):
l = self._lzc(sig, width)
sig = self.alu.shl(sig, l, width)
e_unb -= l
biased = e_unb + f["BIAS"]
if biased <= 0:
sig, st = self._shr_sticky(sig, 1 - biased, width)
sig |= st
biased = 0
lsb, g, r, s = (sig >> 3) & 1, (sig >> 2) & 1, (sig >> 1) & 1, sig & 1
up = self.u.rnd(self.rm, sign, lsb, g, r, s)
mant, _ = self.alu.add(sig >> 3, up, M + 9)
if mant >> (M + 1):
mant >>= 1
biased += 1
if biased == 0 and (mant >> M):
biased = 1
if biased >= maxe:
if self.rm == 1 or (self.rm == 2 and sign) or (self.rm == 3 and not sign):
return (sign << (E + M)) | ((maxe - 1) << M) | ((1 << M) - 1)
return (sign << (E + M)) | (maxe << M)
return (sign << (E + M)) | (biased << M) | (mant & ((1 << M) - 1))
# ---------- arithmetic ----------
def add(self, a, b, f, sub=False):
E, M = f["E"], f["M"]
a0, b0 = a, b
if sub:
b ^= 1 << (E + M)
ca, sa, ea, ma = self.unpack(a, f)
cb, sb, eb, mb = self.unpack(b, f)
if ca in (QNAN, SNAN) or cb in (QNAN, SNAN):
return self._nan2(ca, a0, cb, b0, f) # NaN payload unflipped
if ca == INF and cb == INF:
return f["DNAN"] if sa != sb else a
if ca == INF: return a
if cb == INF: return b
if ca == ZERO and cb == ZERO:
return ((sa & sb) if self.rm != 3 else (sa | sb)) << (E + M)
if ca == ZERO: return b
if cb == ZERO: return a
ea, eb = self._unb(ca, ea, f), self._unb(cb, eb, f)
ma <<= 3; mb <<= 3 # GRS space
width = M + 12
if (ea, ma) < (eb, mb) if ea != eb else ma < mb:
sa, sb, ea, eb, ma, mb = sb, sa, eb, ea, mb, ma
if ea < eb:
sa, sb, ea, eb, ma, mb = sb, sa, eb, ea, mb, ma
elif ea == eb and ma < mb:
sa, sb, ma, mb = sb, sa, mb, ma
mb, st = self._shr_sticky(mb, ea - eb, width)
mb |= st
if sa == sb:
sig, _ = self.alu.add(ma, mb, width)
else:
sig, _ = self.alu.sub(ma, mb, width)
if not self._nz(sig, width):
return (1 if self.rm == 3 else 0) << (E + M)
return self._pack_round(sa, ea, sig, f)
def mul(self, a, b, f):
E, M = f["E"], f["M"]
ca, sa, ea, ma = self.unpack(a, f)
cb, sb, eb, mb = self.unpack(b, f)
sign = sa ^ sb
if ca in (QNAN, SNAN) or cb in (QNAN, SNAN):
return self._nan2(ca, a, cb, b, f)
if (ca == INF and cb == ZERO) or (ca == ZERO and cb == INF):
return f["DNAN"]
if ca == INF or cb == INF:
return (sign << (E + M)) | (((1 << E) - 1) << M)
if ca == ZERO or cb == ZERO:
return sign << (E + M)
ea, eb = self._unb(ca, ea, f), self._unb(cb, eb, f)
nb = ((M + 1 + 7) // 8) * 8
full = self.alu.mul(ma, mb, nb) # (M+1)x(M+1) -> <=2M+2 bits
# |a*b| = full * 2^(ea+eb-2M); want sig*2^(e-M-3): sig=full>>(M-2), e=ea+eb+1
sig, st = self._shr_sticky(full, M - 2, 2 * nb)
sig |= st
return self._pack_round(sign, ea + eb + 1, sig, f)
def div(self, a, b, f):
E, M = f["E"], f["M"]
ca, sa, ea, ma = self.unpack(a, f)
cb, sb, eb, mb = self.unpack(b, f)
sign = sa ^ sb
if ca in (QNAN, SNAN) or cb in (QNAN, SNAN):
return self._nan2(ca, a, cb, b, f)
if (ca == INF and cb == INF) or (ca == ZERO and cb == ZERO):
return f["DNAN"]
if ca == INF or cb == ZERO:
return (sign << (E + M)) | (((1 << E) - 1) << M)
if ca == ZERO or cb == INF:
return sign << (E + M)
ea, eb = self._unb(ca, ea, f), self._unb(cb, eb, f)
e = ea - eb
if ma < mb:
ma <<= 1
e -= 1
# ma/mb in [1,2): q = (ma<<(M+3))/mb has implicit at M+3 -> e unchanged
q, rem = self.alu.divmod_(ma << (M + 3), mb, 2 * M + 8)
if self._nz(rem, M + 8):
q |= 1
return self._pack_round(sign, e, q, f)
def sqrt(self, a, f):
E, M = f["E"], f["M"]
ca, sa, ea, ma = self.unpack(a, f)
if ca in (QNAN, SNAN):
return f["DNAN"] if ca == SNAN else a
if ca == ZERO:
return sa << (E + M)
if sa:
return f["DNAN"]
if ca == INF:
return a
ea = self._unb(ca, ea, f)
if ea & 1:
ma <<= 1
ea -= 1
e = ea >> 1
# x = ma << (M+6): isqrt(x) = sqrt(ma)*2^((M+6)/2)... handled by exact
# pairing: result n=M+4 bits, x has 2n bits, |res| = sqrt(ma)*2^(M/2+3)
# and |sqrt| = sqrt(ma)*2^((ea-M)/2) = res * 2^(e-M-3) [ea = 2e]
n = M + 4
x = ma << (M + 6)
res = 0; rem = 0
wb = 2 * n + 8
for i in reversed(range(n)):
rem = (rem << 2) | ((x >> (2 * i)) & 3)
trial = (res << 2) | 1
diff, borrow = self.alu.sub(rem, trial, wb)
if not borrow:
rem = diff
res = (res << 1) | 1
else:
res <<= 1
if self._nz(rem, wb):
res |= 1
return self._pack_round(0, e, res, f)
# ---------- conversions ----------
def cvt_f2f(self, a, src, dst):
ca, sa, ea, ma = self.unpack(a, src)
E, M, sM = dst["E"], dst["M"], src["M"]
if ca in (QNAN, SNAN):
if ca == SNAN:
return dst["DNAN"]
if dst["M"] > sM: # widen qNaN payload
man = (a & ((1 << sM) - 1)) << (M - sM)
else: # narrow
man = (a & ((1 << sM) - 1)) >> (sM - M)
if not self._nz(man, M) or (man >> (M - 1)) & 1:
return dst["DNAN"]
return (sa << (E + M)) | (((1 << E) - 1) << M) | man
if ca == INF:
return (sa << (E + M)) | (((1 << E) - 1) << M)
if ca == ZERO:
return sa << (E + M)
ea = self._unb(ca, ea, src)
if M >= sM:
sig = ma << (M - sM + 3)
else:
sig, st = self._shr_sticky(ma << 3, sM - M, sM + 4)
sig |= st
return self._pack_round(sa, ea, sig, dst)
def f2i(self, a, f, width, rm):
"""float -> signed int; invalid (NaN/inf/overflow) -> MIPS INT_MAX."""
ca, sa, ea, ma = self.unpack(a, f)
M = f["M"]
lim = 1 << (width - 1)
if ca in (QNAN, SNAN, INF):
return lim - 1
if ca == ZERO:
return 0
ea = self._unb(ca, ea, f)
if ea >= width:
return lim - 1
if ea >= M:
v = ma << (ea - M)
g = s = 0
else:
k = M - ea
v, _ = self._shr_sticky(ma, k, M + 1)
g = (ma >> (k - 1)) & 1 if k - 1 <= M else 0
below = ma & ((1 << max(k - 1, 0)) - 1)
s = self._nz(below, M + 1)
up = self.u.rnd(rm, sa, v & 1, g, 0, s)
v, _ = self.alu.add(v, up, width + 16)
if v > (lim if sa else lim - 1):
return lim - 1
return ((-v) & ((1 << width) - 1)) if sa else v
def i2f(self, v, width, f):
M = f["M"]
sign = (v >> (width - 1)) & 1
mag = ((1 << width) - v) & ((1 << width) - 1) if sign else v
if mag == 0:
return 0
l = self._lzc(mag, width)
top = width - 1 - l # |v| = mag, msb at top
if top > M + 3:
sig, st = self._shr_sticky(mag, top - (M + 3), width)
sig |= st
else:
sig = self.alu.shl(mag, (M + 3) - top, M + 4)
return self._pack_round(sign, top, sig, f)
# ---------- compare ----------
def cmp(self, a, b, f):
"""(less, equal, unordered)"""
ca = self.unpack(a, f)[0]
cb = self.unpack(b, f)[0]
if ca in (QNAN, SNAN) or cb in (QNAN, SNAN):
return 0, 0, 1
if ca == ZERO and cb == ZERO:
return 0, 1, 0
W = f["W"]
def key(x):
s = x >> (W - 1)
mag = x & ((1 << (W - 1)) - 1)
return ((1 << (W - 1)) + mag) if not s else ((1 << (W - 1)) - mag)
ka, kb = key(a), key(b)
if ka == kb:
return 0, 1, 0
_, borrow = self.alu.sub(ka, kb, W + 8)
return borrow, 0, 0