Quazim0t0's picture
Import from Quazim0t0/neural-ddr; repoint refs to NeuralVerified
141e646 verified
Raw History Blame Contribute Delete
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())