use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use num_complex::Complex64;
use crate::{
Laplacian, conv::DynamicConvolver, error::GspError, functions::gaussian_wavelet,
kernel::impulse,
};
use super::{
numerics::{
arrays_close, gradient_complex_axis, gradient_real_axis, peak_local_max,
real_times_complex, scalar_close, stddev,
},
peaks::{ClusterCloud, ModeTable, NetworkAnalysisResult, PeakCloud, Peaks},
};
pub struct Sgma {
laplacian: Laplacian,
pub scales: Array1<f64>,
pub freqs: Array1<f64>,
pub order: usize,
pub w0: f64,
pub ts: Array1<f64>,
pub wavlen: Array1<f64>,
spatial_convolver: DynamicConvolver,
temporal_matrix_cache: Option<Array2<Complex64>>,
temporal_grid_cache: Option<Array1<f64>>,
temporal_target_cache: Option<f64>,
}
impl Sgma {
pub fn new(
l: Laplacian,
scales: Array1<f64>,
freqs: Array1<f64>,
order: usize,
w0: f64,
) -> Result<Self, GspError> {
if scales.is_empty() || freqs.is_empty() {
return Err(GspError::Dimensions(
"scales and freqs must be non-empty".to_string(),
));
}
if scales.iter().any(|s| *s <= 0.0 || !s.is_finite()) {
return Err(GspError::InvalidScales(
"all scales must be finite and > 0".to_string(),
));
}
if freqs.iter().any(|f| *f <= 0.0 || !f.is_finite()) {
return Err(GspError::InvalidScales(
"all frequencies must be finite and > 0".to_string(),
));
}
let ts = freqs.mapv(|f| w0 / (2.0 * std::f64::consts::PI * f));
let wavlen = scales.mapv(|s| s.sqrt());
let spatial_convolver = DynamicConvolver::new(l.clone())?;
Ok(Self {
laplacian: l,
scales,
freqs,
order,
w0,
ts,
wavlen,
spatial_convolver,
temporal_matrix_cache: None,
temporal_grid_cache: None,
temporal_target_cache: None,
})
}
pub fn spectrum(
&mut self,
v: ArrayView2<f64>,
t: ArrayView1<f64>,
bus: usize,
time: f64,
) -> Result<Array2<f64>, GspError> {
Ok(self
.spectrum_core(v, t, bus, time, None)?
.mapv(|z| z.norm()))
}
pub fn spectrum_with_precomputed_temporal(
&mut self,
v: ArrayView2<f64>,
t: ArrayView1<f64>,
bus: usize,
time: f64,
vb: ArrayView2<Complex64>,
) -> Result<Array2<f64>, GspError> {
Ok(self
.spectrum_core(v, t, bus, time, Some(vb))?
.mapv(|z| z.norm()))
}
pub fn spectrum_complex(
&mut self,
v: ArrayView2<f64>,
t: ArrayView1<f64>,
bus: usize,
time: f64,
) -> Result<Array2<Complex64>, GspError> {
self.spectrum_core(v, t, bus, time, None)
}
pub fn spectrum_complex_with_precomputed_temporal(
&mut self,
v: ArrayView2<f64>,
t: ArrayView1<f64>,
bus: usize,
time: f64,
vb: ArrayView2<Complex64>,
) -> Result<Array2<Complex64>, GspError> {
self.spectrum_core(v, t, bus, time, Some(vb))
}
fn spectrum_core(
&mut self,
v: ArrayView2<f64>,
t: ArrayView1<f64>,
bus: usize,
time: f64,
vb: Option<ArrayView2<Complex64>>,
) -> Result<Array2<Complex64>, GspError> {
self.validate_bus(bus)?;
self.validate_signal_frame(v, t)?;
let n_nodes = self.laplacian.rows();
let vb_mat = if let Some(vb_in) = vb {
if vb_in.nrows() != n_nodes || vb_in.ncols() != self.ts.len() {
return Err(GspError::Dimensions(format!(
"precomputed VB shape {:?} does not match expected ({}, {})",
vb_in.raw_dim(),
n_nodes,
self.ts.len()
)));
}
vb_in.to_owned()
} else {
let bmat = self.build_temporal_matrix(t, time)?;
real_times_complex(v, bmat.view())
};
let impulse_signal = impulse(&self.laplacian, bus, 1)?;
let spatial_scales = self.scales.as_slice().ok_or_else(|| {
GspError::Dimensions("SGMA scales must be stored contiguously".to_string())
})?;
let spatial =
self.spatial_convolver
.bandpass(impulse_signal.view(), spatial_scales, self.order)?;
let mut a = Array2::<f64>::zeros((spatial.len(), n_nodes));
for (idx, r) in spatial.iter().enumerate() {
a.row_mut(idx).assign(&r.column(0).to_owned());
}
Ok(real_times_complex(a.view(), vb_mat.view()))
}
pub fn analyze(
&mut self,
v: ArrayView2<f64>,
t: ArrayView1<f64>,
bus: usize,
time: f64,
top_n: usize,
min_dist: usize,
) -> Result<Peaks, GspError> {
if top_n == 0 {
return Ok(Peaks::empty(false));
}
let s = self.spectrum(v, t, bus, time)?;
self.find_peaks(s.view(), top_n, min_dist, false)
}
pub fn analyze_many(
&mut self,
v: ArrayView2<f64>,
t: ArrayView1<f64>,
time: f64,
buses: Option<&[usize]>,
top_n: usize,
min_dist: usize,
) -> Result<NetworkAnalysisResult, GspError> {
self.validate_signal_frame(v, t)?;
if top_n == 0 {
return Ok(NetworkAnalysisResult::empty());
}
let bus_list = if let Some(b) = buses {
b.to_vec()
} else {
(0..v.nrows()).collect::<Vec<_>>()
};
let bmat = self.build_temporal_matrix(t, time)?;
let vb = real_times_complex(v, bmat.view());
let mut all_w = Vec::new();
let mut all_f = Vec::new();
let mut all_m = Vec::new();
let mut all_b = Vec::new();
for &bus in &bus_list {
self.validate_bus(bus)?;
let y = self.spectrum_with_precomputed_temporal(v, t, bus, time, vb.view())?;
let p = self.find_peaks(y.view(), top_n, min_dist, false)?;
for idx in 0..p.wavelength.len() {
all_w.push(p.wavelength[idx]);
all_f.push(p.frequency[idx]);
all_m.push(p.magnitude[idx]);
all_b.push(bus);
}
}
let peaks = PeakCloud {
wavelength: Array1::from(all_w),
frequency: Array1::from(all_f),
magnitude: Array1::from(all_m),
bus_id: Array1::from(all_b),
};
let clusters = self.compute_density_clusters(&peaks, top_n, min_dist);
Ok(NetworkAnalysisResult { peaks, clusters })
}
pub fn find_peaks(
&self,
spectrum: ArrayView2<f64>,
top_n: usize,
min_dist: usize,
return_indices: bool,
) -> Result<Peaks, GspError> {
if top_n == 0 {
return Ok(Peaks::empty(return_indices));
}
if spectrum.nrows() != self.scales.len() || spectrum.ncols() != self.freqs.len() {
return Err(GspError::Dimensions(format!(
"spectrum shape {:?} does not match expected ({}, {})",
spectrum.raw_dim(),
self.scales.len(),
self.freqs.len()
)));
}
let min_distance = min_dist.max(1);
let coords = peak_local_max(spectrum, min_distance, top_n);
if coords.is_empty() {
return Ok(Peaks::empty(return_indices));
}
let mut wavelengths = Vec::with_capacity(coords.len());
let mut freqs = Vec::with_capacity(coords.len());
let mut mags = Vec::with_capacity(coords.len());
let mut sidx = Vec::with_capacity(coords.len());
let mut fidx = Vec::with_capacity(coords.len());
for (si, fi) in coords {
wavelengths.push(self.wavlen[si]);
freqs.push(self.freqs[fi]);
mags.push(spectrum[[si, fi]]);
sidx.push(si);
fidx.push(fi);
}
Ok(Peaks {
wavelength: Array1::from(wavelengths),
frequency: Array1::from(freqs),
magnitude: Array1::from(mags),
scale_idx: return_indices.then(|| Array1::from(sidx)),
freq_idx: return_indices.then(|| Array1::from(fidx)),
})
}
pub fn find_modes(
&self,
spectrum: ArrayView2<Complex64>,
top_n: usize,
min_dist: usize,
) -> Result<ModeTable, GspError> {
if top_n == 0 {
return Ok(ModeTable::empty());
}
let mag = spectrum.mapv(|z| z.norm());
let peaks = self.find_peaks(mag.view(), top_n, min_dist, true)?;
if peaks.wavelength.is_empty() {
return Ok(ModeTable::empty());
}
let si = peaks
.scale_idx
.as_ref()
.ok_or_else(|| GspError::Parse("missing scale indices".to_string()))?;
let fi = peaks
.freq_idx
.as_ref()
.ok_or_else(|| GspError::Parse("missing frequency indices".to_string()))?;
let log_spec = spectrum.mapv(|z| (z + Complex64::new(1e-20, 0.0)).ln());
let grad_f = gradient_complex_axis(log_spec.view(), self.freqs.view(), 1)?;
let grad_s = gradient_complex_axis(log_spec.view(), self.wavlen.view(), 0)?;
let mut damping = Array1::<f64>::zeros(peaks.frequency.len());
for idx in 0..peaks.frequency.len() {
let s_idx = si[idx];
let f_idx = fi[idx];
let f0 = peaks.frequency[idx];
let s0 = peaks.wavelength[idx];
let omega_n = 2.0 * std::f64::consts::PI * f0;
let dphi_df = grad_f[[s_idx, f_idx]].im;
let dphi_ds = grad_s[[s_idx, f_idx]].im;
let mut zeta = if dphi_df.abs() < 1e-8 {
let log_mag = mag.mapv(|v| (v + 1e-20).ln());
let first = gradient_real_axis(log_mag.view(), self.freqs.view(), 1)?;
let d2 = gradient_real_axis(first.view(), self.freqs.view(), 1)?;
let curv = d2[[s_idx, f_idx]].min(-1e-10);
((-2.0 / curv).sqrt()) / (2.0 * f0)
} else {
let zeta_f = -1.0 / (omega_n * dphi_df);
let zeta_s = -s0 / (omega_n * dphi_ds * s0 * s0 + 1e-10);
let w_f = dphi_df.abs() + 1e-10;
let w_s = (dphi_ds * s0).abs() + 1e-10;
(w_f * zeta_f + w_s * zeta_s) / (w_f + w_s)
};
if !zeta.is_finite() {
zeta = 0.0;
}
damping[idx] = zeta.clamp(0.0, 1.0);
}
Ok(ModeTable {
frequency: peaks.frequency,
damping,
wavelength: peaks.wavelength,
magnitude: peaks.magnitude,
})
}
pub fn clear_temporal_cache(&mut self) {
self.temporal_matrix_cache = None;
self.temporal_grid_cache = None;
self.temporal_target_cache = None;
}
fn build_temporal_matrix(
&mut self,
t: ArrayView1<f64>,
time_target: f64,
) -> Result<&Array2<Complex64>, GspError> {
let use_cache = self
.temporal_grid_cache
.as_ref()
.zip(self.temporal_target_cache)
.is_some_and(|(tc, tt)| {
tc.len() == t.len()
&& scalar_close(tt, time_target, 1e-12)
&& arrays_close(tc.view(), t, 1e-10)
});
if !use_cache {
let time_grid = t.to_owned();
let mut temporal_matrix = Array2::<Complex64>::zeros((t.len(), self.ts.len()));
for (idx, &sc) in self.ts.iter().enumerate() {
let wavelet = gaussian_wavelet(&time_grid, sc, time_target, self.w0);
temporal_matrix.column_mut(idx).assign(&wavelet);
}
self.temporal_matrix_cache = Some(temporal_matrix);
self.temporal_grid_cache = Some(time_grid);
self.temporal_target_cache = Some(time_target);
}
self.temporal_matrix_cache
.as_ref()
.ok_or_else(|| GspError::Cache("temporal matrix cache missing".to_string()))
}
fn compute_density_clusters(
&self,
peaks: &PeakCloud,
top_n: usize,
min_dist: usize,
) -> ClusterCloud {
if peaks.wavelength.len() < 2 {
return ClusterCloud::empty();
}
let x = peaks.wavelength.mapv(|w| w.log10());
let y = peaks.frequency.clone();
let n = x.len() as f64;
let hx = (1.06 * stddev(x.view()) * n.powf(-0.2)).max(1e-6);
let hy = (1.06 * stddev(y.view()) * n.powf(-0.2)).max(1e-6);
let gx = self.wavlen.mapv(|w| w.log10());
let gy = self.freqs.clone();
let mut z = Array2::<f64>::zeros((gx.len(), gy.len()));
let norm = 1.0 / (2.0 * std::f64::consts::PI * hx * hy * n);
for i in 0..gx.len() {
for j in 0..gy.len() {
let mut acc = 0.0;
for k in 0..x.len() {
let dx = (gx[i] - x[k]) / hx;
let dy = (gy[j] - y[k]) / hy;
acc += (-0.5 * (dx * dx + dy * dy)).exp();
}
z[[i, j]] = norm * acc;
}
}
self.find_peaks(z.view(), top_n, min_dist, false)
.map(|p| ClusterCloud {
wavelength: p.wavelength,
frequency: p.frequency,
density: p.magnitude,
})
.unwrap_or_else(|_| ClusterCloud::empty())
}
fn validate_signal_frame(
&self,
v: ArrayView2<f64>,
t: ArrayView1<f64>,
) -> Result<(), GspError> {
let n_nodes = self.laplacian.rows();
if v.nrows() != n_nodes {
return Err(GspError::Dimensions(format!(
"signal matrix rows {} does not match graph size {}",
v.nrows(),
n_nodes
)));
}
if v.ncols() != t.len() {
return Err(GspError::Dimensions(format!(
"signal time columns {} does not match time vector length {}",
v.ncols(),
t.len()
)));
}
Ok(())
}
fn validate_bus(&self, bus: usize) -> Result<(), GspError> {
let n_nodes = self.laplacian.rows();
if bus >= n_nodes {
return Err(GspError::Index(format!(
"bus {bus} out of bounds for graph with {n_nodes} nodes"
)));
}
Ok(())
}
}