"""Verified neural INT8 multiplier -- the atom of a GPU tensor core / VNNI lane. A monolithic MLP cannot learn a bit-exact 8x8 multiply (the high-order product bits are too nonlinear). So -- exactly like the byte-slice / ripple trick the other projects use for hard functions -- we shrink the *verified atom* to a 4-bit unsigned multiply and compose everything exactly: atom: NeuralMul4 -- unsigned 4x4 -> 8, domain 16*16 = 256, verified N/N. 8x8: a*b = ah*bh<<8 + (ah*bl + al*bh)<<4 + al*bl (unsigned), exact. signed: a_s = a_u - 256*a7 ; Baugh-Wooley correction, exact integer glue. Only the 4x4 multiply is neural (and N/N-proven); the shifts, adds and sign correction are exact composition. The parallel *throughput* of thousands of such lanes is NOT here -- that needs real silicon (kernel.py, the GUDA-role path). """ from __future__ import annotations import numpy as np import torch from .common import bits_of, int_of, pm, mlp, verify, train class NeuralMul4: """N/N-verified unsigned 4x4 -> 8 multiplier (the finite atom).""" def __init__(self, h: int = 128, layers: int = 3): self.net = mlp(8, 8, h=h, layers=layers) def dataset(self) -> tuple[torch.Tensor, torch.Tensor]: X, Y = [], [] for a in range(16): ab = bits_of(a, 4) for b in range(16): X.append(pm(torch.cat([ab, bits_of(b, 4)]))) Y.append(bits_of(a * b, 8)) return torch.stack(X), torch.stack(Y) def fit(self, steps: int = 4000, lr: float = 2e-3, tag: str = "mul4"): X, Y = self.dataset() train(self.net, X, Y, steps=steps, lr=lr, tag=tag) return self def verify(self) -> tuple[int, int]: X, Y = self.dataset() return verify(self.net, X, Y) @torch.no_grad() def mul(self, a: int, b: int) -> int: self.net.eval() x = pm(torch.cat([bits_of(a & 0xF, 4), bits_of(b & 0xF, 4)])).unsqueeze(0) return int_of((self.net(x)[0] > 0).float()) @torch.no_grad() def mul_array(self, a: np.ndarray, b: np.ndarray) -> np.ndarray: """Batched unsigned 4x4 -> 8 over arrays (one neural forward for all).""" self.net.eval() a = (np.asarray(a).astype(np.int64) & 0xF) b = (np.asarray(b).astype(np.int64) & 0xF) idx = np.arange(4) bits_a = (a[:, None] >> idx) & 1 bits_b = (b[:, None] >> idx) & 1 x = np.concatenate([bits_a, bits_b], axis=1).astype(np.float32) * 2.0 - 1.0 out = (self.net(torch.from_numpy(x)) > 0).to(torch.int64).numpy() from . import instrument instrument.bump("NeuralMul4.forward_calls", 1) instrument.bump("NeuralMul4.products", a.shape[0]) return (out * (1 << np.arange(8))).sum(axis=1) class NeuralMul8: """Signed 8x8 -> 16 multiply, composed exactly from the verified 4x4 atom.""" def __init__(self, h: int = 128, layers: int = 3): self.atom = NeuralMul4(h=h, layers=layers) def fit(self, steps: int = 4000, lr: float = 2e-3, tag: str = "mul4"): self.atom.fit(steps=steps, lr=lr, tag=tag) return self def verify_atom(self) -> tuple[int, int]: return self.atom.verify() def _umul8(self, a_u: int, b_u: int) -> int: """Unsigned 8x8 -> 16 via four verified 4x4 sub-products (exact glue).""" al, ah = a_u & 0xF, (a_u >> 4) & 0xF bl, bh = b_u & 0xF, (b_u >> 4) & 0xF ll = self.atom.mul(al, bl) lh = self.atom.mul(al, bh) hl = self.atom.mul(ah, bl) hh = self.atom.mul(ah, bh) return ll + ((lh + hl) << 4) + (hh << 8) def mul(self, a: int, b: int) -> int: """Signed product via the verified atom + exact Baugh-Wooley correction.""" a_u, b_u = a & 0xFF, b & 0xFF a7, b7 = (a_u >> 7) & 1, (b_u >> 7) & 1 prod = self._umul8(a_u, b_u) - (a7 * b_u << 8) - (b7 * a_u << 8) + (a7 * b7 << 16) prod &= 0xFFFF # 16-bit two's complement return prod - 65536 if prod >= 32768 else prod @torch.no_grad() def mul_array(self, a: np.ndarray, b: np.ndarray) -> np.ndarray: """Batched signed 8x8 -> 16 over arrays, via the verified 4x4 atom. Four nibble sub-products are batched into ONE neural forward, then the exact shift/add glue and Baugh-Wooley sign correction are applied in numpy. Result is bit-exact (the atom is N/N-verified).""" a = np.asarray(a).astype(np.int64).ravel() b = np.asarray(b).astype(np.int64).ravel() au, bu = a & 0xFF, b & 0xFF al, ah = au & 0xF, (au >> 4) & 0xF bl, bh = bu & 0xF, (bu >> 4) & 0xF n = au.shape[0] pa = np.concatenate([al, al, ah, ah]) pb = np.concatenate([bl, bh, bl, bh]) prod = self.atom.mul_array(pa, pb) ll, lh, hl, hh = prod[:n], prod[n:2*n], prod[2*n:3*n], prod[3*n:4*n] u = ll + ((lh + hl) << 4) + (hh << 8) a7, b7 = (au >> 7) & 1, (bu >> 7) & 1 p = (u - (a7 * bu << 8) - (b7 * au << 8) + (a7 * b7 << 16)) & 0xFFFF return np.where(p >= 32768, p - 65536, p) @torch.no_grad() def verify(self, full: bool = True) -> tuple[int, int]: """Exhaustively check the composed signed multiply over all 65536 inputs.""" ok = 0 for a in range(-128, 128): for b in range(-128, 128): if self.mul(a, b) == a * b: ok += 1 return ok, 256 * 256