File size: 2,897 Bytes
141e646
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
"""
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())