neural-ddr / step2to6_train_verify.py
Quazim0t0's picture
Import from Quazim0t0/neural-ddr; repoint refs to NeuralVerified
141e646 verified
Raw
History Blame Contribute Delete
3.48 kB
"""
Steps 2-6: train + exhaustively verify ADDR_MAP, CMD_DECODE, WR_CRC, ODECC,
then compose them into a DDR5 bridge that corrects real bit-flips over host RAM.
"""
import random
import torch
from neural_ddr.common import mlp, train, verify
from neural_ddr import addr_map, cmd_decode, crc, ecc
from neural_ddr.ddr5_bridge import DDR5Bridge
torch.manual_seed(0)
random.seed(0)
R = {}
print("=" * 62)
print("STEPS 2-6 -- neural DDR5 units + bridge")
print("=" * 62)
# ---- Step 2: ADDR_MAP ----
print("Step 2 ADDR_MAP (12 -> 12):")
Xa, Ya = addr_map.domain()
net_addr = train(mlp(12, 12, 256, 2), Xa, Ya, steps=8000, tag="addr")
R["ADDR_MAP"] = verify(net_addr, Xa, Ya)
# ---- Step 3: CMD_DECODE ----
print("Step 3 CMD_DECODE (5 -> 4):")
Xc, Yc = cmd_decode.domain()
net_cmd = train(mlp(5, 4, 64, 2), Xc, Yc, steps=6000, tag="cmd")
R["CMD_DECODE"] = verify(net_cmd, Xc, Yc)
# ---- Step 4: WR_CRC (bit-slice) ----
print("Step 4 WR_CRC bit-slice (9 -> 8):")
Xr, Yr = crc.domain()
net_crc = train(mlp(9, 8, 128, 2), Xr, Yr, steps=6000, tag="crc")
R["WR_CRC"] = verify(net_crc, Xr, Yr)
# ---- Step 5: ODECC (encode + decode/correct) ----
print("Step 5 ODECC encode (8 -> 4):")
Xe, Ye = ecc.enc_domain()
net_enc = train(mlp(8, 4, 256, 3), Xe, Ye, steps=15000, tag="ecc-enc")
R["ODECC_ENC"] = verify(net_enc, Xe, Ye)
print("Step 5 ODECC decode/correct (12 -> 8):")
Xd, Yd = ecc.dec_domain()
net_dec = train(mlp(12, 8, 512, 3), Xd, Yd, steps=20000, tag="ecc-dec")
R["ODECC_DEC"] = verify(net_dec, Xd, Yd)
print("-" * 62)
for name, (ok, tot) in R.items():
print(f" {name:12s} verified {ok}/{tot} -> {'PASS' if ok == tot else 'FAIL'}")
# ---- behavioural checks ----
print("-" * 62)
# CMD_DECODE full table
cmd_ok = all(cmd_decode.run(net_cmd, v) == cmd_decode.golden_cmd(v) for v in range(32))
print(f"CMD_DECODE truth table (32/32): {'PASS' if cmd_ok else 'FAIL'} "
f"e.g. 0b00101 -> {cmd_decode.NAME[cmd_decode.run(net_cmd, 0b00101)]}")
# ADDR_MAP sample
f = addr_map.run(net_addr, 0xABC)
print(f"ADDR_MAP 0xABC -> col={f['col']} bank={f['bank']} bg={f['bg']} row={f['row']}")
# CRC burst vs golden
bad = 0
for _ in range(200):
msg = [random.randint(0, 255) for _ in range(8)]
bad += (crc.crc_burst(net_crc, msg) != crc.golden_burst(msg))
print(f"WR_CRC 8-byte bursts (200): {'PASS' if bad == 0 else f'FAIL({bad})'} "
f"(neural ripple == golden CRC)")
# ---- Step 6: DDR5 bridge with real bit-flip correction ----
N = 2048
br = DDR5Bridge(N, net_enc, net_dec)
for a in range(N):
br.write(a, (a * 89 + 7) & 0xFF)
# inject one random single-bit fault per location, then read back
corrected = 0
for a in range(N):
br.inject_fault(a, random.randint(0, 11))
if br.read(a) == ((a * 89 + 7) & 0xFF):
corrected += 1
print(f"DDR5Bridge SEC: {corrected}/{N} bytes correct after a random bit-flip each "
f"-> {'PASS' if corrected == N else 'FAIL'}")
allpass = all(ok == tot for ok, tot in R.values()) and cmd_ok and bad == 0 and corrected == N
print("=" * 62)
print(f"OVERALL: {'ALL PASS' if allpass else 'some checks failed'}")
if all(ok == tot for ok, tot in R.values()):
torch.save({
"addr_map": net_addr.state_dict(), "cmd_decode": net_cmd.state_dict(),
"wr_crc": net_crc.state_dict(), "odecc_enc": net_enc.state_dict(),
"odecc_dec": net_dec.state_dict(),
"meta": {k: f"{v[0]}/{v[1]}" for k, v in R.items()},
}, "DDR5_units.pt")
print("saved verified units -> DDR5_units.pt")