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