use std::collections::VecDeque;
pub struct DivergenceReport {
pub deltas: Vec<f64>,
pub pre_norms: Option<Vec<f64>>,
pub post_norm: Option<f64>,
}
impl DivergenceReport {
pub fn max_relative_delta(&self) -> f64 {
self.deltas.iter().copied().fold(0.0_f64, f64::max)
}
pub fn cosine_similarities(&self) -> Option<Vec<f64>> {
let pre_norms = self.pre_norms.as_ref()?;
let post_norm = self.post_norm?;
if post_norm < 1e-10 {
return None;
}
Some(
self.deltas
.iter()
.zip(pre_norms)
.map(|(&delta, &pre_norm)| {
if pre_norm < 1e-10 {
return 0.0;
}
let diff_sq = (delta * post_norm).powi(2);
let pre_sq = pre_norm.powi(2);
let post_sq = post_norm.powi(2);
((pre_sq + post_sq - diff_sq) / (2.0 * pre_norm * post_norm)).clamp(-1.0, 1.0)
})
.collect(),
)
}
pub fn magnitude_shifts(&self) -> Option<Vec<f64>> {
let pre_norms = self.pre_norms.as_ref()?;
let post_norm = self.post_norm?;
Some(pre_norms.iter().map(|&pre| pre - post_norm).collect())
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[non_exhaustive]
pub enum ConvergenceAction {
Stable,
SuppressGrowth,
NudgeDown { factor: f64 },
}
pub(crate) fn reset_divergence_signals(
divergence: &mut [Option<f64>],
pre_norm: &mut [Option<f64>],
post_norm: &mut Option<f64>,
) {
for d in divergence {
*d = None;
}
for p in pre_norm {
*p = None;
}
*post_norm = None;
}
pub trait ConvergenceGuard: Send + Sync {
fn report(
&mut self,
report: &DivergenceReport,
k_used: usize,
k_max: usize,
) -> ConvergenceAction;
fn telemetry(&self) -> Vec<(&'static str, f64)> {
Vec::new()
}
fn reset(&mut self) {}
fn trend_history(&self) -> Option<Vec<f64>> {
None
}
fn clone_box(&self) -> Box<dyn ConvergenceGuard>;
}
impl Clone for Box<dyn ConvergenceGuard> {
fn clone(&self) -> Self {
self.clone_box()
}
}
#[derive(Debug, Clone)]
pub struct LambdaSample {
pub d_raw: f64,
pub lambda_raw: Option<f64>,
pub lambda_ema: Option<f64>,
pub k_used: usize,
pub k_max: usize,
}
#[derive(Clone)]
pub struct LambdaEstimator {
prev_d: Option<f64>,
ema_raw: f64,
ema_t: u32,
alpha: f64,
noise_floor: f64,
}
impl Default for LambdaEstimator {
fn default() -> Self {
Self {
prev_d: None,
ema_raw: 0.0,
ema_t: 0,
alpha: 0.9,
noise_floor: 1e-8,
}
}
}
impl LambdaEstimator {
pub fn with_alpha(alpha: f64) -> Self {
let alpha = if alpha.is_finite() {
alpha.clamp(f64::EPSILON, 1.0)
} else {
Self::default().alpha
};
Self {
alpha,
..Default::default()
}
}
pub fn observe(&mut self, d_raw: f64, k_used: usize, k_max: usize) -> LambdaSample {
let lambda_raw = match self.prev_d {
Some(prev) if prev > self.noise_floor && d_raw > self.noise_floor && k_max > 0 => {
Some((d_raw / prev).ln() / k_max as f64)
}
_ => None,
};
if let Some(l) = lambda_raw {
self.ema_raw = self.alpha * self.ema_raw + (1.0 - self.alpha) * l;
self.ema_t = self.ema_t.saturating_add(1);
}
self.prev_d = if d_raw > self.noise_floor {
Some(d_raw)
} else {
None
};
let lambda_ema = if self.ema_t == 0 {
None
} else {
let denom = 1.0 - self.alpha.powi(self.ema_t as i32);
Some(if denom > 0.0 {
self.ema_raw / denom
} else {
self.ema_raw
})
};
LambdaSample {
d_raw,
lambda_raw,
lambda_ema,
k_used,
k_max,
}
}
pub fn reset(&mut self) {
self.prev_d = None;
self.ema_raw = 0.0;
self.ema_t = 0;
}
}
#[derive(Debug, Default, Clone, Copy)]
pub struct NoGuard;
impl ConvergenceGuard for NoGuard {
fn clone_box(&self) -> Box<dyn ConvergenceGuard> {
Box::new(*self)
}
fn report(&mut self, _: &DivergenceReport, _: usize, _: usize) -> ConvergenceAction {
ConvergenceAction::Stable
}
}
#[derive(Clone)]
pub struct TrendGuard {
threshold: f64,
history: VecDeque<f64>,
}
impl Default for TrendGuard {
fn default() -> Self {
Self::new(0.01)
}
}
impl TrendGuard {
pub fn new(threshold: f64) -> Self {
Self {
threshold,
history: VecDeque::with_capacity(6),
}
}
pub fn with_threshold(mut self, threshold: f64) -> Self {
self.threshold = threshold;
self
}
pub fn with_history<I>(mut self, history: I) -> Self
where
I: IntoIterator<Item = f64>,
{
let mut ring: VecDeque<f64> = history.into_iter().collect();
while ring.len() > 5 {
ring.pop_front();
}
self.history = ring;
self
}
pub fn history(&self) -> &VecDeque<f64> {
&self.history
}
fn check_trend(&self) -> ConvergenceAction {
if self.history.len() < 3 {
return ConvergenceAction::Stable;
}
let len = self.history.len();
let rising = self.history[len - 1] > self.history[len - 2]
&& self.history[len - 2] > self.history[len - 3]
&& self.history[len - 1] > self.threshold;
if rising {
crate::verbose!(
" ddp: weight-space divergence trending up | history={:.4?} | suppressing growth",
Vec::from(self.history.clone()),
);
ConvergenceAction::SuppressGrowth
} else {
ConvergenceAction::Stable
}
}
}
impl ConvergenceGuard for TrendGuard {
fn clone_box(&self) -> Box<dyn ConvergenceGuard> {
Box::new(self.clone())
}
fn report(&mut self, report: &DivergenceReport, _: usize, _: usize) -> ConvergenceAction {
let divergence = report.max_relative_delta();
if self.history.len() >= 5 {
self.history.pop_front();
}
self.history.push_back(divergence);
self.check_trend()
}
fn telemetry(&self) -> Vec<(&'static str, f64)> {
match self.history.back() {
Some(&d) => vec![("d_history_last", d)],
None => Vec::new(),
}
}
fn reset(&mut self) {
self.history.clear();
}
fn trend_history(&self) -> Option<Vec<f64>> {
if self.history.is_empty() {
None
} else {
Some(self.history.iter().copied().collect())
}
}
}
#[derive(Clone)]
pub struct MsfGuard {
estimator: LambdaEstimator,
suppress_threshold: f64,
suppress_sustain: usize,
nudge_threshold: f64,
nudge_sustain: usize,
nudge_factor: f64,
suppress_streak: usize,
nudge_streak: usize,
last_sample: Option<LambdaSample>,
}
impl Default for MsfGuard {
fn default() -> Self {
Self {
estimator: LambdaEstimator::default(),
suppress_threshold: 1.0e-3,
suppress_sustain: 3,
nudge_threshold: 1.0e-2,
nudge_sustain: 3,
nudge_factor: 0.5,
suppress_streak: 0,
nudge_streak: 0,
last_sample: None,
}
}
}
impl MsfGuard {
pub fn with_alpha(mut self, alpha: f64) -> Self {
self.estimator = LambdaEstimator::with_alpha(alpha);
self
}
pub fn with_suppress(mut self, threshold: f64, sustain: usize) -> Self {
self.suppress_threshold = threshold;
self.suppress_sustain = sustain;
self
}
pub fn with_nudge(mut self, threshold: f64, sustain: usize, factor: f64) -> Self {
self.nudge_threshold = threshold;
self.nudge_sustain = sustain;
self.nudge_factor = factor;
self
}
pub fn without_nudge(mut self) -> Self {
self.nudge_threshold = f64::INFINITY;
self
}
pub fn last_sample(&self) -> Option<&LambdaSample> {
self.last_sample.as_ref()
}
}
impl ConvergenceGuard for MsfGuard {
fn clone_box(&self) -> Box<dyn ConvergenceGuard> {
Box::new(self.clone())
}
fn report(
&mut self,
report: &DivergenceReport,
k_used: usize,
k_max: usize,
) -> ConvergenceAction {
let d_raw = report.max_relative_delta();
let sample = self.estimator.observe(d_raw, k_used, k_max);
let lambda_ema = sample.lambda_ema;
self.last_sample = Some(sample);
let lema = lambda_ema.unwrap_or(f64::NEG_INFINITY);
if lema > self.suppress_threshold {
self.suppress_streak += 1;
} else {
self.suppress_streak = 0;
}
if lema > self.nudge_threshold {
self.nudge_streak += 1;
} else {
self.nudge_streak = 0;
}
if self.nudge_streak >= self.nudge_sustain && self.nudge_threshold.is_finite() {
self.nudge_streak = 0;
crate::verbose!(
" ddp: msf λ_ema={:.4e} sustained > nudge_threshold {:.4e} | nudging anchor down by {:.2}",
lema, self.nudge_threshold, self.nudge_factor,
);
return ConvergenceAction::NudgeDown {
factor: self.nudge_factor,
};
}
if self.suppress_streak >= self.suppress_sustain {
self.suppress_streak = 0;
crate::verbose!(
" ddp: msf λ_ema={:.4e} sustained > suppress_threshold {:.4e} | suppressing growth",
lema, self.suppress_threshold,
);
return ConvergenceAction::SuppressGrowth;
}
ConvergenceAction::Stable
}
fn telemetry(&self) -> Vec<(&'static str, f64)> {
let s = match &self.last_sample {
Some(s) => s,
None => return Vec::new(),
};
let mut out = Vec::with_capacity(2);
if let Some(l) = s.lambda_raw {
out.push(("lambda_raw", l));
}
if let Some(l) = s.lambda_ema {
out.push(("lambda_ema", l));
}
out
}
fn reset(&mut self) {
self.estimator.reset();
self.suppress_streak = 0;
self.nudge_streak = 0;
self.last_sample = None;
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_report(deltas: &[f64]) -> DivergenceReport {
DivergenceReport {
deltas: deltas.to_vec(),
pre_norms: None,
post_norm: None,
}
}
fn make_full_report(deltas: &[f64], pre_norms: &[f64], post_norm: f64) -> DivergenceReport {
DivergenceReport {
deltas: deltas.to_vec(),
pre_norms: Some(pre_norms.to_vec()),
post_norm: Some(post_norm),
}
}
#[test]
fn max_relative_delta_picks_worst_rank() {
let r = make_report(&[0.01, 0.05, 0.03]);
assert!((r.max_relative_delta() - 0.05).abs() < 1e-10);
}
#[test]
fn max_relative_delta_empty_is_zero() {
let r = make_report(&[]);
assert_eq!(r.max_relative_delta(), 0.0);
}
#[test]
fn cosine_similarities_none_when_missing_norms() {
let r = make_report(&[0.01, 0.02]);
assert!(r.cosine_similarities().is_none());
assert!(r.magnitude_shifts().is_none());
}
#[test]
fn cosine_similarities_correct_for_small_delta() {
let r = make_full_report(&[0.001], &[10.0], 10.0);
let cos = r.cosine_similarities().unwrap();
assert!(cos[0] > 0.999, "expected cos near 1.0, got {}", cos[0]);
}
#[test]
fn magnitude_shifts_correct() {
let r = make_full_report(&[0.01, 0.02], &[10.5, 9.8], 10.0);
let shifts = r.magnitude_shifts().unwrap();
assert!((shifts[0] - 0.5).abs() < 1e-10);
assert!((shifts[1] - (-0.2)).abs() < 1e-10);
}
#[test]
fn no_guard_is_always_stable() {
let mut g = NoGuard;
for _ in 0..10 {
assert_eq!(
g.report(&make_report(&[0.5, 0.5]), 8, 4),
ConvergenceAction::Stable
);
}
}
#[test]
fn no_guard_emits_no_telemetry() {
let g = NoGuard;
assert!(g.telemetry().is_empty());
}
#[test]
fn trend_default_threshold_is_0_01() {
let g = TrendGuard::default();
let mut g = g;
g.report(&make_report(&[0.011]), 8, 4);
g.report(&make_report(&[0.012]), 8, 4);
assert_eq!(
g.report(&make_report(&[0.013]), 8, 4),
ConvergenceAction::SuppressGrowth
);
}
#[test]
fn trend_3_rises_above_threshold_suppress() {
let mut g = TrendGuard::new(0.01);
assert_eq!(
g.report(&make_report(&[0.02]), 8, 4),
ConvergenceAction::Stable
);
assert_eq!(
g.report(&make_report(&[0.03]), 8, 4),
ConvergenceAction::Stable
);
assert_eq!(
g.report(&make_report(&[0.04]), 8, 4),
ConvergenceAction::SuppressGrowth
);
}
#[test]
fn trend_non_rising_is_stable() {
let mut g = TrendGuard::new(0.01);
g.report(&make_report(&[0.05]), 8, 4);
g.report(&make_report(&[0.04]), 8, 4); assert_eq!(
g.report(&make_report(&[0.06]), 8, 4),
ConvergenceAction::Stable
);
}
#[test]
fn trend_below_threshold_is_stable() {
let mut g = TrendGuard::new(0.10);
g.report(&make_report(&[0.01]), 8, 4);
g.report(&make_report(&[0.02]), 8, 4);
assert_eq!(
g.report(&make_report(&[0.03]), 8, 4),
ConvergenceAction::Stable
);
}
#[test]
fn trend_history_capped_at_5() {
let mut g = TrendGuard::new(0.01);
for i in 0..10 {
g.report(&make_report(&[i as f64 * 0.01]), 8, 4);
}
assert_eq!(g.history().len(), 5);
}
#[test]
fn trend_reset_clears_history() {
let mut g = TrendGuard::new(0.01);
for _ in 0..5 {
g.report(&make_report(&[0.05]), 8, 4);
}
g.reset();
assert!(g.history().is_empty());
}
#[test]
fn trend_history_snapshot_roundtrip() {
let mut g = TrendGuard::new(0.01);
for v in [0.01, 0.02, 0.03, 0.04, 0.05] {
g.report(&make_report(&[v]), 8, 4);
}
let snap = g.trend_history().expect("non-empty history");
assert_eq!(snap, vec![0.01, 0.02, 0.03, 0.04, 0.05]);
let restored = TrendGuard::new(0.01).with_history(snap);
assert_eq!(g.history(), restored.history());
}
#[test]
fn trend_with_history_skips_warmup_after_restore() {
let pre_crash = vec![0.02, 0.03, 0.04];
let mut cold = TrendGuard::new(0.01);
assert_eq!(
cold.report(&make_report(&[0.05]), 8, 4),
ConvergenceAction::Stable
);
let mut warm = TrendGuard::new(0.01).with_history(pre_crash);
assert_eq!(
warm.report(&make_report(&[0.05]), 8, 4),
ConvergenceAction::SuppressGrowth
);
}
#[test]
fn trend_with_history_truncates_oversize_input() {
let oversize: Vec<f64> = (0..10).map(|i| i as f64).collect();
let g = TrendGuard::new(0.01).with_history(oversize);
assert_eq!(g.history().len(), 5);
let restored: Vec<f64> = g.history().iter().copied().collect();
assert_eq!(restored, vec![5.0, 6.0, 7.0, 8.0, 9.0]);
}
#[test]
fn trend_history_empty_returns_none() {
let g = TrendGuard::default();
assert!(g.trend_history().is_none());
}
#[test]
fn no_guard_trend_history_is_none() {
let g = NoGuard;
assert!(g.trend_history().is_none());
}
#[test]
fn msf_guard_trend_history_is_none() {
let g = MsfGuard::default();
assert!(g.trend_history().is_none());
}
#[test]
fn lambda_first_event_is_none() {
let mut e = LambdaEstimator::default();
let s = e.observe(0.05, 8, 4);
assert!(s.lambda_raw.is_none());
assert!(s.lambda_ema.is_none());
assert!((s.d_raw - 0.05).abs() < 1e-12);
assert_eq!(s.k_used, 8);
assert_eq!(s.k_max, 4);
}
#[test]
fn lambda_growth_positive_uses_k_max() {
let mut e = LambdaEstimator::default();
e.observe(0.01, 8, 4);
let s = e.observe(0.02, 8, 4);
let expected = std::f64::consts::LN_2 / 4.0;
let got = s.lambda_raw.expect("second event must produce lambda");
assert!((got - expected).abs() < 1e-10);
}
#[test]
fn lambda_decay_negative_uses_k_max() {
let mut e = LambdaEstimator::default();
e.observe(0.04, 8, 4);
let s = e.observe(0.02, 8, 4);
let expected = -std::f64::consts::LN_2 / 4.0;
let got = s.lambda_raw.expect("second event must produce lambda");
assert!((got - expected).abs() < 1e-10);
}
#[test]
fn lambda_k_max_unchanged_by_world_size() {
let mut e_two = LambdaEstimator::default();
e_two.observe(0.01, 8, 4);
let s_two = e_two.observe(0.02, 8, 4);
let mut e_three = LambdaEstimator::default();
e_three.observe(0.01, 12, 4);
let s_three = e_three.observe(0.02, 12, 4);
let l_two = s_two.lambda_raw.unwrap();
let l_three = s_three.lambda_raw.unwrap();
assert!((l_two - l_three).abs() < 1e-12, "{l_two} vs {l_three}");
}
#[test]
fn lambda_noise_floor_resets_estimator() {
let mut e = LambdaEstimator::default();
e.observe(0.01, 8, 4);
let s = e.observe(1e-12, 8, 4);
assert!(s.lambda_raw.is_none());
let s2 = e.observe(0.02, 8, 4);
assert!(s2.lambda_raw.is_none());
let s3 = e.observe(0.04, 8, 4);
let expected = std::f64::consts::LN_2 / 4.0;
assert!((s3.lambda_raw.unwrap() - expected).abs() < 1e-10);
}
#[test]
fn lambda_ema_bias_corrected_at_warmup() {
let mut e = LambdaEstimator::default();
e.observe(1.0, 1, 1);
let s = e.observe(std::f64::consts::E, 1, 1);
let lambda = s.lambda_raw.unwrap();
let ema = s.lambda_ema.unwrap();
assert!((lambda - 1.0).abs() < 1e-10, "lambda={lambda}");
assert!((ema - 1.0).abs() < 1e-10, "ema={ema}");
}
#[test]
fn lambda_ema_smooths_over_constant_input() {
let mut e = LambdaEstimator::default();
let mut prev = 1.0_f64;
for _ in 0..200 {
let next = prev * std::f64::consts::E;
e.observe(next, 1, 1);
prev = next;
}
assert!((e.ema_raw - 1.0).abs() < 1e-3, "raw_ema = {}", e.ema_raw);
}
#[test]
fn msf_default_starts_stable() {
let mut g = MsfGuard::default();
let s = g.report(&make_report(&[0.01]), 8, 4);
assert_eq!(s, ConvergenceAction::Stable);
}
#[test]
fn msf_suppress_fires_after_sustain() {
let mut g = MsfGuard::default()
.with_suppress(1.0e-3, 3)
.with_nudge(f64::INFINITY, 3, 0.5); let mut prev = 1.0e-4;
let mut fired = false;
for _ in 0..20 {
let next = prev * 2.0;
let action = g.report(&make_report(&[next]), 8, 4);
if matches!(action, ConvergenceAction::SuppressGrowth) {
fired = true;
break;
}
prev = next;
}
assert!(fired, "expected SuppressGrowth within 20 events");
}
#[test]
fn msf_nudge_fires_on_sustained_steeper_growth() {
let mut g = MsfGuard::default()
.with_suppress(1.0e-3, 3)
.with_nudge(1.0e-3, 3, 0.5);
let mut prev = 1.0e-4;
let mut nudge_fires = 0;
for _ in 0..20 {
let next = prev * 2.0;
if matches!(
g.report(&make_report(&[next]), 8, 4),
ConvergenceAction::NudgeDown { .. }
) {
nudge_fires += 1;
}
prev = next;
}
assert!(nudge_fires >= 1, "expected NudgeDown to fire at least once");
}
#[test]
fn msf_decaying_lambda_does_not_fire() {
let mut g = MsfGuard::default();
let mut prev = 0.5;
for _ in 0..30 {
let next = prev * 0.5;
let action = g.report(&make_report(&[next]), 8, 4);
assert_eq!(action, ConvergenceAction::Stable);
prev = next;
}
}
#[test]
fn msf_without_nudge_disables_hard_trigger() {
let mut g = MsfGuard::default()
.with_suppress(1.0e-3, 3)
.without_nudge();
let mut prev = 1.0e-4;
let mut nudge_fires = 0;
for _ in 0..20 {
let next = prev * 2.0;
if matches!(
g.report(&make_report(&[next]), 8, 4),
ConvergenceAction::NudgeDown { .. }
) {
nudge_fires += 1;
}
prev = next;
}
assert_eq!(nudge_fires, 0);
}
#[test]
fn msf_telemetry_carries_lambda_after_first_observation() {
let mut g = MsfGuard::default();
g.report(&make_report(&[0.01]), 8, 4);
g.report(&make_report(&[0.02]), 8, 4);
let t = g.telemetry();
assert!(t.iter().any(|(k, _)| *k == "lambda_raw"));
assert!(t.iter().any(|(k, _)| *k == "lambda_ema"));
}
#[test]
fn msf_reset_clears_state() {
let mut g = MsfGuard::default();
g.report(&make_report(&[0.01]), 8, 4);
g.report(&make_report(&[0.02]), 8, 4);
g.reset();
let s = g.report(&make_report(&[0.04]), 8, 4);
assert_eq!(s, ConvergenceAction::Stable);
assert!(g.last_sample().unwrap().lambda_raw.is_none());
}
}