Download neural_ddr/ecc.py from NeuralVerified/neural-ddr: direct link, hf CLI and curl.
- Browser
- Download file 2.9 kB
-
https://huggingface.co/NeuralVerified/neural-ddr/resolve/main/neural_ddr/ecc.py
- Command line
-
hf download hf://NeuralVerified/neural-ddr/neural_ddr/ecc.py
-
curl -L -o ecc.py https://huggingface.co/NeuralVerified/neural-ddr/resolve/main/neural_ddr/ecc.py
2.9 kB
| """ | |
| Step 5: ODECC -- DDR5 on-die ECC (the marquee DDR5 feature). | |
| DDR5 adds on-die single-error-correction. We implement a concrete SEC Hamming | |
| code -- Hamming(12,8): 8 data bits protected by 4 parity bits, correcting any | |
| single-bit flip in the 12-bit codeword. Real DDR5 protects wider words; the | |
| logic is identical, just more bits. | |
| Two verified units: | |
| ENC : 8 data bits -> 4 parity bits (domain 256, N/N) | |
| DEC : 12-bit codeword -> 8 corrected data bits (all no-error + every | |
| single-bit flip, 256*13 = 3328 inputs, N/N) | |
| The DEC unit genuinely CORRECTS a flipped bit -- that is what lets the bridge | |
| present ECC-protected memory over ordinary host RAM. | |
| """ | |
| from __future__ import annotations | |
| import torch | |
| from .common import bits_of, int_of, pm | |
| DATA_POS = [3, 5, 6, 7, 9, 10, 11, 12] # 1-indexed Hamming positions | |
| PAR_POS = [1, 2, 4, 8] | |
| ENC_IN, ENC_OUT = 8, 4 | |
| DEC_IN, DEC_OUT = 12, 8 | |
| def encode(data8: int) -> list[int]: | |
| d = bits_of(data8, 8) | |
| cw = {p: 0 for p in range(1, 13)} | |
| for i, pos in enumerate(DATA_POS): | |
| cw[pos] = d[i] | |
| par = [] | |
| for j in PAR_POS: | |
| p = 0 | |
| for pos in range(1, 13): | |
| if pos in PAR_POS: | |
| continue | |
| if pos & j: | |
| p ^= cw[pos] | |
| par.append(p) | |
| return par # for positions 1,2,4,8 | |
| def codeword_int(data8: int) -> int: | |
| d = bits_of(data8, 8) | |
| par = encode(data8) | |
| cw = [0] * 12 # index k = position k+1 | |
| for i, pos in enumerate(DATA_POS): | |
| cw[pos - 1] = d[i] | |
| for i, j in enumerate(PAR_POS): | |
| cw[j - 1] = par[i] | |
| return int_of(cw) | |
| def decode(cw_int: int) -> int: | |
| cw = bits_of(cw_int, 12) | |
| syn = 0 | |
| for j in PAR_POS: | |
| s = 0 | |
| for pos in range(1, 13): | |
| if pos & j: | |
| s ^= cw[pos - 1] | |
| if s: | |
| syn |= j | |
| if 1 <= syn <= 12: | |
| cw[syn - 1] ^= 1 # correct the flipped bit | |
| return int_of([cw[pos - 1] for pos in DATA_POS]) | |
| def enc_domain(): | |
| X = torch.stack([pm(bits_of(d, 8)) for d in range(256)]) | |
| Y = torch.tensor([[float(b) for b in encode(d)] for d in range(256)]) | |
| return X, Y | |
| def dec_domain(): | |
| xs, ys = [], [] | |
| for d in range(256): | |
| base = codeword_int(d) | |
| for e in range(13): # 0 = no error, 1..12 flip pos e | |
| v = base if e == 0 else base ^ (1 << (e - 1)) | |
| xs.append(pm(bits_of(v, 12))) | |
| ys.append([float(b) for b in bits_of(d, 8)]) | |
| return torch.stack(xs), torch.tensor(ys) | |
| def enc_run(net, d: int) -> list[int]: | |
| return (net(pm(bits_of(d, 8)).unsqueeze(0))[0] > 0).int().tolist() | |
| def dec_run(net, cw_int: int) -> int: | |
| return int_of((net(pm(bits_of(cw_int, 12)).unsqueeze(0))[0] > 0).int().tolist()) | |