#[inline(always)]
fn parity(x: u8) -> u8 {
let x = x ^ (x >> 4);
let x = x ^ (x >> 2);
(x ^ (x >> 1)) & 1
}
pub fn conv_encode(bits: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(bits.len() * 2);
let mut sr: u8 = 0;
for &b in bits {
let window = ((b & 1) << 4) | (sr & 0x0F);
out.push(parity(window & 0b10101)); out.push(parity(window & 0b10011)); sr = (sr >> 1) | ((b & 1) << 3);
}
out
}
const NUM_STATES: usize = 16;
#[inline]
fn branch_bits(s: u8, b: u8) -> (u8, u8) {
let window = ((b & 1) << 4) | (s & 0x0F);
(parity(window & 0b10101), parity(window & 0b10011))
}
#[inline]
fn next_state(s: u8, b: u8) -> u8 {
(s >> 1) | ((b & 1) << 3)
}
pub fn viterbi_decode(soft: &[f32]) -> Vec<u8> {
let n_syms = soft.len() / 2;
if n_syms == 0 {
return Vec::new();
}
let inf = f32::MAX / 2.0;
let mut pm = [inf; NUM_STATES];
pm[0] = 0.0;
let mut prev_state_table: Vec<[u8; NUM_STATES]> = vec![[0u8; NUM_STATES]; n_syms];
for t in 0..n_syms {
let s0 = soft[t * 2];
let s1 = soft[t * 2 + 1];
let mut new_pm = [inf; NUM_STATES];
for prev in 0..NUM_STATES {
if pm[prev] >= inf { continue; }
for &bit in &[0u8, 1u8] {
let (c0, c1) = branch_bits(prev as u8, bit);
let exp0 = if c0 == 0 { 1.0f32 } else { -1.0f32 };
let exp1 = if c1 == 0 { 1.0f32 } else { -1.0f32 };
let bm = (s0 - exp0) * (s0 - exp0) + (s1 - exp1) * (s1 - exp1);
let ns = next_state(prev as u8, bit) as usize;
let cand = pm[prev] + bm;
if cand < new_pm[ns] {
new_pm[ns] = cand;
prev_state_table[t][ns] = prev as u8;
}
}
}
pm = new_pm;
}
let mut state = pm
.iter()
.enumerate()
.min_by(|a, b| a.1.partial_cmp(b.1).unwrap_or(std::cmp::Ordering::Equal))
.map(|(i, _)| i)
.unwrap_or(0);
let mut bits_out = vec![0u8; n_syms];
for t in (0..n_syms).rev() {
let prev = prev_state_table[t][state] as usize;
let b = (state >> 3) as u8 & 1;
bits_out[t] = b;
state = prev;
}
bits_out
}
pub fn viterbi_decode_hard(bits: &[u8]) -> Vec<u8> {
let soft: Vec<f32> = bits.iter().map(|&b| if b == 0 { 1.0f32 } else { -1.0f32 }).collect();
viterbi_decode(&soft)
}