| """ |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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'}") |
|
|
| |
| print("-" * 62) |
| |
| 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)]}") |
|
|
| |
| f = addr_map.run(net_addr, 0xABC) |
| print(f"ADDR_MAP 0xABC -> col={f['col']} bank={f['bank']} bg={f['bg']} row={f['row']}") |
|
|
| |
| 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)") |
|
|
| |
| N = 2048 |
| br = DDR5Bridge(N, net_enc, net_dec) |
| for a in range(N): |
| br.write(a, (a * 89 + 7) & 0xFF) |
| |
| 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") |
|
|