use sekirei_core::nnue::NnueWeights;
use crate::trainer::ConflictGroupStats;
#[derive(Debug, Clone)]
pub struct EpochDiagnostics {
pub param_update_norm: Option<f32>,
pub ft_active_ratio: f32,
pub ft_saturation_ratio: f32,
pub output_mean: f64,
pub output_std: f64,
pub quantized_ft_zero_ratio: f32,
pub l2_ever_active_ratio: f32,
pub l2_ever_saturated_ratio: f32,
pub l2_dead_neurons: usize,
pub l2_activation_frequency_mean: f32,
pub l2_saturation_frequency_mean: f32,
pub l2_activation_frequency_per_neuron: Vec<f32>,
pub l2_saturation_frequency_per_neuron: Vec<f32>,
pub l2_preactivation_p01: f32,
pub l2_preactivation_p10: f32,
pub l2_preactivation_p50: f32,
pub l2_preactivation_p90: f32,
pub l2_preactivation_p99: f32,
pub l2_bias_per_neuron: Vec<f32>,
pub l2_row_weight_norm_per_neuron: Vec<f32>,
pub output_weight_norm: f32,
pub output_bias: f32,
pub ft_grad_norm_mean: f64,
pub ft_grad_norm_std: f64,
pub l2_grad_norm_mean: f64,
pub l2_grad_norm_std: f64,
pub out_grad_norm_mean: f64,
pub out_grad_norm_std: f64,
pub global_grad_norm_p50: f32,
pub global_grad_norm_p90: f32,
pub global_grad_norm_p95: f32,
pub global_grad_norm_p99: f32,
pub ft_update_norm_mean: f64,
pub ft_update_norm_std: f64,
pub l2_update_norm_mean: f64,
pub l2_update_norm_std: f64,
pub out_update_norm_mean: f64,
pub out_update_norm_std: f64,
pub target_mean: f64,
pub target_std: f64,
pub pred_eval_correlation: f64,
pub train_cp_component: f64,
pub train_wdl_component: Option<f64>,
pub grad_clip_count: u64,
pub ft_clip_trigger_rate: f64,
pub l2_clip_trigger_rate: f64,
pub out_clip_trigger_rate: f64,
pub out_grad_norm_p95: f32,
pub out_grad_norm_p99: f32,
pub out_grad_norm_after_mean: f64,
pub out_grad_norm_after_std: f64,
pub masked_position_count: u64,
pub eligible_position_count: u64,
pub ft_dead_neurons: usize,
pub ft_activation_frequency_mean: f32,
pub conflict_group: ConflictGroupSummary,
pub nonconflict_group: ConflictGroupSummary,
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct ConflictGroupSummary {
pub count: u64,
pub cp_residual_abs_mean: f64,
pub cp_residual_abs_std: f64,
pub wdl_residual_abs_mean: f64,
pub wdl_residual_abs_std: f64,
pub ft_grad_norm_mean: f64,
pub ft_grad_norm_std: f64,
pub l2_grad_norm_mean: f64,
pub l2_grad_norm_std: f64,
pub new_dead_ft_mean: f64,
pub new_dead_l2_mean: f64,
}
pub fn build_conflict_group_summary(stats: &ConflictGroupStats) -> ConflictGroupSummary {
let (cp_residual_abs_mean, cp_residual_abs_std) = mean_std(
stats.cp_residual_abs_sum,
stats.cp_residual_abs_sq_sum,
stats.count,
);
let (wdl_residual_abs_mean, wdl_residual_abs_std) = mean_std(
stats.wdl_residual_abs_sum,
stats.wdl_residual_abs_sq_sum,
stats.count,
);
let (ft_grad_norm_mean, ft_grad_norm_std) = mean_std(
stats.ft_grad_norm_sum,
stats.ft_grad_norm_sq_sum,
stats.count,
);
let (l2_grad_norm_mean, l2_grad_norm_std) = mean_std(
stats.l2_grad_norm_sum,
stats.l2_grad_norm_sq_sum,
stats.count,
);
let new_dead_ft_mean = if stats.count > 0 {
stats.new_dead_ft_sum as f64 / stats.count as f64
} else {
0.0
};
let new_dead_l2_mean = if stats.count > 0 {
stats.new_dead_l2_sum as f64 / stats.count as f64
} else {
0.0
};
ConflictGroupSummary {
count: stats.count,
cp_residual_abs_mean,
cp_residual_abs_std,
wdl_residual_abs_mean,
wdl_residual_abs_std,
ft_grad_norm_mean,
ft_grad_norm_std,
l2_grad_norm_mean,
l2_grad_norm_std,
new_dead_ft_mean,
new_dead_l2_mean,
}
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct TraceLayerSnapshot {
pub preactivation_p10: Vec<f32>,
pub preactivation_p50: Vec<f32>,
pub preactivation_p90: Vec<f32>,
pub weighted_input_p10: Vec<f32>,
pub weighted_input_p50: Vec<f32>,
pub weighted_input_p90: Vec<f32>,
pub dead_frequency: Vec<f32>,
pub saturation_frequency: Vec<f32>,
pub weight_row_norm: Vec<f32>,
pub bias: Vec<f32>,
pub gradient_mean: Vec<f32>,
pub gradient_rms: Vec<f32>,
pub gradient_sign_consistency: Vec<f32>,
pub update_norm: Vec<f32>,
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct TraceSnapshot {
pub position_index: u64,
pub l2: TraceLayerSnapshot,
pub ft: TraceLayerSnapshot,
pub l2_input_norm_mean: f64,
pub l2_input_norm_std: f64,
pub ft_output_mean: f64,
pub ft_output_std: f64,
pub cp_wdl: Option<CpWdlTrace>,
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct CpWdlLayerTrace {
pub cp_gradient_mean: Vec<f32>,
pub wdl_gradient_mean: Vec<f32>,
pub cp_gradient_sign_consistency: Vec<f32>,
pub wdl_gradient_sign_consistency: Vec<f32>,
pub cosine_similarity: Vec<f32>,
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct CpWdlTrace {
pub l2: CpWdlLayerTrace,
pub ft: CpWdlLayerTrace,
pub cp_ft_grad_rms: f64,
pub wdl_ft_grad_rms: f64,
pub cp_l2_grad_rms: f64,
pub wdl_l2_grad_rms: f64,
pub cp_out_grad_rms: f64,
pub wdl_out_grad_rms: f64,
pub cp_target_mean: f64,
pub cp_target_std: f64,
pub wdl_target_mean: f64,
pub wdl_target_std: f64,
pub prediction_mean: f64,
pub prediction_std: f64,
pub cp_residual_mean: f64,
pub cp_residual_std: f64,
pub wdl_residual_mean: f64,
pub wdl_residual_std: f64,
pub cp_d_output_mean: f64,
pub cp_d_output_std: f64,
pub wdl_d_output_mean: f64,
pub wdl_d_output_std: f64,
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct SampleGradRecord {
pub game_id: u64,
pub game_result: String,
pub position_index: u64,
pub prediction: f32,
pub cp_target: f32,
pub wdl_target: Option<f32>,
pub cp_d_output: f32,
pub wdl_d_output: Option<f32>,
pub l2_grad_vector: Vec<f32>,
pub l2_grad_norm: f64,
pub cosine_prev: Option<f32>,
pub cosine_running_mean: Option<f32>,
pub l2_gate: Vec<i8>,
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct ShadowTraceRecord {
pub position_index: u64,
pub g_cp_norm: f64,
pub g_wdl_norm: f64,
pub cos_g_cp_wdl: f64,
pub delta_cp_norm: f64,
pub delta_wdl_norm: f64,
pub delta_blend_norm: f64,
pub cos_delta_cp_wdl: f64,
pub contingency_cp_wdl_blend: [u64; 8],
pub blend_dead_linpred_alive: u64,
pub blend_dead_linpred_dead: u64,
pub blend_alive_linpred_dead: u64,
pub blend_alive_linpred_alive: u64,
pub n_alive_at_anchor: u64,
pub l2_dead_frac_cp: f64,
pub l2_dead_frac_wdl: f64,
pub l2_dead_frac_blend: f64,
pub l2_weighted_input_mean_cp: f64,
pub l2_weighted_input_mean_wdl: f64,
pub l2_weighted_input_mean_blend: f64,
pub blend_matches_real_ft: bool,
pub blend_matches_real_l2: bool,
}
fn sign_consistency(pos_count: u64, neg_count: u64) -> f32 {
let total = pos_count + neg_count;
if total == 0 {
return 0.0;
}
(pos_count as f32 - neg_count as f32).abs() / total as f32
}
#[allow(clippy::too_many_arguments)]
pub fn build_trace_layer_snapshot(
values: &[Vec<f32>],
weighted_input_values: &[Vec<f32>],
zero_count: &[u64],
sat_count: &[u64],
sample_count: u64,
weight_row_norm: Vec<f32>,
bias: Vec<f32>,
dacc_sum: &[f64],
dacc_sq_sum: &[f64],
dacc_pos_count: &[u64],
dacc_neg_count: &[u64],
bias_update_sq_sum: &[f64],
) -> TraceLayerSnapshot {
let n = values.len();
let mut preactivation_p10 = Vec::with_capacity(n);
let mut preactivation_p50 = Vec::with_capacity(n);
let mut preactivation_p90 = Vec::with_capacity(n);
for v in values {
let p = percentiles(v, &[0.10, 0.50, 0.90]);
preactivation_p10.push(p[0]);
preactivation_p50.push(p[1]);
preactivation_p90.push(p[2]);
}
let mut weighted_input_p10 = Vec::new();
let mut weighted_input_p50 = Vec::new();
let mut weighted_input_p90 = Vec::new();
for v in weighted_input_values {
let p = percentiles(v, &[0.10, 0.50, 0.90]);
weighted_input_p10.push(p[0]);
weighted_input_p50.push(p[1]);
weighted_input_p90.push(p[2]);
}
let gradient_mean: Vec<f32> = dacc_sum
.iter()
.map(|&s| {
if sample_count > 0 {
(s / sample_count as f64) as f32
} else {
0.0
}
})
.collect();
let gradient_rms: Vec<f32> = dacc_sq_sum
.iter()
.map(|&s| {
if sample_count > 0 {
(s / sample_count as f64).sqrt() as f32
} else {
0.0
}
})
.collect();
let gradient_sign_consistency: Vec<f32> = dacc_pos_count
.iter()
.zip(dacc_neg_count)
.map(|(&p, &n)| sign_consistency(p, n))
.collect();
let update_norm: Vec<f32> = bias_update_sq_sum
.iter()
.map(|&s| (s.sqrt()) as f32)
.collect();
let dead_frequency: Vec<f32> = if sample_count == 0 {
vec![0.0; zero_count.len()]
} else {
zero_count
.iter()
.map(|&z| z as f32 / sample_count as f32)
.collect()
};
TraceLayerSnapshot {
preactivation_p10,
preactivation_p50,
preactivation_p90,
weighted_input_p10,
weighted_input_p50,
weighted_input_p90,
dead_frequency,
saturation_frequency: l2_saturation_frequency_per_neuron(sat_count, sample_count),
weight_row_norm,
bias,
gradient_mean,
gradient_rms,
gradient_sign_consistency,
update_norm,
}
}
fn cosine_similarity(dot_sum: f64, a_sq_sum: f64, b_sq_sum: f64) -> f32 {
let denom = (a_sq_sum * b_sq_sum).sqrt();
if denom == 0.0 || !denom.is_finite() || !dot_sum.is_finite() {
return 0.0;
}
let result = (dot_sum / denom) as f32;
if result.is_finite() { result } else { 0.0 }
}
#[allow(clippy::too_many_arguments)]
pub fn build_cp_wdl_layer_trace(
cp_sum: &[f64],
cp_sq_sum: &[f64],
cp_pos_count: &[u64],
cp_neg_count: &[u64],
wdl_sum: &[f64],
wdl_sq_sum: &[f64],
wdl_pos_count: &[u64],
wdl_neg_count: &[u64],
dot_sum: &[f64],
sample_count: u64,
) -> CpWdlLayerTrace {
let n = cp_sum.len();
let lengths_match = [
cp_sq_sum.len(),
cp_pos_count.len(),
cp_neg_count.len(),
wdl_sum.len(),
wdl_sq_sum.len(),
wdl_pos_count.len(),
wdl_neg_count.len(),
dot_sum.len(),
]
.into_iter()
.all(|length| length == n);
if !lengths_match {
return CpWdlLayerTrace {
cp_gradient_mean: vec![0.0; n],
wdl_gradient_mean: vec![0.0; n],
cp_gradient_sign_consistency: vec![0.0; n],
wdl_gradient_sign_consistency: vec![0.0; n],
cosine_similarity: vec![0.0; n],
};
}
let mean = |sum: &[f64]| -> Vec<f32> {
sum.iter()
.map(|&s| {
if sample_count > 0 {
let result = (s / sample_count as f64) as f32;
if result.is_finite() { result } else { 0.0 }
} else {
0.0
}
})
.collect()
};
CpWdlLayerTrace {
cp_gradient_mean: mean(cp_sum),
wdl_gradient_mean: mean(wdl_sum),
cp_gradient_sign_consistency: (0..n)
.map(|i| sign_consistency(cp_pos_count[i], cp_neg_count[i]))
.collect(),
wdl_gradient_sign_consistency: (0..n)
.map(|i| sign_consistency(wdl_pos_count[i], wdl_neg_count[i]))
.collect(),
cosine_similarity: (0..n)
.map(|i| cosine_similarity(dot_sum[i], cp_sq_sum[i], wdl_sq_sum[i]))
.collect(),
}
}
pub fn ratio(flags: &[bool]) -> f32 {
if flags.is_empty() {
return 0.0;
}
flags.iter().filter(|&&b| b).count() as f32 / flags.len() as f32
}
pub fn mean_std(sum: f64, sum_sq: f64, n: u64) -> (f64, f64) {
if n == 0 || !sum.is_finite() || !sum_sq.is_finite() {
return (0.0, 0.0);
}
let n = n as f64;
let mean = sum / n;
let variance = (sum_sq / n - mean * mean).max(0.0);
let std = variance.sqrt();
if mean.is_finite() && std.is_finite() {
(mean, std)
} else {
(0.0, 0.0)
}
}
pub fn l2_diff_norm(prev: &[f32], curr: &[f32]) -> f32 {
if prev.len() != curr.len() {
return 0.0;
}
let result = prev
.iter()
.zip(curr.iter())
.map(|(a, b)| (a - b) * (a - b))
.sum::<f32>()
.sqrt();
if result.is_finite() { result } else { 0.0 }
}
pub fn quantized_ft_zero_ratio(w: &NnueWeights) -> f32 {
let total: usize = w.ft.iter().map(|row| row.len()).sum();
if total == 0 {
return 0.0;
}
let zeros = w.ft.iter().flatten().filter(|&&v| v == 0).count();
zeros as f32 / total as f32
}
pub fn l2_activation_frequency_per_neuron(zero_count: &[u64], sample_count: u64) -> Vec<f32> {
if sample_count == 0 {
return vec![0.0; zero_count.len()];
}
zero_count
.iter()
.map(|&z| 1.0 - z as f32 / sample_count as f32)
.collect()
}
pub fn l2_saturation_frequency_per_neuron(sat_count: &[u64], sample_count: u64) -> Vec<f32> {
if sample_count == 0 {
return vec![0.0; sat_count.len()];
}
sat_count
.iter()
.map(|&s| s as f32 / sample_count as f32)
.collect()
}
pub fn l2_dead_neurons(zero_count: &[u64], sample_count: u64) -> usize {
if sample_count == 0 {
return 0;
}
zero_count.iter().filter(|&&z| z == sample_count).count()
}
pub fn percentiles(values: &[f32], qs: &[f32]) -> Vec<f32> {
if values.is_empty() {
return vec![0.0; qs.len()];
}
let mut sorted: Vec<f32> = values.to_vec();
sorted.retain(|value| value.is_finite());
if sorted.is_empty() {
return vec![0.0; qs.len()];
}
sorted.sort_by(|a, b| a.total_cmp(b));
qs.iter()
.map(|&q| {
let idx = (q.clamp(0.0, 1.0) * (sorted.len() - 1) as f32).round() as usize;
sorted[idx]
})
.collect()
}
pub fn l2_row_weight_norm_per_neuron(l2: &[f32], rows: usize, cols: usize) -> Vec<f32> {
if rows.checked_mul(cols) != Some(l2.len()) {
return vec![0.0; cols];
}
(0..cols)
.map(|o| {
let result = (0..rows)
.map(|i| {
let v = l2[i * cols + o];
v * v
})
.sum::<f32>()
.sqrt();
if result.is_finite() { result } else { 0.0 }
})
.collect()
}
pub fn output_weight_norm(out: &[f32]) -> f32 {
let result = out.iter().map(|&x| x * x).sum::<f32>().sqrt();
if result.is_finite() { result } else { 0.0 }
}
#[allow(clippy::too_many_arguments)]
pub fn pearson_correlation(
n: u64,
sum_x: f64,
sum_x2: f64,
sum_y: f64,
sum_y2: f64,
sum_xy: f64,
) -> f64 {
if n < 2
|| !sum_x.is_finite()
|| !sum_x2.is_finite()
|| !sum_y.is_finite()
|| !sum_y2.is_finite()
|| !sum_xy.is_finite()
{
return 0.0;
}
let n = n as f64;
let cov = sum_xy - sum_x * sum_y / n;
let var_x = sum_x2 - sum_x * sum_x / n;
let var_y = sum_y2 - sum_y * sum_y / n;
let denom = (var_x * var_y).max(0.0).sqrt();
if denom <= 0.0 || !denom.is_finite() || !cov.is_finite() {
return 0.0;
}
let result = (cov / denom).clamp(-1.0, 1.0);
if result.is_finite() { result } else { 0.0 }
}
pub fn vector_cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() {
return 0.0;
}
let dot: f32 = a.iter().zip(b).map(|(&x, &y)| x * y).sum();
let norm_a: f32 = a.iter().map(|&x| x * x).sum::<f32>().sqrt();
let norm_b: f32 = b.iter().map(|&x| x * x).sum::<f32>().sqrt();
if norm_a <= 0.0
|| norm_b <= 0.0
|| !norm_a.is_finite()
|| !norm_b.is_finite()
|| !dot.is_finite()
{
return 0.0;
}
let result = (dot / (norm_a * norm_b)).clamp(-1.0, 1.0);
if result.is_finite() { result } else { 0.0 }
}
#[cfg(test)]
mod tests {
use super::*;
use sekirei_core::nnue::{L1, L2};
#[test]
fn ratio_of_empty_slice_is_zero() {
assert_eq!(ratio(&[]), 0.0);
}
#[test]
fn ratio_counts_true_fraction() {
assert_eq!(ratio(&[true, false, true, true]), 0.75);
}
#[test]
fn cp_wdl_trace_shape_mismatch_is_zero_filled() {
let trace = build_cp_wdl_layer_trace(
&[1.0, 2.0],
&[1.0],
&[1, 1],
&[0, 0],
&[1.0, 2.0],
&[1.0, 1.0],
&[1, 1],
&[0, 0],
&[1.0, 1.0],
2,
);
assert_eq!(trace.cp_gradient_mean, vec![0.0, 0.0]);
assert_eq!(trace.wdl_gradient_mean, vec![0.0, 0.0]);
assert_eq!(trace.cosine_similarity, vec![0.0, 0.0]);
}
#[test]
fn mean_std_zero_count_is_zero() {
assert_eq!(mean_std(0.0, 0.0, 0), (0.0, 0.0));
}
#[test]
fn mean_std_matches_hand_computed_values() {
let (sum, sum_sq) = [1.0f64, 2.0, 3.0]
.iter()
.fold((0.0, 0.0), |(s, sq), &x| (s + x, sq + x * x));
let (mean, std) = mean_std(sum, sum_sq, 3);
assert!((mean - 2.0).abs() < 1e-9);
assert!((std - (2.0f64 / 3.0).sqrt()).abs() < 1e-9);
}
#[test]
fn mean_std_non_finite_input_is_zero() {
assert_eq!(mean_std(f64::NAN, 1.0, 3), (0.0, 0.0));
assert_eq!(mean_std(1.0, f64::INFINITY, 3), (0.0, 0.0));
}
#[test]
fn l2_diff_norm_zero_for_identical_snapshots() {
let a = [1.0f32, 2.0, 3.0];
assert_eq!(l2_diff_norm(&a, &a), 0.0);
}
#[test]
fn l2_diff_norm_matches_hand_computed_euclidean_distance() {
let a = [0.0f32, 0.0];
let b = [3.0f32, 4.0];
assert_eq!(l2_diff_norm(&a, &b), 5.0); }
#[test]
fn l2_diff_norm_non_finite_input_is_zero() {
assert_eq!(l2_diff_norm(&[f32::NAN], &[1.0]), 0.0);
assert_eq!(l2_diff_norm(&[f32::INFINITY], &[1.0]), 0.0);
}
#[test]
fn l2_diff_norm_length_mismatch_is_zero() {
assert_eq!(l2_diff_norm(&[1.0], &[1.0, 2.0]), 0.0);
}
#[test]
fn cosine_similarity_identical_vectors_is_one() {
let a = [1.0f32, 2.0, -3.0];
assert!((vector_cosine_similarity(&a, &a) - 1.0).abs() < 1e-6);
}
#[test]
fn accumulated_cosine_similarity_non_finite_input_is_zero() {
assert_eq!(cosine_similarity(f64::NAN, 1.0, 1.0), 0.0);
assert_eq!(cosine_similarity(1.0, f64::INFINITY, 1.0), 0.0);
}
#[test]
fn cosine_similarity_opposite_vectors_is_minus_one() {
let a = [1.0f32, 2.0, -3.0];
let b = [-1.0f32, -2.0, 3.0];
assert!((vector_cosine_similarity(&a, &b) - (-1.0)).abs() < 1e-6);
}
#[test]
fn vector_cosine_similarity_non_finite_input_is_zero() {
assert_eq!(vector_cosine_similarity(&[f32::NAN], &[1.0]), 0.0);
assert_eq!(vector_cosine_similarity(&[f32::INFINITY], &[1.0]), 0.0);
}
#[test]
fn vector_cosine_similarity_length_mismatch_is_zero() {
assert_eq!(vector_cosine_similarity(&[1.0], &[1.0, 2.0]), 0.0);
}
#[test]
fn cosine_similarity_orthogonal_vectors_is_zero() {
let a = [1.0f32, 0.0];
let b = [0.0f32, 1.0];
assert_eq!(vector_cosine_similarity(&a, &b), 0.0);
}
#[test]
fn cosine_similarity_zero_vector_is_zero_not_nan() {
let a = [0.0f32, 0.0, 0.0];
let b = [1.0f32, 2.0, 3.0];
assert_eq!(vector_cosine_similarity(&a, &b), 0.0);
}
#[test]
fn quantized_ft_zero_ratio_counts_exact_zeros() {
let mut w = NnueWeights {
ft: vec![[1i16; L1]; 10], ft_bias: [0i16; L1],
l2: vec![[0.0f32; L2]; 2 * L1],
l2_bias: [0.0f32; L2],
out: [0.0f32; L2],
out_bias: 0.0,
};
assert_eq!(quantized_ft_zero_ratio(&w), 0.0);
w.ft[0] = [0i16; L1]; let expected = L1 as f32 / (10 * L1) as f32;
assert!((quantized_ft_zero_ratio(&w) - expected).abs() < 1e-6);
}
#[test]
fn quantized_ft_zero_ratio_of_empty_ft_is_zero() {
let w = NnueWeights {
ft: vec![],
ft_bias: [0i16; L1],
l2: vec![[0.0f32; L2]; 2 * L1],
l2_bias: [0.0f32; L2],
out: [0.0f32; L2],
out_bias: 0.0,
};
assert_eq!(quantized_ft_zero_ratio(&w), 0.0);
}
#[test]
fn l2_activation_frequency_matches_hand_computed_values() {
let zero_count = [2u64, 0, 4];
assert_eq!(
l2_activation_frequency_per_neuron(&zero_count, 4),
vec![0.5, 1.0, 0.0]
);
}
#[test]
fn l2_activation_frequency_zero_samples_is_zero_filled() {
assert_eq!(
l2_activation_frequency_per_neuron(&[1, 2], 0),
vec![0.0, 0.0]
);
}
#[test]
fn l2_saturation_frequency_matches_hand_computed_values() {
let sat_count = [1u64, 4, 0];
assert_eq!(
l2_saturation_frequency_per_neuron(&sat_count, 4),
vec![0.25, 1.0, 0.0]
);
}
#[test]
fn l2_dead_neurons_counts_only_fully_dead() {
let zero_count = [4u64, 3, 4, 0];
assert_eq!(l2_dead_neurons(&zero_count, 4), 2);
}
#[test]
fn l2_dead_neurons_zero_samples_is_zero() {
assert_eq!(l2_dead_neurons(&[0, 0], 0), 0);
}
#[test]
fn percentiles_matches_hand_computed_median_and_extremes() {
let values = [5.0f32, 1.0, 3.0, 2.0, 4.0]; assert_eq!(percentiles(&values, &[0.0, 0.5, 1.0]), vec![1.0, 3.0, 5.0]);
}
#[test]
fn percentiles_of_empty_is_zero_filled() {
assert_eq!(percentiles(&[], &[0.5, 0.9]), vec![0.0, 0.0]);
}
#[test]
fn percentiles_ignores_non_finite_values() {
let values = [f32::NAN, 3.0, f32::INFINITY, 1.0];
assert_eq!(percentiles(&values, &[0.0, 0.5, 1.0]), vec![1.0, 3.0, 3.0]);
assert_eq!(percentiles(&[f32::NAN], &[0.5]), vec![0.0]);
}
#[test]
fn l2_row_weight_norm_matches_hand_computed_euclidean_distance() {
let l2 = [3.0f32, 0.0, 4.0, 0.0];
assert_eq!(l2_row_weight_norm_per_neuron(&l2, 2, 2), vec![5.0, 0.0]);
}
#[test]
fn weight_norms_non_finite_input_are_zero() {
assert_eq!(l2_row_weight_norm_per_neuron(&[f32::NAN], 1, 1), vec![0.0]);
assert_eq!(output_weight_norm(&[f32::INFINITY]), 0.0);
}
#[test]
fn l2_row_weight_norm_length_mismatch_is_zero_filled() {
assert_eq!(
l2_row_weight_norm_per_neuron(&[1.0, 2.0], 1, 3),
vec![0.0, 0.0, 0.0]
);
}
#[test]
fn output_weight_norm_matches_hand_computed_euclidean_norm() {
assert_eq!(output_weight_norm(&[3.0, 4.0]), 5.0);
assert_eq!(output_weight_norm(&[]), 0.0);
}
#[test]
fn pearson_correlation_perfect_positive_line_is_one() {
let (n, mut sx, mut sx2, mut sy, mut sy2, mut sxy) = (3u64, 0.0, 0.0, 0.0, 0.0, 0.0);
for (x, y) in [(1.0, 2.0), (2.0, 4.0), (3.0, 6.0)] {
sx += x;
sx2 += x * x;
sy += y;
sy2 += y * y;
sxy += x * y;
}
let r = pearson_correlation(n, sx, sx2, sy, sy2, sxy);
assert!((r - 1.0).abs() < 1e-9);
}
#[test]
fn pearson_correlation_perfect_negative_line_is_minus_one() {
let (n, mut sx, mut sx2, mut sy, mut sy2, mut sxy) = (3u64, 0.0, 0.0, 0.0, 0.0, 0.0);
for (x, y) in [(1.0, 6.0), (2.0, 4.0), (3.0, 2.0)] {
sx += x;
sx2 += x * x;
sy += y;
sy2 += y * y;
sxy += x * y;
}
let r = pearson_correlation(n, sx, sx2, sy, sy2, sxy);
assert!((r - (-1.0)).abs() < 1e-9);
}
#[test]
fn pearson_correlation_constant_series_is_zero_not_nan() {
let (n, sx, sx2, sy, sy2, sxy) = (3u64, 6.0, 14.0, 15.0, 75.0, 30.0);
assert_eq!(pearson_correlation(n, sx, sx2, sy, sy2, sxy), 0.0);
}
#[test]
fn pearson_correlation_fewer_than_two_samples_is_zero() {
assert_eq!(pearson_correlation(0, 0.0, 0.0, 0.0, 0.0, 0.0), 0.0);
assert_eq!(pearson_correlation(1, 5.0, 25.0, 5.0, 25.0, 25.0), 0.0);
}
#[test]
fn pearson_correlation_non_finite_input_is_zero() {
assert_eq!(pearson_correlation(2, f64::NAN, 1.0, 1.0, 1.0, 1.0), 0.0);
assert_eq!(
pearson_correlation(2, 1.0, 1.0, 1.0, f64::INFINITY, 1.0),
0.0
);
}
}