neural-ddr / step1_train_verify.py
Quazim0t0's picture
Import from Quazim0t0/neural-ddr; repoint refs to NeuralVerified
141e646 verified
Raw
History Blame Contribute Delete
3.89 kB
"""
Step 1 for neural-DDR: train + exhaustively verify the DBI units, then prove the
memory bridge round-trips real data through them.
(a) encode N/N -> DBI encode matches golden on all 256 bytes
(b) decode N/N -> DBI decode matches golden on all 512 (data,flag) inputs
(c) round-trip -> bridge stores every byte at many addresses via the neural
units and reads it back bit-exact (data integrity)
(d) DBI works -> DDR4/5 bridge never drives >4 DQ lines LOW; DDR3 does
(e) generation -> same host RAM, DDR3 vs DDR5 behavior, both correct
"""
import torch
from neural_ddr.dbi import (
NeuralDBIEncode, NeuralDBIDecode, encode_domain, decode_domain,
verify_encode, verify_decode,
)
from neural_ddr.bridge import MemoryBridge
torch.manual_seed(0)
dev = "cuda" if torch.cuda.is_available() else "cpu"
def train(unit, X, Y, steps=12000, lr=2e-3, tag="", verify=None):
X, Y = X.to(dev), Y.to(dev)
unit = unit.to(dev)
opt = torch.optim.Adam(unit.parameters(), lr=lr)
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=steps)
for epoch in range(steps):
opt.zero_grad()
loss = (torch.tanh(unit(X)) - torch.tanh(Y * 4)).pow(2).mean()
loss.backward(); opt.step(); sched.step()
if epoch % 2000 == 0 or epoch == steps - 1:
n_ok, n_tot = verify(unit.cpu()); unit.to(dev)
print(f" [{tag}] epoch {epoch:5d} loss {loss.item():.2e} verified {n_ok}/{n_tot}")
if n_ok == n_tot:
print(f" [{tag}] -> full N/N reached")
break
return unit.cpu()
if __name__ == "__main__":
print("=" * 60)
print("STEP 1 -- neural DBI units + memory bridge")
print("=" * 60)
Xe, Ye = encode_domain()
Xd, Yd = decode_domain()
print("training DBI encode (8 -> 9):")
enc = train(NeuralDBIEncode(), Xe, Ye, tag="enc", verify=verify_encode)
print("training DBI decode (9 -> 8):")
dec = train(NeuralDBIDecode(), Xd, Yd, tag="dec", verify=verify_decode)
ne, te = verify_encode(enc)
nd, td = verify_decode(dec)
print("-" * 60)
print(f"(a) DBI encode verified {ne}/{te} -> {'PASS' if ne==te else 'FAIL'}")
print(f"(b) DBI decode verified {nd}/{td} -> {'PASS' if nd==td else 'FAIL'}")
# (c) round-trip integrity through the bridge (DDR5 mode)
N = 4096
br5 = MemoryBridge(N, enc, dec, generation="DDR5")
bad = 0
for a in range(N):
v = (a * 37 + 11) & 0xFF
br5.write(a, v)
for a in range(N):
v = (a * 37 + 11) & 0xFF
bad += (br5.read(a) != v)
print(f"(c) bridge round-trip {N-bad}/{N} bytes exact -> {'PASS' if bad==0 else 'FAIL'}")
# (d) DBI bus effect: DDR5 caps DQ-low at 4/byte; DDR3 does not
br3 = MemoryBridge(256, enc, dec, generation="DDR3")
worst5 = 0
for v in range(256):
br5b = MemoryBridge(1, enc, dec, generation="DDR5"); br5b.write(0, v)
worst5 = max(worst5, int(br5b.dq_low_total))
br3.write(v % 256, v)
print(f"(d) DBI bus effect: max DQ-low/byte DDR5={worst5} (<=4?), "
f"avg DQ-low DDR3={br3.avg_dq_low():.2f} -> {'PASS' if worst5<=4 else 'FAIL'}")
# (e) generation swap over the SAME data, both correct
okg = True
for gen in ("DDR3", "DDR4", "DDR5"):
b = MemoryBridge(256, enc, dec, generation=gen)
for v in range(256):
b.write(v, v)
okg = okg and all(b.read(v) == v for v in range(256))
print(f"(e) generation swap DDR3/4/5 all round-trip -> {'PASS' if okg else 'FAIL'}")
if ne == te and nd == td:
torch.save({"encode": enc.state_dict(), "decode": dec.state_dict(),
"meta": {"unit": "DBI", "encode_verified": f"{ne}/{te}",
"decode_verified": f"{nd}/{td}"}}, "DBI.pt")
print(" saved verified unit -> DBI.pt")