use crate::celt_tf_adjust::TfDirection;
const INV_SQRT2: f64 = core::f64::consts::FRAC_1_SQRT_2;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TfHadamardError {
ZeroBlocks,
BlocksNotPowerOfTwo {
nb_blocks: usize,
},
BlocksDoNotDivideLength {
len: usize,
nb_blocks: usize,
},
LevelsExceedBlocks {
levels: u8,
nb_blocks: usize,
},
}
impl core::fmt::Display for TfHadamardError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match *self {
TfHadamardError::ZeroBlocks => {
write!(
f,
"oxideav-opus: CELT §4.3.4.5 TF transform requires nb_blocks >= 1"
)
}
TfHadamardError::BlocksNotPowerOfTwo { nb_blocks } => write!(
f,
"oxideav-opus: CELT §4.3.4.5 TF transform requires a power-of-two block \
count, got nb_blocks = {nb_blocks}"
),
TfHadamardError::BlocksDoNotDivideLength { len, nb_blocks } => write!(
f,
"oxideav-opus: CELT §4.3.4.5 TF transform: vector length {len} is not a \
multiple of nb_blocks = {nb_blocks}"
),
TfHadamardError::LevelsExceedBlocks { levels, nb_blocks } => write!(
f,
"oxideav-opus: CELT §4.3.4.5 TF transform: {levels} Hadamard levels exceed \
log2(nb_blocks) for nb_blocks = {nb_blocks}"
),
}
}
}
impl std::error::Error for TfHadamardError {}
fn fwht_natural_inplace(x: &mut [f64]) {
let n = x.len();
debug_assert!(n.is_power_of_two());
let mut stride = 1;
while stride < n {
let mut base = 0;
while base < n {
for i in base..base + stride {
let a = x[i];
let b = x[i + stride];
x[i] = (a + b) * INV_SQRT2;
x[i + stride] = (a - b) * INV_SQRT2;
}
base += stride << 1;
}
stride <<= 1;
}
}
#[inline]
fn bit_reverse(mut v: usize, bits: u32) -> usize {
let mut r = 0;
for _ in 0..bits {
r = (r << 1) | (v & 1);
v >>= 1;
}
r
}
#[inline]
fn inverse_gray(mut v: usize) -> usize {
let mut mask = v >> 1;
while mask != 0 {
v ^= mask;
mask >>= 1;
}
v
}
fn sequency_permutation(bits: u32) -> Vec<usize> {
let n = 1usize << bits;
let mut rank = vec![0usize; n];
for (r, slot) in rank.iter_mut().enumerate() {
*slot = inverse_gray(bit_reverse(r, bits));
}
let mut perm = vec![0usize; n];
for (r, &s) in rank.iter().enumerate() {
perm[s] = r;
}
perm
}
pub fn apply_tf_hadamard(
x: &mut [f64],
nb_blocks: usize,
direction: TfDirection,
) -> Result<(), TfHadamardError> {
if nb_blocks == 0 {
return Err(TfHadamardError::ZeroBlocks);
}
if !nb_blocks.is_power_of_two() {
return Err(TfHadamardError::BlocksNotPowerOfTwo { nb_blocks });
}
let len = x.len();
if len % nb_blocks != 0 {
return Err(TfHadamardError::BlocksDoNotDivideLength { len, nb_blocks });
}
let levels = direction.levels();
if levels == 0 {
return Ok(()); }
let span = 1usize << levels;
if span > nb_blocks {
return Err(TfHadamardError::LevelsExceedBlocks { levels, nb_blocks });
}
let bins = len / nb_blocks; let sequency = matches!(direction, TfDirection::IncreaseTime(_));
let perm = if sequency {
Some(sequency_permutation(levels as u32))
} else {
None
};
let mut scratch = vec![0.0f64; span];
for m in 0..bins {
let base = m * nb_blocks;
let mut off = 0;
while off + span <= nb_blocks {
let group = &mut x[base + off..base + off + span];
scratch.copy_from_slice(group);
fwht_natural_inplace(&mut scratch);
if let Some(perm) = &perm {
for (dst, &src) in group.iter_mut().zip(perm.iter()) {
*dst = scratch[src];
}
} else {
group.copy_from_slice(&scratch);
}
off += span;
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
const TOL: f64 = 1e-12;
fn l2(x: &[f64]) -> f64 {
x.iter().map(|v| v * v).sum::<f64>().sqrt()
}
#[test]
fn unchanged_is_identity() {
let mut x = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let orig = x.clone();
apply_tf_hadamard(&mut x, 4, TfDirection::Unchanged).unwrap();
assert_eq!(x, orig);
}
#[test]
fn single_block_no_op_even_with_levels() {
let mut x = vec![1.0, 2.0, 3.0];
let r = apply_tf_hadamard(&mut x, 1, TfDirection::IncreaseTime(1));
assert_eq!(
r,
Err(TfHadamardError::LevelsExceedBlocks {
levels: 1,
nb_blocks: 1
})
);
}
#[test]
fn two_point_frequency_butterfly() {
let mut x = vec![3.0, 1.0];
apply_tf_hadamard(&mut x, 2, TfDirection::IncreaseFrequency(1)).unwrap();
assert!((x[0] - 4.0 * INV_SQRT2).abs() < TOL);
assert!((x[1] - 2.0 * INV_SQRT2).abs() < TOL);
}
#[test]
fn orthonormal_preserves_l2_norm() {
for dir in [
TfDirection::IncreaseFrequency(1),
TfDirection::IncreaseFrequency(2),
TfDirection::IncreaseTime(1),
TfDirection::IncreaseTime(2),
] {
let mut x: Vec<f64> = (0..16).map(|i| (i as f64 * 0.37).sin()).collect();
let before = l2(&x);
apply_tf_hadamard(&mut x, 4, dir).unwrap();
let after = l2(&x);
assert!(
(before - after).abs() < 1e-9,
"dir {dir:?}: {before} vs {after}"
);
}
}
#[test]
fn transform_is_self_inverse() {
for (dir, blocks) in [
(TfDirection::IncreaseFrequency(2), 4usize),
(TfDirection::IncreaseTime(2), 4usize),
(TfDirection::IncreaseFrequency(3), 8usize),
(TfDirection::IncreaseTime(3), 8usize),
] {
let orig: Vec<f64> = (0..blocks * 3).map(|i| (i as f64 * 0.13).cos()).collect();
let mut x = orig.clone();
apply_tf_hadamard(&mut x, blocks, dir).unwrap();
apply_tf_hadamard(&mut x, blocks, dir).unwrap();
for (a, b) in x.iter().zip(orig.iter()) {
assert!((a - b).abs() < 1e-9, "self-inverse {dir:?}: {a} vs {b}");
}
}
}
#[test]
fn frequency_and_time_differ_for_multilevel() {
let base: Vec<f64> = (0..4).map(|i| (i as f64 + 1.0) * 0.5).collect();
let mut freq = base.clone();
let mut time = base.clone();
apply_tf_hadamard(&mut freq, 4, TfDirection::IncreaseFrequency(2)).unwrap();
apply_tf_hadamard(&mut time, 4, TfDirection::IncreaseTime(2)).unwrap();
assert_ne!(
freq.iter().map(|v| format!("{v:.9}")).collect::<Vec<_>>(),
time.iter().map(|v| format!("{v:.9}")).collect::<Vec<_>>()
);
}
#[test]
fn level1_sequency_equals_natural() {
let base = vec![2.0, 5.0];
let mut freq = base.clone();
let mut time = base.clone();
apply_tf_hadamard(&mut freq, 2, TfDirection::IncreaseFrequency(1)).unwrap();
apply_tf_hadamard(&mut time, 2, TfDirection::IncreaseTime(1)).unwrap();
for (a, b) in freq.iter().zip(time.iter()) {
assert!((a - b).abs() < TOL);
}
}
#[test]
fn sequency_permutation_level2_is_correct() {
let perm = sequency_permutation(2);
assert_eq!(perm, vec![0, 2, 3, 1]);
}
#[test]
fn partial_levels_transform_subgroups_independently() {
let mut x = vec![1.0, 3.0, 10.0, 6.0]; apply_tf_hadamard(&mut x, 4, TfDirection::IncreaseFrequency(1)).unwrap();
assert!((x[0] - 4.0 * INV_SQRT2).abs() < TOL);
assert!((x[1] + 2.0 * INV_SQRT2).abs() < TOL);
assert!((x[2] - 16.0 * INV_SQRT2).abs() < TOL);
assert!((x[3] - 4.0 * INV_SQRT2).abs() < TOL);
}
#[test]
fn multi_bin_each_bin_independent() {
let mut x = vec![3.0, 1.0, 8.0, 2.0];
apply_tf_hadamard(&mut x, 2, TfDirection::IncreaseFrequency(1)).unwrap();
assert!((x[0] - 4.0 * INV_SQRT2).abs() < TOL);
assert!((x[1] - 2.0 * INV_SQRT2).abs() < TOL);
assert!((x[2] - 10.0 * INV_SQRT2).abs() < TOL);
assert!((x[3] - 6.0 * INV_SQRT2).abs() < TOL);
}
#[test]
fn errors_on_non_power_of_two_blocks() {
let mut x = vec![1.0; 6];
assert_eq!(
apply_tf_hadamard(&mut x, 3, TfDirection::IncreaseFrequency(1)),
Err(TfHadamardError::BlocksNotPowerOfTwo { nb_blocks: 3 })
);
}
#[test]
fn errors_on_zero_blocks() {
let mut x = vec![1.0; 4];
assert_eq!(
apply_tf_hadamard(&mut x, 0, TfDirection::IncreaseFrequency(1)),
Err(TfHadamardError::ZeroBlocks)
);
}
#[test]
fn errors_when_blocks_do_not_divide_length() {
let mut x = vec![1.0; 5];
assert_eq!(
apply_tf_hadamard(&mut x, 2, TfDirection::IncreaseFrequency(1)),
Err(TfHadamardError::BlocksDoNotDivideLength {
len: 5,
nb_blocks: 2
})
);
}
#[test]
fn fwht_natural_matches_hand_computed_order4() {
let mut x = vec![1.0, 0.0, 0.0, 0.0];
fwht_natural_inplace(&mut x);
for v in &x {
assert!((v - 0.5).abs() < TOL);
}
}
}