""" 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")