use std::f64::consts::PI;
use rand::prelude::*;
use crate::error::FdarError;
use crate::iter_maybe_parallel;
use crate::matrix::FdMatrix;
use crate::regression::fdata_to_pc_1d;
#[cfg(feature = "parallel")]
use rayon::iter::ParallelIterator;
pub(crate) const CO_CLUSTER_INIT_PARALLEL_THRESHOLD: usize = 3;
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct BlockParams {
pub mean: Vec<f64>,
pub variance: Vec<f64>,
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct CoClusterResult {
pub row_labels: Vec<usize>,
pub col_labels: Vec<usize>,
pub n_row_blocks: usize,
pub n_col_blocks: usize,
pub block_params: Vec<BlockParams>,
pub row_props: Vec<f64>,
pub col_props: Vec<f64>,
pub log_likelihood: f64,
pub icl: f64,
pub iterations: usize,
pub converged: bool,
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct CoClusterConfig {
pub n_row_blocks: usize,
pub n_col_blocks: usize,
pub ncomp: usize,
pub max_iter: usize,
pub tol: f64,
pub n_init: usize,
pub seed: u64,
}
impl Default for CoClusterConfig {
fn default() -> Self {
Self {
n_row_blocks: 2,
n_col_blocks: 2,
ncomp: 5,
max_iter: 200,
tol: 1e-6,
n_init: 3,
seed: 42,
}
}
}
#[inline]
fn log_gaussian_1d(x: f64, mu: f64, var: f64) -> f64 {
if var <= 0.0 {
return f64::NEG_INFINITY;
}
-0.5 * ((x - mu).powi(2) / var + var.ln() + (2.0 * PI).ln())
}
fn build_block_scores(
data: &FdMatrix,
rotation: &FdMatrix,
mean: &[f64],
weights: &[f64],
col_labels: &[usize],
n: usize,
m: usize,
l_blocks: usize,
eff_ncomp: usize,
) -> Vec<f64> {
let total = n * l_blocks * eff_ncomp;
let mut buf = vec![0.0_f64; total];
for j in 0..m {
let l = col_labels[j];
let w = weights[j];
let mean_j = mean[j];
for i in 0..n {
let val = data[(i, j)] - mean_j;
let base = (i * l_blocks + l) * eff_ncomp;
for k in 0..eff_ncomp {
buf[base + k] += w * val * rotation[(j, k)];
}
}
}
buf
}
fn block_score_reg(block_scores: &[f64], n: usize, l_blocks: usize, eff_ncomp: usize) -> f64 {
const REG_REL: f64 = 1e-6;
if n == 0 || l_blocks == 0 || eff_ncomp == 0 {
return REG_REL;
}
let total_blocks = l_blocks * eff_ncomp;
let mut total_var = 0.0_f64;
let mut n_dims = 0u64;
for l in 0..l_blocks {
for comp in 0..eff_ncomp {
let mut sum = 0.0_f64;
let mut ss = 0.0_f64;
for i in 0..n {
let v = block_scores[(i * l_blocks + l) * eff_ncomp + comp];
sum += v;
ss += v * v;
}
let mean = sum / n as f64;
let var = ss / n as f64 - mean * mean;
total_var += var;
n_dims += 1;
}
}
let _ = total_blocks; let mean_var = if n_dims > 0 {
total_var / n_dims as f64
} else {
0.0
};
if mean_var > 0.0 {
REG_REL * mean_var
} else {
REG_REL
}
}
fn m_step(
block_scores: &[f64],
row_labels: &[usize],
col_labels: &[usize],
n: usize,
m: usize,
k_blocks: usize,
l_blocks: usize,
eff_ncomp: usize,
reg: f64,
) -> (Vec<f64>, Vec<f64>, Vec<BlockParams>) {
let mut row_counts = vec![0usize; k_blocks];
for &r in row_labels {
row_counts[r] += 1;
}
let row_props: Vec<f64> = row_counts.iter().map(|&c| c as f64 / n as f64).collect();
let mut col_counts = vec![0usize; l_blocks];
for &c in col_labels {
col_counts[c] += 1;
}
let col_props: Vec<f64> = col_counts.iter().map(|&c| c as f64 / m as f64).collect();
let mut block_params = Vec::with_capacity(k_blocks * l_blocks);
for k in 0..k_blocks {
for l in 0..l_blocks {
let mut mean = vec![0.0_f64; eff_ncomp];
let mut var = vec![0.0_f64; eff_ncomp];
let mut cnt = 0u64;
for i in 0..n {
if row_labels[i] != k {
continue;
}
cnt += 1;
let base = (i * l_blocks + l) * eff_ncomp;
for comp in 0..eff_ncomp {
mean[comp] += block_scores[base + comp];
}
}
if cnt > 0 {
let nf = cnt as f64;
for comp in 0..eff_ncomp {
mean[comp] /= nf;
}
for i in 0..n {
if row_labels[i] != k {
continue;
}
let base = (i * l_blocks + l) * eff_ncomp;
for comp in 0..eff_ncomp {
let d = block_scores[base + comp] - mean[comp];
var[comp] += d * d;
}
}
for comp in 0..eff_ncomp {
var[comp] = var[comp] / nf + reg;
}
} else {
for comp in 0..eff_ncomp {
var[comp] = reg;
}
}
block_params.push(BlockParams {
mean,
variance: var,
});
}
}
(row_props, col_props, block_params)
}
fn classification_log_likelihood(
block_scores: &[f64],
row_labels: &[usize],
_col_labels: &[usize],
row_props: &[f64],
col_props: &[f64],
block_params: &[BlockParams],
n: usize,
_m: usize,
_k_blocks: usize,
l_blocks: usize,
eff_ncomp: usize,
) -> f64 {
let mut ll = 0.0_f64;
for i in 0..n {
let k = row_labels[i];
let rp = row_props[k];
if rp < 1e-15 {
continue;
}
ll += rp.ln();
for l in 0..l_blocks {
let cp = col_props[l];
if cp < 1e-15 {
continue;
}
let bp = &block_params[k * l_blocks + l];
let base = (i * l_blocks + l) * eff_ncomp;
let mut block_ld = 0.0_f64;
for comp in 0..eff_ncomp {
block_ld +=
log_gaussian_1d(block_scores[base + comp], bp.mean[comp], bp.variance[comp]);
}
ll += cp.ln() + block_ld;
}
}
ll
}
fn e_row_step(
block_scores: &[f64],
row_props: &[f64],
col_props: &[f64],
block_params: &[BlockParams],
n: usize,
k_blocks: usize,
l_blocks: usize,
eff_ncomp: usize,
) -> Vec<usize> {
let mut row_labels = vec![0usize; n];
for i in 0..n {
let mut best_k = 0usize;
let mut best_score = f64::NEG_INFINITY;
for k in 0..k_blocks {
let rp = row_props[k];
if rp < 1e-15 {
continue;
}
let mut score = rp.ln();
for l in 0..l_blocks {
let cp = col_props[l];
if cp < 1e-15 {
continue;
}
let bp = &block_params[k * l_blocks + l];
let base = (i * l_blocks + l) * eff_ncomp;
let mut block_ld = 0.0_f64;
for comp in 0..eff_ncomp {
block_ld += log_gaussian_1d(
block_scores[base + comp],
bp.mean[comp],
bp.variance[comp],
);
}
score += cp.ln() + block_ld;
}
if score > best_score {
best_score = score;
best_k = k;
}
}
row_labels[i] = best_k;
}
row_labels
}
fn e_col_step(
data: &FdMatrix,
rotation: &FdMatrix,
mean: &[f64],
weights: &[f64],
col_labels: &[usize],
row_labels: &[usize],
row_props: &[f64],
col_props: &[f64],
block_params: &[BlockParams],
n: usize,
m: usize,
_k_blocks: usize,
l_blocks: usize,
eff_ncomp: usize,
) -> Vec<usize> {
let mut new_col_labels = col_labels.to_vec();
for j in 0..m {
let w_j = weights[j];
let mean_j = mean[j];
let mut s = vec![0.0_f64; n * eff_ncomp];
for i in 0..n {
let val = w_j * (data[(i, j)] - mean_j);
for comp in 0..eff_ncomp {
s[i * eff_ncomp + comp] = val * rotation[(j, comp)];
}
}
let mut best_l = 0usize;
let mut best_gain = f64::NEG_INFINITY;
for l_cand in 0..l_blocks {
let cp = col_props[l_cand];
if cp < 1e-15 {
continue;
}
let l_curr = col_labels[j];
let mut gain = 0.0_f64;
for i in 0..n {
let k = row_labels[i];
let rp = row_props[k];
if rp < 1e-15 {
continue;
}
let bp_cand = &block_params[k * l_blocks + l_cand];
let mut ld_cand_new = 0.0_f64;
for comp in 0..eff_ncomp {
ld_cand_new += log_gaussian_1d(
s[i * eff_ncomp + comp],
bp_cand.mean[comp],
bp_cand.variance[comp],
);
}
gain += cp.ln() + ld_cand_new;
if l_curr != l_cand {
let bp_curr = &block_params[k * l_blocks + l_curr];
let mut ld_curr = 0.0_f64;
for comp in 0..eff_ncomp {
ld_curr += log_gaussian_1d(
s[i * eff_ncomp + comp],
bp_curr.mean[comp],
bp_curr.variance[comp],
);
}
let cp_curr = col_props[l_curr];
if cp_curr >= 1e-15 {
gain -= cp_curr.ln() + ld_curr;
}
}
}
if gain > best_gain {
best_gain = gain;
best_l = l_cand;
}
}
new_col_labels[j] = best_l;
}
new_col_labels
}
fn col_kmeans_init(data: &FdMatrix, n: usize, m: usize, l_blocks: usize, seed: u64) -> Vec<usize> {
if l_blocks >= m {
return (0..m).map(|j| j % l_blocks).collect();
}
let mut rng = StdRng::seed_from_u64(seed);
let profile_l2sq = |j1: usize, j2: usize| -> f64 {
let c1 = data.column(j1);
let c2 = data.column(j2);
c1.iter().zip(c2.iter()).map(|(a, b)| (a - b).powi(2)).sum()
};
let first = rng.gen_range(0..m);
let mut centers: Vec<usize> = vec![first];
for _ in 1..l_blocks {
let dists: Vec<f64> = (0..m)
.map(|j| {
centers
.iter()
.map(|&c| profile_l2sq(j, c))
.fold(f64::INFINITY, f64::min)
})
.collect();
let total: f64 = dists.iter().sum();
if total < 1e-15 {
centers.push(centers.len() % m);
continue;
}
let threshold = rng.gen::<f64>() * total;
let mut cum = 0.0;
let mut next = m - 1;
for (j, &d) in dists.iter().enumerate() {
cum += d;
if cum >= threshold {
next = j;
break;
}
}
centers.push(next);
}
let mut col_labels: Vec<usize> = (0..m)
.map(|j| {
centers
.iter()
.enumerate()
.map(|(ci, &c)| (ci, profile_l2sq(j, c)))
.min_by(|a, b| a.1.partial_cmp(&b.1).unwrap())
.map(|(ci, _)| ci)
.unwrap_or(0)
})
.collect();
for _ in 0..10 {
let mut cent = vec![0.0_f64; n * l_blocks];
let mut cnt = vec![0u64; l_blocks];
for j in 0..m {
let l = col_labels[j];
cnt[l] += 1;
let col = data.column(j);
for i in 0..n {
cent[l * n + i] += col[i];
}
}
for l in 0..l_blocks {
if cnt[l] > 0 {
let c = cnt[l] as f64;
for i in 0..n {
cent[l * n + i] /= c;
}
}
}
let mut changed = false;
for j in 0..m {
let col = data.column(j);
let mut best_l = 0usize;
let mut best_d = f64::INFINITY;
for l in 0..l_blocks {
let d: f64 = (0..n).map(|i| (col[i] - cent[l * n + i]).powi(2)).sum();
if d < best_d {
best_d = d;
best_l = l;
}
}
if col_labels[j] != best_l {
changed = true;
col_labels[j] = best_l;
}
}
if !changed {
break;
}
}
col_labels
}
#[allow(clippy::too_many_arguments)]
fn cem_single_fit(
data: &FdMatrix,
rotation: &FdMatrix,
mean: &[f64],
weights: &[f64],
init_row_labels: Vec<usize>,
init_col_labels: Vec<usize>,
n: usize,
m: usize,
k_blocks: usize,
l_blocks: usize,
eff_ncomp: usize,
max_iter: usize,
tol: f64,
) -> (CoClusterResult, Vec<f64>) {
let mut row_labels = init_row_labels;
let mut col_labels = init_col_labels;
let mut block_scores = build_block_scores(
data,
rotation,
mean,
weights,
&col_labels,
n,
m,
l_blocks,
eff_ncomp,
);
let reg = block_score_reg(&block_scores, n, l_blocks, eff_ncomp);
let (mut row_props, mut col_props, mut block_params) = m_step(
&block_scores,
&row_labels,
&col_labels,
n,
m,
k_blocks,
l_blocks,
eff_ncomp,
reg,
);
let mut prev_ll = f64::NEG_INFINITY;
let mut per_iter_ll: Vec<f64> = Vec::with_capacity(max_iter);
let mut iterations = 0usize;
let mut converged = false;
for iter in 0..max_iter {
row_labels = e_row_step(
&block_scores,
&row_props,
&col_props,
&block_params,
n,
k_blocks,
l_blocks,
eff_ncomp,
);
col_labels = e_col_step(
data,
rotation,
mean,
weights,
&col_labels,
&row_labels,
&row_props,
&col_props,
&block_params,
n,
m,
k_blocks,
l_blocks,
eff_ncomp,
);
block_scores = build_block_scores(
data,
rotation,
mean,
weights,
&col_labels,
n,
m,
l_blocks,
eff_ncomp,
);
let (rp, cp, bp) = m_step(
&block_scores,
&row_labels,
&col_labels,
n,
m,
k_blocks,
l_blocks,
eff_ncomp,
reg,
);
row_props = rp;
col_props = cp;
block_params = bp;
let ll = classification_log_likelihood(
&block_scores,
&row_labels,
&col_labels,
&row_props,
&col_props,
&block_params,
n,
m,
k_blocks,
l_blocks,
eff_ncomp,
);
per_iter_ll.push(ll);
iterations = iter + 1;
if iter > 0 && (ll - prev_ll).abs() < tol {
converged = true;
break;
}
prev_ll = ll;
}
let log_likelihood = per_iter_ll.last().copied().unwrap_or(f64::NEG_INFINITY);
let p_kl = (k_blocks.saturating_sub(1))
+ (l_blocks.saturating_sub(1))
+ 2 * k_blocks * l_blocks * eff_ncomp;
let icl = log_likelihood - 0.5 * (p_kl as f64) * ((n as f64).ln() + (m as f64).ln());
let result = CoClusterResult {
row_labels,
col_labels,
n_row_blocks: k_blocks,
n_col_blocks: l_blocks,
block_params,
row_props,
col_props,
log_likelihood,
icl,
iterations,
converged,
};
(result, per_iter_ll)
}
#[must_use = "expensive computation whose result should not be discarded"]
pub fn co_cluster(
data: &FdMatrix,
argvals: &[f64],
config: &CoClusterConfig,
) -> Result<CoClusterResult, FdarError> {
let (n, m) = data.shape();
if config.ncomp < 1 {
return Err(FdarError::InvalidParameter {
parameter: "ncomp",
message: format!("ncomp must be >= 1, got {}", config.ncomp),
});
}
if config.n_row_blocks > n {
return Err(FdarError::InvalidParameter {
parameter: "n_row_blocks",
message: format!(
"n_row_blocks={} exceeds number of observations n={}",
config.n_row_blocks, n
),
});
}
if config.n_row_blocks == 0 {
return Err(FdarError::InvalidParameter {
parameter: "n_row_blocks",
message: "n_row_blocks must be >= 1".to_string(),
});
}
if config.n_col_blocks > m {
return Err(FdarError::InvalidParameter {
parameter: "n_col_blocks",
message: format!(
"n_col_blocks={} exceeds number of argument points m={}",
config.n_col_blocks, m
),
});
}
if config.n_col_blocks == 0 {
return Err(FdarError::InvalidParameter {
parameter: "n_col_blocks",
message: "n_col_blocks must be >= 1".to_string(),
});
}
let k_blocks = config.n_row_blocks;
let l_blocks = config.n_col_blocks;
let fpca = fdata_to_pc_1d(data, config.ncomp, argvals)?;
let eff_ncomp = fpca.scores.ncols();
let rotation = &fpca.rotation; let mean = &fpca.mean; let weights = &fpca.weights;
let n_init = config.n_init.max(1);
let run_init = |init: usize| -> Result<CoClusterResult, FdarError> {
let seed = config.seed.wrapping_add(init as u64 * 1000);
use crate::clustering::kmeans_fd;
let km = kmeans_fd(data, argvals, k_blocks, 100, 1e-4, seed)?;
let init_row_labels = km.cluster;
let init_col_labels = col_kmeans_init(data, n, m, l_blocks, seed.wrapping_add(1));
let (result, _per_iter_ll) = cem_single_fit(
data,
rotation,
mean,
weights,
init_row_labels,
init_col_labels,
n,
m,
k_blocks,
l_blocks,
eff_ncomp,
config.max_iter,
config.tol,
);
Ok(result)
};
let results: Vec<CoClusterResult> = if n_init >= CO_CLUSTER_INIT_PARALLEL_THRESHOLD {
iter_maybe_parallel!(0..n_init)
.map(run_init)
.collect::<Result<Vec<_>, _>>()?
} else {
(0..n_init).map(run_init).collect::<Result<Vec<_>, _>>()?
};
let best = results.into_iter().reduce(|acc, r| {
if r.log_likelihood > acc.log_likelihood {
r
} else {
acc
}
});
best.ok_or_else(|| FdarError::ComputationFailed {
operation: "co_cluster",
detail: "all initializations failed".to_string(),
})
}
#[derive(Debug, Clone)]
#[non_exhaustive]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct CoClusterSelectResult {
pub best: CoClusterResult,
pub best_k: usize,
pub best_l: usize,
pub grid_scores: Vec<(usize, usize, f64, usize, f64)>,
pub slope_estimate: f64,
pub penalty_rate: f64,
}
#[must_use = "expensive grid sweep whose result should not be discarded"]
pub fn co_cluster_select(
data: &FdMatrix,
argvals: &[f64],
k_range: &[usize],
l_range: &[usize],
config: &CoClusterConfig,
) -> Result<CoClusterSelectResult, FdarError> {
if k_range.is_empty() {
return Err(FdarError::InvalidParameter {
parameter: "k_range",
message: "k_range must be non-empty".to_string(),
});
}
if l_range.is_empty() {
return Err(FdarError::InvalidParameter {
parameter: "l_range",
message: "l_range must be non-empty".to_string(),
});
}
let grid: Vec<(usize, usize)> = k_range
.iter()
.flat_map(|&k| l_range.iter().map(move |&l| (k, l)))
.collect();
let mut cell_results: Vec<(usize, usize, CoClusterResult)> = Vec::with_capacity(grid.len());
for &(k, l) in &grid {
let mut cell_cfg = config.clone();
cell_cfg.n_row_blocks = k;
cell_cfg.n_col_blocks = l;
let result = co_cluster(data, argvals, &cell_cfg)?;
cell_results.push((k, l, result));
}
struct CellInfo {
k: usize,
l: usize,
ll: f64,
dim: usize,
result_idx: usize,
}
let infos: Vec<CellInfo> = cell_results
.iter()
.enumerate()
.map(|(idx, (k, l, res))| {
let eff_ncomp = if res.block_params.is_empty() {
0
} else {
res.block_params[0].mean.len()
};
let dim = k.saturating_sub(1) + l.saturating_sub(1) + 2 * k * l * eff_ncomp;
CellInfo {
k: *k,
l: *l,
ll: res.log_likelihood,
dim,
result_idx: idx,
}
})
.collect();
let n_grid = infos.len();
let (slope_estimate, penalty_rate) = if n_grid < 4 {
(0.0_f64, 0.0_f64)
} else {
let mut sorted_by_dim: Vec<usize> = (0..n_grid).collect();
sorted_by_dim.sort_by(|&a, &b| infos[b].dim.cmp(&infos[a].dim));
let n_top = (n_grid / 2).max(4).min(n_grid);
let top_idxs = &sorted_by_dim[..n_top];
let d_mean: f64 = top_idxs.iter().map(|&i| infos[i].dim as f64).sum::<f64>() / n_top as f64;
let l_mean: f64 = top_idxs.iter().map(|&i| infos[i].ll).sum::<f64>() / n_top as f64;
let numerator: f64 = top_idxs
.iter()
.map(|&i| (infos[i].dim as f64 - d_mean) * (infos[i].ll - l_mean))
.sum();
let denominator: f64 = top_idxs
.iter()
.map(|&i| (infos[i].dim as f64 - d_mean).powi(2))
.sum();
if denominator.abs() < 1e-10 {
(0.0_f64, 0.0_f64)
} else {
let slope = numerator / denominator;
let pen = 2.0 * slope.abs();
if pen <= 0.0 {
(slope, 0.0_f64)
} else {
(slope, pen)
}
}
};
let penalised: Vec<f64> = infos
.iter()
.map(|ci| {
if penalty_rate > 0.0 {
ci.ll - penalty_rate * ci.dim as f64
} else {
ci.ll
}
})
.collect();
let best_pos = penalised
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap_or(std::cmp::Ordering::Less))
.map(|(i, _)| i)
.unwrap_or(0);
let best_k = infos[best_pos].k;
let best_l = infos[best_pos].l;
let best_result_idx = infos[best_pos].result_idx;
let grid_scores: Vec<(usize, usize, f64, usize, f64)> = infos
.iter()
.enumerate()
.map(|(pos, ci)| (ci.k, ci.l, ci.ll, ci.dim, penalised[pos]))
.collect();
let best = cell_results.remove(best_result_idx).2;
Ok(CoClusterSelectResult {
best,
best_k,
best_l,
grid_scores,
slope_estimate,
penalty_rate,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_helpers::{adjusted_rand_index, uniform_grid};
fn make_block_data(
n: usize,
m: usize,
seed: u64,
) -> (FdMatrix, Vec<f64>, Vec<usize>, Vec<usize>) {
use rand::prelude::*;
use rand_distr::Normal;
let argvals = uniform_grid(m);
let mut rng = StdRng::seed_from_u64(seed);
let noise_dist = Normal::new(0.0_f64, 0.1).unwrap();
let m_half = m / 2;
let mut data = FdMatrix::zeros(n, m);
let mut true_row_labels = vec![0usize; n];
let mut true_col_labels = vec![0usize; m];
for j in m_half..m {
true_col_labels[j] = 1;
}
for i in 0..n {
let row_group = if i < n / 2 { 0 } else { 1 };
true_row_labels[i] = row_group;
let signal = if row_group == 0 { 5.0_f64 } else { -5.0_f64 };
for j in 0..m {
let noise: f64 = rng.sample(noise_dist);
let base = if j < m_half { signal } else { 0.0 };
data[(i, j)] = base + noise;
}
}
(data, argvals, true_row_labels, true_col_labels)
}
fn run_single_cem_with_ll(
data: &FdMatrix,
argvals: &[f64],
k: usize,
l: usize,
ncomp: usize,
seed: u64,
) -> (CoClusterResult, Vec<f64>) {
let (n, m) = data.shape();
let fpca = fdata_to_pc_1d(data, ncomp, argvals).unwrap();
let eff_ncomp = fpca.scores.ncols();
use crate::clustering::kmeans_fd;
let km = kmeans_fd(data, argvals, k, 100, 1e-4, seed).unwrap();
let init_row = km.cluster;
let init_col = col_kmeans_init(data, n, m, l, seed.wrapping_add(1));
cem_single_fit(
data,
&fpca.rotation,
&fpca.mean,
&fpca.weights,
init_row,
init_col,
n,
m,
k,
l,
eff_ncomp,
200,
1e-6,
)
}
#[test]
fn test_co_cluster_smoke() {
let n = 8;
let m = 6;
let argvals = uniform_grid(m);
let data = FdMatrix::zeros(n, m);
let config = CoClusterConfig {
n_row_blocks: 2,
n_col_blocks: 2,
ncomp: 3,
n_init: 1,
..Default::default()
};
let result = co_cluster(&data, &argvals, &config).unwrap();
assert_eq!(result.row_labels.len(), n);
assert_eq!(result.col_labels.len(), m);
assert_eq!(result.block_params.len(), 4);
assert!(result.log_likelihood.is_finite() || result.log_likelihood == f64::NEG_INFINITY);
}
#[test]
fn test_classification_ll_nondecreasing() {
let (data, argvals, _, _) = make_block_data(16, 10, 7777);
let (_result, per_iter_ll) = run_single_cem_with_ll(&data, &argvals, 2, 2, 3, 42);
for w in per_iter_ll.windows(2) {
assert!(
w[1] >= w[0] - 1e-6,
"LL decreased: iter[i]={:.6} -> iter[i+1]={:.6}",
w[0],
w[1]
);
}
}
#[test]
fn test_coclustering_recovers_block_structure() {
let (data, argvals, true_row, true_col) = make_block_data(20, 12, 1234);
let config = CoClusterConfig {
n_row_blocks: 2,
n_col_blocks: 2,
ncomp: 3,
n_init: 3,
seed: 42,
..Default::default()
};
let result = co_cluster(&data, &argvals, &config).unwrap();
let ari_row = adjusted_rand_index(&true_row, &result.row_labels);
let ari_col = adjusted_rand_index(&true_col, &result.col_labels);
assert!(
ari_row > 0.8,
"Row ARI too low: {ari_row:.3} (expected > 0.8)"
);
assert!(
ari_col > 0.8,
"Col ARI too low: {ari_col:.3} (expected > 0.8)"
);
}
#[test]
fn test_determinism_under_seed() {
let (data, argvals, _, _) = make_block_data(16, 10, 999);
let config = CoClusterConfig {
n_row_blocks: 2,
n_col_blocks: 2,
ncomp: 3,
n_init: 2,
seed: 77,
..Default::default()
};
let r1 = co_cluster(&data, &argvals, &config).unwrap();
let r2 = co_cluster(&data, &argvals, &config).unwrap();
assert_eq!(
r1.row_labels, r2.row_labels,
"row_labels differ across runs"
);
assert_eq!(
r1.col_labels, r2.col_labels,
"col_labels differ across runs"
);
assert_eq!(
r1.log_likelihood, r2.log_likelihood,
"log_likelihood differs"
);
assert_eq!(r1.icl, r2.icl, "ICL differs");
}
#[test]
fn test_icl_is_finite() {
let (data, argvals, _, _) = make_block_data(16, 10, 42);
let config = CoClusterConfig {
n_row_blocks: 2,
n_col_blocks: 2,
ncomp: 3,
n_init: 1,
..Default::default()
};
let result = co_cluster(&data, &argvals, &config).unwrap();
assert!(result.icl.is_finite(), "ICL is not finite: {}", result.icl);
}
#[test]
fn test_error_k_exceeds_n() {
let n = 8;
let m = 6;
let data = FdMatrix::zeros(n, m);
let argvals = uniform_grid(m);
let config = CoClusterConfig {
n_row_blocks: 99,
n_col_blocks: 2,
ncomp: 3,
..Default::default()
};
let err = co_cluster(&data, &argvals, &config).unwrap_err();
assert!(
matches!(
err,
FdarError::InvalidParameter {
parameter: "n_row_blocks",
..
}
),
"Expected InvalidParameter(n_row_blocks), got: {err:?}"
);
}
#[test]
fn test_error_l_exceeds_m() {
let n = 8;
let m = 6;
let data = FdMatrix::zeros(n, m);
let argvals = uniform_grid(m);
let config = CoClusterConfig {
n_row_blocks: 2,
n_col_blocks: 99,
ncomp: 3,
..Default::default()
};
let err = co_cluster(&data, &argvals, &config).unwrap_err();
assert!(
matches!(
err,
FdarError::InvalidParameter {
parameter: "n_col_blocks",
..
}
),
"Expected InvalidParameter(n_col_blocks), got: {err:?}"
);
}
#[test]
fn test_error_zero_ncomp() {
let n = 8;
let m = 6;
let data = FdMatrix::zeros(n, m);
let argvals = uniform_grid(m);
let config = CoClusterConfig {
n_row_blocks: 2,
n_col_blocks: 2,
ncomp: 0,
..Default::default()
};
let err = co_cluster(&data, &argvals, &config).unwrap_err();
assert!(
matches!(
err,
FdarError::InvalidParameter {
parameter: "ncomp",
..
}
),
"Expected InvalidParameter(ncomp), got: {err:?}"
);
}
#[test]
fn test_error_argvals_mismatch() {
let n = 8;
let m = 6;
let data = FdMatrix::zeros(n, m);
let argvals = uniform_grid(m + 3); let config = CoClusterConfig {
n_row_blocks: 2,
n_col_blocks: 2,
ncomp: 3,
..Default::default()
};
let err = co_cluster(&data, &argvals, &config).unwrap_err();
assert!(
matches!(err, FdarError::InvalidDimension { .. }),
"Expected InvalidDimension, got: {err:?}"
);
}
#[test]
fn test_co_cluster_select_smoke() {
let n = 8;
let m = 6;
let argvals = uniform_grid(m);
let data = FdMatrix::zeros(n, m);
let config = CoClusterConfig {
ncomp: 2,
n_init: 1,
..Default::default()
};
let result = co_cluster_select(&data, &argvals, &[2, 3], &[2], &config).unwrap();
assert_eq!(
result.grid_scores.len(),
2,
"Expected 2 grid cells (K in {{2,3}}, L=2)"
);
assert_eq!(
result.best.row_labels.len(),
n,
"best.row_labels.len() should equal n"
);
assert_eq!(
result.best.col_labels.len(),
m,
"best.col_labels.len() should equal m"
);
}
#[test]
fn test_slope_heuristic_selects_correct_kl() {
let (data, argvals, true_row, _) = make_block_data(24, 12, 2024);
let config = CoClusterConfig {
ncomp: 3,
n_init: 3,
seed: 42,
..Default::default()
};
let result = co_cluster_select(&data, &argvals, &[2, 3, 4], &[2, 3], &config).unwrap();
assert_eq!(result.grid_scores.len(), 6, "Expected 6 grid cells");
for &(k, l, ll, dim, pen) in &result.grid_scores {
assert!(
ll.is_finite() || ll == f64::NEG_INFINITY,
"grid entry (K={k}, L={l}) has non-finite ll={ll}"
);
let _ = (dim, pen); }
assert_eq!(result.best.row_labels.len(), 24);
let ari = adjusted_rand_index(&true_row, &result.best.row_labels);
assert!(
ari > 0.6,
"Row ARI too low: {ari:.3}. best_k={}, best_l={}",
result.best_k,
result.best_l
);
}
#[test]
fn test_select_single_cell() {
let n = 10;
let m = 8;
let (data, argvals, _, _) = make_block_data(n, m, 42);
let config = CoClusterConfig {
ncomp: 2,
n_init: 1,
seed: 1,
..Default::default()
};
let result = co_cluster_select(&data, &argvals, &[2], &[2], &config).unwrap();
assert_eq!(
result.grid_scores.len(),
1,
"Single-cell grid should have 1 entry"
);
assert_eq!(result.best_k, 2);
assert_eq!(result.best_l, 2);
assert_eq!(
result.slope_estimate, 0.0,
"slope_estimate should be 0 for single-cell"
);
assert_eq!(
result.penalty_rate, 0.0,
"penalty_rate should be 0 for single-cell"
);
}
#[test]
fn test_select_empty_range_errors() {
let n = 8;
let m = 6;
let data = FdMatrix::zeros(n, m);
let argvals = uniform_grid(m);
let config = CoClusterConfig::default();
let err = co_cluster_select(&data, &argvals, &[], &[2], &config).unwrap_err();
assert!(
matches!(
err,
FdarError::InvalidParameter {
parameter: "k_range",
..
}
),
"Expected InvalidParameter(k_range), got: {err:?}"
);
let err = co_cluster_select(&data, &argvals, &[2], &[], &config).unwrap_err();
assert!(
matches!(
err,
FdarError::InvalidParameter {
parameter: "l_range",
..
}
),
"Expected InvalidParameter(l_range), got: {err:?}"
);
}
#[test]
fn test_select_determinism() {
let (data, argvals, _, _) = make_block_data(16, 10, 12345);
let config = CoClusterConfig {
ncomp: 3,
n_init: 2,
seed: 99,
..Default::default()
};
let r1 = co_cluster_select(&data, &argvals, &[2, 3], &[2, 3], &config).unwrap();
let r2 = co_cluster_select(&data, &argvals, &[2, 3], &[2, 3], &config).unwrap();
assert_eq!(r1.best_k, r2.best_k, "best_k differs across runs");
assert_eq!(r1.best_l, r2.best_l, "best_l differs across runs");
assert_eq!(
r1.grid_scores.len(),
r2.grid_scores.len(),
"grid_scores.len() differs"
);
for (a, b) in r1.grid_scores.iter().zip(r2.grid_scores.iter()) {
assert_eq!(a.0, b.0, "K differs in grid_scores");
assert_eq!(a.1, b.1, "L differs in grid_scores");
assert_eq!(a.2, b.2, "log_lik differs in grid_scores");
assert_eq!(a.3, b.3, "model_dim differs in grid_scores");
assert_eq!(a.4, b.4, "penalised_score differs in grid_scores");
}
}
#[test]
fn test_result_surface_populated() {
let n = 10;
let m = 8;
let (data, argvals, _, _) = make_block_data(n, m, 555);
let config = CoClusterConfig {
n_row_blocks: 2,
n_col_blocks: 2,
ncomp: 3,
n_init: 1,
..Default::default()
};
let result = co_cluster(&data, &argvals, &config).unwrap();
assert_eq!(result.row_labels.len(), n, "row_labels.len() != n");
assert_eq!(result.col_labels.len(), m, "col_labels.len() != m");
assert_eq!(
result.block_params.len(),
result.n_row_blocks * result.n_col_blocks,
"block_params.len() != K*L"
);
assert_eq!(result.row_props.len(), result.n_row_blocks);
assert_eq!(result.col_props.len(), result.n_col_blocks);
for bp in &result.block_params {
assert!(!bp.mean.is_empty(), "block_param.mean is empty");
assert_eq!(
bp.mean.len(),
bp.variance.len(),
"mean/variance length mismatch"
);
}
}
}