use crate::error::FdarError;
use crate::iter_maybe_parallel;
use crate::matrix::FdMatrix;
#[cfg(feature = "parallel")]
use rayon::iter::ParallelIterator;
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct GakConfig {
pub sigma: Option<f64>,
}
impl GakConfig {
#[must_use]
pub fn with_sigma(sigma: f64) -> Self {
Self { sigma: Some(sigma) }
}
}
#[inline]
pub(crate) fn logsumexp3(a: f64, b: f64, c: f64) -> f64 {
let max_val = a.max(b).max(c);
if max_val == f64::NEG_INFINITY {
return f64::NEG_INFINITY;
}
if !max_val.is_finite() {
return max_val;
}
let ea = (a - max_val).exp();
let eb = (b - max_val).exp();
let ec = (c - max_val).exp();
max_val + (ea + eb + ec).ln()
}
#[inline]
fn log_local(xi: f64, yj: f64, inv_two_sigma_sq: f64) -> f64 {
let d = xi - yj;
let neg_half_dist = -(d * d) * inv_two_sigma_sq; let h = neg_half_dist.exp(); neg_half_dist - (2.0 - h).ln()
}
pub(crate) fn loggak(x: &[f64], y: &[f64], sigma: f64) -> f64 {
let n = x.len();
let m = y.len();
if n == 0 || m == 0 {
return f64::NEG_INFINITY;
}
let inv_two_sigma_sq = 1.0 / (2.0 * sigma * sigma);
let mut prev = vec![f64::NEG_INFINITY; m + 1];
let mut curr = vec![f64::NEG_INFINITY; m + 1];
prev[0] = 0.0;
for i in 1..=n {
curr[0] = f64::NEG_INFINITY;
let xi = x[i - 1];
for j in 1..=m {
let ll = log_local(xi, y[j - 1], inv_two_sigma_sq);
curr[j] = ll + logsumexp3(prev[j], curr[j - 1], prev[j - 1]);
}
std::mem::swap(&mut prev, &mut curr);
}
prev[m]
}
#[must_use]
pub fn gak(x: &[f64], y: &[f64], sigma: f64) -> f64 {
if sigma <= 0.0 || x.is_empty() || y.is_empty() {
return 0.0;
}
let log_xy = loggak(x, y, sigma);
let log_xx = loggak(x, x, sigma);
let log_yy = loggak(y, y, sigma);
normalize_log(log_xy, log_xx, log_yy)
}
#[inline]
fn normalize_log(log_xy: f64, log_xx: f64, log_yy: f64) -> f64 {
let log_norm = log_xy - 0.5 * (log_xx + log_yy);
if log_norm == f64::NEG_INFINITY {
0.0
} else {
log_norm.exp()
}
}
#[must_use]
pub fn sigma_gak(data: &FdMatrix) -> f64 {
const SIGMA_FLOOR: f64 = 1e-8;
let n = data.nrows();
let m = data.ncols();
if n < 2 || m == 0 {
return SIGMA_FLOOR.max(1.0);
}
let rows: Vec<Vec<f64>> = (0..n).map(|i| data.row(i)).collect();
let mut dists: Vec<f64> = Vec::with_capacity(n * (n - 1) / 2);
for i in 0..n {
for j in (i + 1)..n {
let mut sum = 0.0;
for k in 0..m {
let d = rows[i][k] - rows[j][k];
sum += d * d;
}
dists.push(sum.sqrt());
}
}
if dists.is_empty() {
return SIGMA_FLOOR.max(1.0);
}
dists.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let mid = dists.len() / 2;
let median = if dists.len() % 2 == 0 {
0.5 * (dists[mid - 1] + dists[mid])
} else {
dists[mid]
};
median.max(SIGMA_FLOOR)
}
#[must_use = "expensive computation whose result should not be discarded"]
pub fn gak_gram_matrix(data: &FdMatrix, config: &GakConfig) -> Result<FdMatrix, FdarError> {
let (gram, _diag_log, _sigma, _rows) = build_train_gram(data, config)?;
Ok(gram)
}
#[allow(clippy::type_complexity)]
fn build_train_gram(
data: &FdMatrix,
config: &GakConfig,
) -> Result<(FdMatrix, Vec<f64>, f64, Vec<Vec<f64>>), FdarError> {
let n = data.nrows();
let m = data.ncols();
if n == 0 || m == 0 {
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: "non-empty matrix (nrows > 0, ncols > 0)".to_string(),
actual: format!("{n}x{m}"),
});
}
let sigma = match config.sigma {
Some(s) => s,
None => sigma_gak(data),
};
if sigma.is_nan() || sigma <= 0.0 {
return Err(FdarError::InvalidParameter {
parameter: "sigma",
message: format!("bandwidth must be > 0, got {sigma}"),
});
}
let rows: Vec<Vec<f64>> = (0..n).map(|i| data.row(i)).collect();
let diag_log: Vec<f64> = iter_maybe_parallel!(0..n)
.map(|i| loggak(&rows[i], &rows[i], sigma))
.collect();
let upper_vals: Vec<f64> = iter_maybe_parallel!(0..n)
.flat_map(|i| {
((i + 1)..n)
.map(|j| {
let log_xy = loggak(&rows[i], &rows[j], sigma);
normalize_log(log_xy, diag_log[i], diag_log[j])
})
.collect::<Vec<_>>()
})
.collect();
let mut gram = FdMatrix::zeros(n, n);
for i in 0..n {
gram[(i, i)] = 1.0; }
let mut idx = 0;
for i in 0..n {
for j in (i + 1)..n {
let v = upper_vals[idx];
gram[(i, j)] = v;
gram[(j, i)] = v; idx += 1;
}
}
Ok((gram, diag_log, sigma, rows))
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct GakGramTrain {
pub gram: FdMatrix,
pub(crate) log_self: Vec<f64>,
pub sigma: f64,
pub(crate) train_rows: Vec<Vec<f64>>,
}
impl GakGramTrain {
#[must_use]
pub fn log_self(&self) -> &[f64] {
&self.log_self
}
}
#[must_use = "expensive computation whose result should not be discarded"]
pub fn gak_gram_train(data: &FdMatrix, config: &GakConfig) -> Result<GakGramTrain, FdarError> {
let (gram, log_self, sigma, train_rows) = build_train_gram(data, config)?;
Ok(GakGramTrain {
gram,
log_self,
sigma,
train_rows,
})
}
#[must_use = "expensive computation whose result should not be discarded"]
pub fn gak_gram_predict(train: &GakGramTrain, new_data: &FdMatrix) -> Result<FdMatrix, FdarError> {
let n_train = train.train_rows.len();
debug_assert_eq!(
train.log_self.len(),
n_train,
"log_self length must equal n_train"
);
debug_assert_eq!(
train.gram.nrows(),
n_train,
"training Gram row count must equal n_train"
);
let n_test = new_data.nrows();
let m_test = new_data.ncols();
let sigma = train.sigma;
if n_test == 0 || m_test == 0 {
return Err(FdarError::InvalidDimension {
parameter: "new_data",
expected: "non-empty matrix (nrows > 0, ncols > 0)".to_string(),
actual: format!("{n_test}x{m_test}"),
});
}
if sigma.is_nan() || sigma <= 0.0 {
return Err(FdarError::InvalidParameter {
parameter: "sigma",
message: format!("stored training bandwidth must be > 0, got {sigma}"),
});
}
let m_train = train.train_rows.first().map_or(0, Vec::len);
if m_test != m_train {
return Err(FdarError::InvalidDimension {
parameter: "new_data.ncols",
expected: format!("{m_train} (training evaluation-grid width)"),
actual: format!("{m_test}"),
});
}
let test_rows: Vec<Vec<f64>> = (0..n_test).map(|t| new_data.row(t)).collect();
let test_self: Vec<f64> = iter_maybe_parallel!(0..n_test)
.map(|t| loggak(&test_rows[t], &test_rows[t], sigma))
.collect();
let row_blocks: Vec<Vec<f64>> = iter_maybe_parallel!(0..n_test)
.map(|t| {
let mut row = vec![0.0; n_train];
for j in 0..n_train {
let log_xy = loggak(&test_rows[t], &train.train_rows[j], sigma);
row[j] = normalize_log(log_xy, test_self[t], train.log_self[j]);
}
row
})
.collect();
let mut out = FdMatrix::zeros(n_test, n_train);
for (t, row) in row_blocks.iter().enumerate() {
for (j, &v) in row.iter().enumerate() {
out[(t, j)] = v;
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
fn matrix_from_rows(rows: &[Vec<f64>]) -> FdMatrix {
let n = rows.len();
let m = rows[0].len();
let mut data = vec![0.0; n * m];
for (i, r) in rows.iter().enumerate() {
for (j, &v) in r.iter().enumerate() {
data[i + j * n] = v; }
}
FdMatrix::from_slice(&data, n, m).unwrap()
}
#[test]
fn test_logsumexp3_basic() {
assert!((logsumexp3(0.0, 0.0, 0.0) - 3.0_f64.ln()).abs() < 1e-15);
let a = logsumexp3(1.0, f64::NEG_INFINITY, f64::NEG_INFINITY);
assert!((a - 1.0).abs() < 1e-15);
let z = logsumexp3(f64::NEG_INFINITY, f64::NEG_INFINITY, f64::NEG_INFINITY);
assert_eq!(z, f64::NEG_INFINITY);
assert!(!z.is_nan());
}
#[test]
fn test_gak_no_underflow() {
let m = 200;
let x: Vec<f64> = (0..m).map(|k| (k as f64 * 0.05).sin()).collect();
let y: Vec<f64> = (0..m).map(|k| (k as f64 * 0.05 + 0.3).sin()).collect();
let k = gak(&x, &y, 1.0);
assert!(k > 1e-10, "GAK underflowed to {k} on m={m} series");
assert!(k <= 1.0 + 1e-12);
assert!(k.is_finite());
}
#[test]
fn test_gak_normalized_range() {
let rows = vec![
(0..80).map(|k| (k as f64 * 0.1).sin()).collect::<Vec<_>>(),
(0..80).map(|k| (k as f64 * 0.1).cos()).collect::<Vec<_>>(),
(0..80).map(|k| k as f64).collect::<Vec<_>>(), (0..80).map(|k| -(k as f64) * 3.0).collect::<Vec<_>>(),
];
let data = matrix_from_rows(&rows);
let gram = gak_gram_matrix(&data, &GakConfig::with_sigma(2.0)).unwrap();
let n = gram.nrows();
for i in 0..n {
assert!((gram[(i, i)] - 1.0).abs() < 1e-12, "diag[{i}] != 1");
for j in 0..n {
let v = gram[(i, j)];
assert!(v.is_finite(), "non-finite entry ({i},{j}) = {v}");
assert!(
(0.0..=1.0 + 1e-12).contains(&v),
"entry ({i},{j}) = {v} out of [0,1]"
);
}
}
}
#[test]
fn test_gak_gram_symmetric() {
let rows: Vec<Vec<f64>> = (0..6)
.map(|i| (0..40).map(|k| ((k + i) as f64 * 0.2).sin()).collect())
.collect();
let data = matrix_from_rows(&rows);
let gram = gak_gram_matrix(&data, &GakConfig::with_sigma(1.5)).unwrap();
let n = gram.nrows();
for i in 0..n {
for j in 0..n {
assert_eq!(
gram[(i, j)].to_bits(),
gram[(j, i)].to_bits(),
"asymmetry at ({i},{j})"
);
}
}
}
#[test]
fn test_gak_gram_psd() {
let rows: Vec<Vec<f64>> = (0..8)
.map(|i| {
(0..50)
.map(|k| (k as f64 * 0.15 + i as f64 * 0.4).sin() + 0.1 * i as f64)
.collect()
})
.collect();
let data = matrix_from_rows(&rows);
let gram = gak_gram_matrix(&data, &GakConfig::with_sigma(2.0)).unwrap();
let dm = gram.to_dmatrix();
let eig = dm.symmetric_eigenvalues();
let min_eig = eig.iter().cloned().fold(f64::INFINITY, f64::min);
assert!(
min_eig >= -1e-8,
"min eigenvalue {min_eig} < -1e-8 (not PSD)"
);
}
#[test]
fn test_gak_parallel_matches_sequential() {
let rows: Vec<Vec<f64>> = (0..7)
.map(|i| (0..60).map(|k| ((k * 3 + i) as f64 * 0.07).cos()).collect())
.collect();
let data = matrix_from_rows(&rows);
let cfg = GakConfig::with_sigma(1.2);
let g1 = gak_gram_matrix(&data, &cfg).unwrap();
let g2 = gak_gram_matrix(&data, &cfg).unwrap();
let n = g1.nrows();
for i in 0..n {
for j in 0..n {
assert_eq!(
g1[(i, j)].to_bits(),
g2[(i, j)].to_bits(),
"nondeterministic entry ({i},{j})"
);
}
}
}
#[test]
fn test_sigma_gak_healthy() {
let rows: Vec<Vec<f64>> = (0..10)
.map(|i| {
(0..60)
.map(|k| (k as f64 * 0.12 + i as f64 * 0.15).sin())
.collect()
})
.collect();
let data = matrix_from_rows(&rows);
let sigma = sigma_gak(&data);
assert!(sigma > 0.0, "sigma heuristic returned non-positive {sigma}");
let gram = gak_gram_matrix(&data, &GakConfig::default()).unwrap();
let n = gram.nrows();
let mut min_off = f64::INFINITY;
let mut max_off = f64::NEG_INFINITY;
for i in 0..n {
for j in 0..n {
if i != j {
min_off = min_off.min(gram[(i, j)]);
max_off = max_off.max(gram[(i, j)]);
}
}
}
assert!(
max_off < 0.999,
"Gram near-constant (max off-diag {max_off})"
);
assert!(
min_off > 1e-4,
"Gram near-identity (min off-diag {min_off})"
);
assert!(max_off - min_off > 0.05, "off-diagonal range too narrow");
}
#[test]
fn test_sigma_gak_floor_on_identical() {
let rows = vec![vec![1.0; 20], vec![1.0; 20], vec![1.0; 20]];
let data = matrix_from_rows(&rows);
let sigma = sigma_gak(&data);
assert!(sigma > 0.0, "sigma floor failed: {sigma}");
let gram = gak_gram_matrix(&data, &GakConfig::default()).unwrap();
for i in 0..3 {
for j in 0..3 {
assert!((gram[(i, j)] - 1.0).abs() < 1e-9);
}
}
}
#[test]
fn test_gak_vs_reference() {
let x = [0.0, 1.0];
let y = [0.0, 2.0];
let got = gak(&x, &y, 1.0);
let expected = 0.448_432_219_612_369_95_f64;
assert!(
(got - expected).abs() < 1e-9,
"GAK reference mismatch: got {got}, expected {expected}"
);
let x3 = [0.0, 1.0, 2.0];
let y3 = [0.0, 1.0, 3.0];
let got3 = gak(&x3, &y3, 2.0);
let expected3 = 0.805_752_775_914_924_f64;
assert!(
(got3 - expected3).abs() < 1e-9,
"GAK length-3 reference mismatch: got {got3}, expected {expected3}"
);
assert!((gak(&x, &x, 1.0) - 1.0).abs() < 1e-12);
}
#[test]
fn test_gak_gram_empty_errors() {
let empty = FdMatrix::zeros(0, 0);
assert!(matches!(
gak_gram_matrix(&empty, &GakConfig::with_sigma(1.0)),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn test_gak_gram_bad_sigma_errors() {
let data = matrix_from_rows(&[vec![0.0, 1.0], vec![1.0, 2.0]]);
assert!(matches!(
gak_gram_matrix(&data, &GakConfig::with_sigma(-1.0)),
Err(FdarError::InvalidParameter { .. })
));
assert!(matches!(
gak_gram_matrix(&data, &GakConfig::with_sigma(0.0)),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn test_gram_train_shape_psd() {
let rows: Vec<Vec<f64>> = (0..8)
.map(|i| {
(0..50)
.map(|k| (k as f64 * 0.15 + i as f64 * 0.4).sin() + 0.1 * i as f64)
.collect()
})
.collect();
let data = matrix_from_rows(&rows);
let fit = gak_gram_train(&data, &GakConfig::with_sigma(2.0)).unwrap();
let n = data.nrows();
assert_eq!(fit.gram.shape(), (n, n));
assert_eq!(fit.log_self().len(), n);
assert!(fit.sigma > 0.0);
for i in 0..n {
assert!((fit.gram[(i, i)] - 1.0).abs() < 1e-12);
for j in 0..n {
assert_eq!(fit.gram[(i, j)].to_bits(), fit.gram[(j, i)].to_bits());
}
}
let eig = fit.gram.to_dmatrix().symmetric_eigenvalues();
let min_eig = eig.iter().cloned().fold(f64::INFINITY, f64::min);
assert!(min_eig >= -1e-8, "min eig {min_eig} < -1e-8 (not PSD)");
}
#[test]
fn test_gram_predict_shape() {
let train_rows: Vec<Vec<f64>> = (0..5)
.map(|i| (0..30).map(|k| ((k + i) as f64 * 0.2).sin()).collect())
.collect();
let train = matrix_from_rows(&train_rows);
let fit = gak_gram_train(&train, &GakConfig::with_sigma(1.5)).unwrap();
let test_rows: Vec<Vec<f64>> = (0..3)
.map(|i| (0..30).map(|k| ((k + i) as f64 * 0.25).cos()).collect())
.collect();
let test = matrix_from_rows(&test_rows);
let k = gak_gram_predict(&fit, &test).unwrap();
assert_eq!(k.shape(), (3, 5), "predict Gram must be n_test × n_train");
}
#[test]
fn test_gram_predict_normalized() {
let train_rows: Vec<Vec<f64>> = (0..4)
.map(|i| (0..40).map(|k| ((k + i * 3) as f64 * 0.1).sin()).collect())
.collect();
let train = matrix_from_rows(&train_rows);
let fit = gak_gram_train(&train, &GakConfig::with_sigma(1.0)).unwrap();
let test_rows = vec![
train_rows[2].clone(),
(0..40).map(|k| (k as f64 * 0.07).cos()).collect(),
];
let test = matrix_from_rows(&test_rows);
let k = gak_gram_predict(&fit, &test).unwrap();
let (nt, ntr) = k.shape();
for t in 0..nt {
for j in 0..ntr {
let v = k[(t, j)];
assert!(v.is_finite(), "non-finite ({t},{j}) = {v}");
assert!((0.0..=1.0 + 1e-12).contains(&v), "({t},{j}) = {v} ∉ [0,1]");
}
}
assert!(
(k[(0, 2)] - 1.0).abs() < 1e-9,
"identical test curve should score ≈1 in col 2, got {}",
k[(0, 2)]
);
}
#[test]
fn test_gram_predict_reproduces_train() {
let rows: Vec<Vec<f64>> = (0..6)
.map(|i| (0..45).map(|k| ((k + i * 2) as f64 * 0.13).sin()).collect())
.collect();
let data = matrix_from_rows(&rows);
let fit = gak_gram_train(&data, &GakConfig::with_sigma(1.8)).unwrap();
let k = gak_gram_predict(&fit, &data).unwrap();
let n = data.nrows();
assert_eq!(k.shape(), (n, n));
for i in 0..n {
for j in 0..n {
assert!(
(k[(i, j)] - fit.gram[(i, j)]).abs() < 1e-12,
"predict({i},{j})={} vs train={}",
k[(i, j)],
fit.gram[(i, j)]
);
}
}
}
#[test]
fn test_gram_predict_sigma_consistency() {
let train_rows: Vec<Vec<f64>> = (0..5)
.map(|i| (0..40).map(|k| ((k + i) as f64 * 0.1).sin()).collect())
.collect();
let train = matrix_from_rows(&train_rows);
let explicit_sigma = 3.7;
let fit = gak_gram_train(&train, &GakConfig::with_sigma(explicit_sigma)).unwrap();
assert!((fit.sigma - explicit_sigma).abs() < 1e-15);
let test_rows: Vec<Vec<f64>> = (0..3)
.map(|i| {
(0..40)
.map(|k| ((k + i) as f64 * 0.1).sin() * 50.0)
.collect()
})
.collect();
let test = matrix_from_rows(&test_rows);
let sigma_test = sigma_gak(&test);
assert!(
(sigma_test - explicit_sigma).abs() > 1.0,
"test set's own σ ({sigma_test}) should differ from train σ"
);
let k = gak_gram_predict(&fit, &test).unwrap();
let t0 = test.row(0);
let tr0 = train.row(0);
let expected = {
let log_xy = loggak(&t0, &tr0, explicit_sigma);
let log_xx = loggak(&t0, &t0, explicit_sigma);
normalize_log(log_xy, log_xx, fit.log_self()[0])
};
assert!(
(k[(0, 0)] - expected).abs() < 1e-12,
"predict did not use train.sigma: got {}, expected {expected}",
k[(0, 0)]
);
}
#[test]
fn test_gram_predict_empty_and_grid_errors() {
let train = matrix_from_rows(&[vec![0.0, 1.0, 2.0], vec![1.0, 2.0, 3.0]]);
let fit = gak_gram_train(&train, &GakConfig::with_sigma(1.0)).unwrap();
let empty = FdMatrix::zeros(0, 0);
assert!(matches!(
gak_gram_predict(&fit, &empty),
Err(FdarError::InvalidDimension { .. })
));
let bad = matrix_from_rows(&[vec![0.0, 1.0]]);
assert!(matches!(
gak_gram_predict(&fit, &bad),
Err(FdarError::InvalidDimension { .. })
));
}
}