#[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)]
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>>,
}
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);
}
}
Self {
code,
n,
k,
m,
msg_col_rows,
check_bits,
bit_checks,
}
}
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) {
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 mut msg: Vec<Vec<f32>> = self
.check_bits
.iter()
.map(|bits| bits.iter().map(|&b| llr[b]).collect())
.collect();
let mut ext: Vec<Vec<f32>> = self.check_bits.iter().map(|b| vec![0.0; b.len()]).collect();
let mut min_unsat = init_unsat;
let mut best = hard.clone();
for _iter in 0..max_iter {
for (c, bits) in self.check_bits.iter().enumerate() {
let deg = bits.len();
#[allow(clippy::needless_range_loop)]
for i1 in 0..deg {
let mut prod = 1.0f32;
for (i2, _) in bits.iter().enumerate() {
if i2 != i1 {
prod *= fast_tanh(msg[c][i2] / 2.0);
}
}
ext[c][i1] = 2.0 * fast_atanh(prod.clamp(-1.0, 1.0));
}
}
for (bit, checks) in self.bit_checks.iter().enumerate() {
let mut l = llr[bit];
for &c in checks {
let idx = self.check_bits[c].iter().position(|&b| b == bit).unwrap();
l += ext[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 total: f32 = llr[bit]
+ checks
.iter()
.map(|&c| {
let idx = self.check_bits[c].iter().position(|&b| b == bit).unwrap();
ext[c][idx]
})
.sum::<f32>();
for &c in checks {
let idx = self.check_bits[c].iter().position(|&b| b == bit).unwrap();
msg[c][idx] = total - ext[c][idx];
}
}
}
(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
}