use crate::error::SpectralError;
use crate::transform::is_power_of_two_len;
fn require_pow2(a: &[f64]) -> Result<usize, SpectralError> {
if !is_power_of_two_len(a.len()) {
return Err(SpectralError::NonPowerOfTwoLength { len: a.len() });
}
Ok(a.len().trailing_zeros() as usize)
}
pub fn zeta_transform(a: &[f64]) -> Result<Vec<f64>, SpectralError> {
let bits = require_pow2(a)?;
let mut f = a.to_vec();
let n = f.len();
for i in 0..bits {
for mask in 0..n {
if mask & (1 << i) != 0 {
f[mask] += f[mask ^ (1 << i)];
}
}
}
Ok(f)
}
pub fn mobius_transform(a: &[f64]) -> Result<Vec<f64>, SpectralError> {
let bits = require_pow2(a)?;
let mut f = a.to_vec();
let n = f.len();
for i in 0..bits {
for mask in 0..n {
if mask & (1 << i) != 0 {
f[mask] -= f[mask ^ (1 << i)];
}
}
}
Ok(f)
}
pub fn or_convolution(a: &[f64], b: &[f64]) -> Result<Vec<f64>, SpectralError> {
if a.len() != b.len() {
return Err(SpectralError::ShapeMismatch(
"or_convolution operands".to_string(),
));
}
let za = zeta_transform(a)?;
let zb = zeta_transform(b)?;
let prod: Vec<f64> = za.iter().zip(&zb).map(|(x, y)| x * y).collect();
mobius_transform(&prod)
}
fn zeta_in_place(f: &mut [f64], bits: usize) {
let n = f.len();
for i in 0..bits {
for mask in 0..n {
if mask & (1 << i) != 0 {
f[mask] += f[mask ^ (1 << i)];
}
}
}
}
fn mobius_in_place(f: &mut [f64], bits: usize) {
let n = f.len();
for i in 0..bits {
for mask in 0..n {
if mask & (1 << i) != 0 {
f[mask] -= f[mask ^ (1 << i)];
}
}
}
}
pub fn subset_convolution(a: &[f64], b: &[f64]) -> Result<Vec<f64>, SpectralError> {
if a.len() != b.len() {
return Err(SpectralError::ShapeMismatch(
"subset_convolution operands".to_string(),
));
}
let bits = require_pow2(a)?;
let n = a.len();
let mut fa = vec![vec![0.0; n]; bits + 1];
let mut fb = vec![vec![0.0; n]; bits + 1];
for (mask, (&av, &bv)) in a.iter().zip(b).enumerate() {
let pc = (mask as u32).count_ones() as usize;
fa[pc][mask] = av;
fb[pc][mask] = bv;
}
for layer in 0..=bits {
zeta_in_place(&mut fa[layer], bits);
zeta_in_place(&mut fb[layer], bits);
}
let mut fh = vec![vec![0.0; n]; bits + 1];
for i in 0..=bits {
for j in 0..=i {
for mask in 0..n {
fh[i][mask] += fa[j][mask] * fb[i - j][mask];
}
}
}
let mut h = vec![0.0; n];
for (i, layer) in fh.iter_mut().enumerate() {
mobius_in_place(layer, bits);
for mask in 0..n {
if (mask as u32).count_ones() as usize == i {
h[mask] = layer[mask];
}
}
}
Ok(h)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn zeta_mobius_are_inverse() {
let a = vec![3.0, -1.0, 4.0, 1.0, 5.0, 9.0, 2.0, 6.0];
let recovered = mobius_transform(&zeta_transform(&a).unwrap()).unwrap();
for (x, y) in recovered.iter().zip(&a) {
assert!((x - y).abs() < 1e-9);
}
}
#[test]
fn zeta_is_subset_sum() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let z = zeta_transform(&a).unwrap();
assert_eq!(z[3], 10.0);
assert_eq!(z[1], 1.0 + 2.0);
}
#[test]
fn subset_convolution_matches_brute_force() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let b = vec![5.0, 6.0, 7.0, 8.0];
let got = subset_convolution(&a, &b).unwrap();
let n = a.len();
let mut want = vec![0.0; n];
for s in 0..n {
let mut sub = s;
loop {
want[s] += a[sub] * b[s ^ sub];
if sub == 0 {
break;
}
sub = (sub - 1) & s;
}
}
for (x, y) in got.iter().zip(&want) {
assert!((x - y).abs() < 1e-9, "{x} vs {y}");
}
}
}