gspx 0.1.2

Sparse graph signal processing and spectral graph wavelets in Rust
Documentation
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 {
    /// Creates a Chebyshev convolver for a fixed Laplacian.
    ///
    /// # Errors
    /// Returns an error when `l` is not square.
    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(),
        })
    }

    /// Convolves a real-valued signal matrix with a Chebyshev kernel.
    ///
    /// # Errors
    /// Returns an error when dimensions/kernel parameters are invalid.
    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)
    }

    /// 1D wrapper around [`ChebyConvolver::convolve`].
    ///
    /// # Errors
    /// Returns an error when dimensions/kernel parameters are invalid.
    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))
    }

    /// Convolves complex-valued signals by splitting real/imaginary parts.
    ///
    /// # Errors
    /// Returns an error when dimensions/kernel parameters are invalid.
    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)
    }

    /// 1D wrapper around [`ChebyConvolver::convolve_complex`].
    ///
    /// # Errors
    /// Returns an error when dimensions/kernel parameters are invalid.
    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))
    }

    /// Applies multiple Chebyshev kernels to the same real-valued signal matrix.
    ///
    /// # Errors
    /// Returns an error when dimensions/kernel parameters are invalid.
    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(())
    }
}