#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LdpcCode {
N512R12,
N576R23,
N512R34,
}
impl LdpcCode {
pub fn n(self) -> usize {
match self {
LdpcCode::N512R12 => 512,
LdpcCode::N576R23 => 576,
LdpcCode::N512R34 => 512,
}
}
pub fn k(self) -> usize {
match self {
LdpcCode::N512R12 => 256,
LdpcCode::N576R23 => 384,
LdpcCode::N512R34 => 384,
}
}
pub fn m(self) -> usize {
self.n() - self.k()
}
fn col_weight(self) -> usize {
3
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum DecodeRule {
SumProduct,
MinSum,
ScaledMinSum(f32),
}
#[derive(Debug, Clone)]
pub struct Ldpc {
code: LdpcCode,
n: usize,
k: usize,
m: usize,
msg_col_rows: Vec<Vec<usize>>,
check_bits: Vec<Vec<usize>>,
bit_checks: Vec<Vec<usize>>,
bit_check_edge_idx: Vec<Vec<usize>>,
}
impl Ldpc {
pub fn new(code: LdpcCode) -> Self {
let n = code.n();
let k = code.k();
let m = code.m();
assert!(m >= 1 && k >= 1 && n == k + m);
let cw = code.col_weight();
let mut msg_col_rows: Vec<Vec<usize>> = Vec::with_capacity(k);
let mut row_load = vec![0usize; m];
let mut used_pairs: std::collections::HashSet<(usize, usize)> =
std::collections::HashSet::new();
let mut state: u64 = code_seed(code);
let mut next = || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for _col in 0..k {
let mut rows: Vec<usize> = Vec::with_capacity(cw);
while rows.len() < cw {
let offset = (next() % m as u64) as usize;
let mut best: Option<usize> = None;
let mut best_load = usize::MAX;
for step in 0..m {
let r = (offset + step) % m;
if rows.contains(&r) {
continue;
}
let makes_cycle = rows
.iter()
.any(|&q| used_pairs.contains(&ordered_pair(q, r)));
if makes_cycle {
continue;
}
if row_load[r] < best_load {
best_load = row_load[r];
best = Some(r);
}
}
match best {
Some(r) => rows.push(r),
None => {
let r = (0..m)
.map(|s| (offset + s) % m)
.find(|r| !rows.contains(r))
.expect("m > col_weight guarantees a free row");
rows.push(r);
}
}
}
for i in 0..rows.len() {
row_load[rows[i]] += 1;
for j in (i + 1)..rows.len() {
used_pairs.insert(ordered_pair(rows[i], rows[j]));
}
}
rows.sort_unstable();
msg_col_rows.push(rows);
}
let mut check_bits: Vec<Vec<usize>> = vec![Vec::new(); m];
let mut bit_checks: Vec<Vec<usize>> = vec![Vec::new(); n];
for (col, rows) in msg_col_rows.iter().enumerate() {
for &r in rows {
check_bits[r].push(col);
bit_checks[col].push(r);
}
}
#[allow(clippy::needless_range_loop)]
for i in 0..m {
let pcol = k + i;
check_bits[i].push(pcol);
bit_checks[pcol].push(i);
if i > 0 {
let prev = k + i - 1;
check_bits[i].push(prev);
bit_checks[prev].push(i);
}
}
let bit_check_edge_idx: Vec<Vec<usize>> = bit_checks
.iter()
.enumerate()
.map(|(b, checks)| {
checks
.iter()
.map(|&c| {
check_bits[c]
.iter()
.position(|&x| x == b)
.expect("bit is incident to its check")
})
.collect()
})
.collect();
Self {
code,
n,
k,
m,
msg_col_rows,
check_bits,
bit_checks,
bit_check_edge_idx,
}
}
pub fn code(&self) -> LdpcCode {
self.code
}
pub fn n(&self) -> usize {
self.n
}
pub fn k(&self) -> usize {
self.k
}
pub fn m(&self) -> usize {
self.m
}
pub fn encode(&self, message: &[u8]) -> Vec<u8> {
assert_eq!(message.len(), self.k, "LDPC message must be exactly K bits");
let mut cw = vec![0u8; self.n];
cw[..self.k].copy_from_slice(message);
let mut s = vec![0u8; self.m];
for (col, rows) in self.msg_col_rows.iter().enumerate() {
let bit = message[col] & 1;
if bit != 0 {
for &r in rows {
s[r] ^= 1;
}
}
}
let mut prev = 0u8;
for i in 0..self.m {
let p = s[i] ^ prev;
cw[self.k + i] = p;
prev = p;
}
cw
}
pub fn syndrome_weight(&self, hard: &[u8]) -> usize {
let mut unsat = 0;
for bits in &self.check_bits {
let mut x = 0u8;
for &b in bits {
x ^= hard[b] & 1;
}
if x != 0 {
unsat += 1;
}
}
unsat
}
pub fn decode_soft(&self, llr: &[f32], max_iter: usize) -> (Vec<u8>, usize) {
self.decode_soft_with(llr, max_iter, DecodeRule::SumProduct)
}
pub fn decode_soft_with(
&self,
llr: &[f32],
max_iter: usize,
rule: DecodeRule,
) -> (Vec<u8>, usize) {
assert_eq!(llr.len(), self.n, "LDPC LLR slice must be N long");
let mut hard = vec![0u8; self.n];
for (h, &l) in hard.iter_mut().zip(llr) {
*h = u8::from(l <= 0.0);
}
let init_unsat = self.syndrome_weight(&hard);
if init_unsat == 0 {
return (hard[..self.k].to_vec(), 0);
}
let n_edges: usize = self.check_bits.iter().map(Vec::len).sum();
let mut check_start = vec![0usize; self.m + 1];
for (c, bits) in self.check_bits.iter().enumerate() {
check_start[c + 1] = check_start[c] + bits.len();
}
let mut msg = vec![0.0f32; n_edges];
for (c, bits) in self.check_bits.iter().enumerate() {
let base = check_start[c];
for (i, &b) in bits.iter().enumerate() {
msg[base + i] = llr[b];
}
}
let mut ext = vec![0.0f32; n_edges];
let mut min_unsat = init_unsat;
let mut best = hard.clone();
let max_deg = self.check_bits.iter().map(Vec::len).max().unwrap_or(0);
let mut tanh_half = vec![0.0f32; max_deg];
for _iter in 0..max_iter {
for (c, bits) in self.check_bits.iter().enumerate() {
let deg = bits.len();
let base = check_start[c];
let msg_c = &msg[base..base + deg];
let ext_c = &mut ext[base..base + deg];
match rule {
DecodeRule::SumProduct => {
for j in 0..deg {
tanh_half[j] = fast_tanh(msg_c[j] / 2.0);
}
#[allow(clippy::needless_range_loop)]
for i1 in 0..deg {
let mut prod = 1.0f32;
for i2 in 0..deg {
if i2 != i1 {
prod *= tanh_half[i2];
}
}
ext_c[i1] = 2.0 * fast_atanh(prod.clamp(-1.0, 1.0));
}
}
DecodeRule::MinSum | DecodeRule::ScaledMinSum(_) => {
let scale = match rule {
DecodeRule::ScaledMinSum(a) => a,
_ => 1.0,
};
let mut min1 = f32::INFINITY; let mut min2 = f32::INFINITY; let mut argmin = 0usize; let mut sign_parity = 1.0f32; for (j, &v) in msg_c.iter().enumerate() {
if v < 0.0 {
sign_parity = -sign_parity;
}
let a = v.abs();
if a < min1 {
min2 = min1;
min1 = a;
argmin = j;
} else if a < min2 {
min2 = a;
}
}
for i1 in 0..deg {
let s_other = if msg_c[i1] < 0.0 {
-sign_parity
} else {
sign_parity
};
let mag = if i1 == argmin { min2 } else { min1 };
ext_c[i1] = scale * s_other * mag;
}
}
}
}
for (bit, checks) in self.bit_checks.iter().enumerate() {
let edge_idx = &self.bit_check_edge_idx[bit];
let mut l = llr[bit];
for (&c, &idx) in checks.iter().zip(edge_idx) {
l += ext[check_start[c] + idx];
}
hard[bit] = u8::from(l <= 0.0);
}
let unsat = self.syndrome_weight(&hard);
if unsat < min_unsat {
min_unsat = unsat;
best.copy_from_slice(&hard);
if unsat == 0 {
break;
}
}
for (bit, checks) in self.bit_checks.iter().enumerate() {
let edge_idx = &self.bit_check_edge_idx[bit];
let total: f32 = llr[bit]
+ checks
.iter()
.zip(edge_idx)
.map(|(&c, &idx)| ext[check_start[c] + idx])
.sum::<f32>();
for (&c, &idx) in checks.iter().zip(edge_idx) {
let e = check_start[c] + idx;
msg[e] = total - ext[e];
}
}
}
(best[..self.k].to_vec(), min_unsat)
}
}
#[inline]
fn ordered_pair(a: usize, b: usize) -> (usize, usize) {
if a <= b { (a, b) } else { (b, a) }
}
#[inline]
fn code_seed(code: LdpcCode) -> u64 {
match code {
LdpcCode::N512R12 => 0x4C44_5043_3531_3200,
LdpcCode::N576R23 => 0x4C44_5043_3531_3201,
LdpcCode::N512R34 => 0x4C44_5043_3531_3202,
}
}
#[inline]
fn fast_tanh(x: f32) -> f32 {
if x < -4.97 {
return -1.0;
}
if x > 4.97 {
return 1.0;
}
let x2 = x * x;
let a = x * (945.0 + x2 * (105.0 + x2));
let b = 945.0 + x2 * (420.0 + x2 * 15.0);
a / b
}
#[inline]
fn fast_atanh(x: f32) -> f32 {
let x2 = x * x;
let a = x * (945.0 + x2 * (-735.0 + x2 * 64.0));
let b = 945.0 + x2 * (-1050.0 + x2 * 225.0);
a / b
}