use std::collections::HashMap;
use ndarray::{Array2, Array3, ArrayView1, ArrayView2, Axis, Zip};
use num_complex::Complex64;
use crate::{Laplacian, error::GspError, kernel::ChebyKernel};
use super::shared::{accumulate, affine_sparse, sparse_dense_mul_into};
pub struct ChebyConvolver {
l: Laplacian,
n_vertices: usize,
rec_cache: HashMap<(u64, u64), Laplacian>,
}
impl ChebyConvolver {
pub fn new(l: Laplacian) -> Result<Self, GspError> {
if l.rows() != l.cols() {
return Err(GspError::Dimensions(format!(
"Laplacian must be square, got {}x{}",
l.rows(),
l.cols()
)));
}
Ok(Self {
n_vertices: l.rows(),
l,
rec_cache: HashMap::new(),
})
}
pub fn convolve(
&mut self,
b: ArrayView2<f64>,
kernel: &ChebyKernel,
) -> Result<Array3<f64>, GspError> {
self.check_signal_2d(b)?;
kernel.validate()?;
self.convolve_real(b, kernel)
}
pub fn convolve_1d(
&mut self,
b: ArrayView1<f64>,
kernel: &ChebyKernel,
) -> Result<Array2<f64>, GspError> {
self.check_signal_1d(b)?;
let b2 = b.insert_axis(Axis(1));
let w = self.convolve(b2, kernel)?;
Ok(w.index_axis_move(Axis(1), 0))
}
pub fn convolve_complex(
&mut self,
b: ArrayView2<Complex64>,
kernel: &ChebyKernel,
) -> Result<Array3<Complex64>, GspError> {
self.check_signal_2d_complex(b)?;
let wr = self.convolve(b.mapv(|v| v.re).view(), kernel)?;
let wi = self.convolve(b.mapv(|v| v.im).view(), kernel)?;
let mut out = Array3::<Complex64>::zeros(wr.raw_dim());
Zip::from(&mut out)
.and(&wr)
.and(&wi)
.for_each(|o, &r, &i| *o = Complex64::new(r, i));
Ok(out)
}
pub fn convolve_complex_1d(
&mut self,
b: ArrayView1<Complex64>,
kernel: &ChebyKernel,
) -> Result<Array2<Complex64>, GspError> {
let b2 = b.insert_axis(Axis(1));
let w = self.convolve_complex(b2, kernel)?;
Ok(w.index_axis_move(Axis(1), 0))
}
pub fn convolve_multi(
&mut self,
b: ArrayView2<f64>,
kernels: &[ChebyKernel],
) -> Result<Vec<Array3<f64>>, GspError> {
self.check_signal_2d(b)?;
let mut out = Vec::with_capacity(kernels.len());
for k in kernels {
out.push(self.convolve(b, k)?);
}
Ok(out)
}
fn convolve_real(
&mut self,
b: ArrayView2<f64>,
kernel: &ChebyKernel,
) -> Result<Array3<f64>, GspError> {
let n_order = kernel.coefficients.nrows();
let n_dim = kernel.coefficients.ncols();
let n_signals = b.ncols();
if n_order == 0 || n_dim == 0 {
return Ok(Array3::zeros((self.n_vertices, n_signals, n_dim)));
}
let mut w = Array3::<f64>::zeros((self.n_vertices, n_signals, n_dim));
let mut tkm2 = b.to_owned();
accumulate(&mut w, &tkm2, kernel.coefficients.row(0));
if n_order == 1 {
return Ok(w);
}
let m = self.get_recurrence_matrix(kernel.spectrum_bound, kernel.min_lambda)?;
let mut tkm1 = Array2::<f64>::zeros(tkm2.raw_dim());
sparse_dense_mul_into(m, &tkm2, &mut tkm1);
accumulate(&mut w, &tkm1, kernel.coefficients.row(1));
let mut scratch = Array2::<f64>::zeros(tkm2.raw_dim());
for k in 2..n_order {
sparse_dense_mul_into(m, &tkm1, &mut scratch);
Zip::from(&mut scratch)
.and(&tkm2)
.for_each(|value, &prev| *value = 2.0 * *value - prev);
std::mem::swap(&mut tkm2, &mut tkm1);
std::mem::swap(&mut tkm1, &mut scratch);
accumulate(&mut w, &tkm1, kernel.coefficients.row(k));
}
Ok(w)
}
fn get_recurrence_matrix(
&mut self,
spectrum_bound: f64,
min_lambda: f64,
) -> Result<&Laplacian, GspError> {
if spectrum_bound <= min_lambda {
return Err(GspError::InvalidKernel(
"spectrum_bound must be greater than min_lambda".to_string(),
));
}
let key = (spectrum_bound.to_bits(), min_lambda.to_bits());
if !self.rec_cache.contains_key(&key) {
let range = spectrum_bound - min_lambda;
let alpha = 2.0 / range;
let beta = -(spectrum_bound + min_lambda) / range;
let m = affine_sparse(&self.l, alpha, beta);
self.rec_cache.insert(key, m);
}
self.rec_cache
.get(&key)
.ok_or_else(|| GspError::Factorization("failed to cache recurrence matrix".to_string()))
}
fn check_signal_1d(&self, b: ArrayView1<f64>) -> Result<(), GspError> {
if b.len() != self.n_vertices {
return Err(GspError::Dimensions(format!(
"signal length {} does not match graph size {}",
b.len(),
self.n_vertices
)));
}
Ok(())
}
fn check_signal_2d(&self, b: ArrayView2<f64>) -> Result<(), GspError> {
if b.nrows() != self.n_vertices {
return Err(GspError::Dimensions(format!(
"signal rows {} does not match graph size {}",
b.nrows(),
self.n_vertices
)));
}
Ok(())
}
fn check_signal_2d_complex(&self, b: ArrayView2<Complex64>) -> Result<(), GspError> {
if b.nrows() != self.n_vertices {
return Err(GspError::Dimensions(format!(
"signal rows {} does not match graph size {}",
b.nrows(),
self.n_vertices
)));
}
Ok(())
}
}