Download step1_train_verify.py from Quazim0t0/neural-ddr: direct link, hf CLI and curl.
- Browser
- Download file 3.89 kB
-
https://huggingface.co/Quazim0t0/neural-ddr/resolve/main/step1_train_verify.py
- Command line
-
hf download hf://Quazim0t0/neural-ddr/step1_train_verify.py
-
curl -L -o step1_train_verify.py https://huggingface.co/Quazim0t0/neural-ddr/resolve/main/step1_train_verify.py
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") | |