use std::collections::VecDeque;
use crate::distributed::Phase;
use crate::distributed::ddp_run::convergence::ConvergenceAction;
#[derive(Debug, Clone, Copy)]
pub struct LrEventMetaConfig {
pub lr_window_cap: usize,
pub anchor_window_cap: usize,
pub guard_window_cap: usize,
pub sharp_drop_threshold: f64,
pub effective_factor_min: f64,
pub effective_factor_max: f64,
}
impl Default for LrEventMetaConfig {
fn default() -> Self {
Self {
lr_window_cap: 3,
anchor_window_cap: 5,
guard_window_cap: 5,
sharp_drop_threshold: 0.3,
effective_factor_min: 0.1,
effective_factor_max: 0.95,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[non_exhaustive]
pub enum MetaAction {
Noop,
NudgeDown { factor: f64 },
}
pub struct LrEventMeta {
config: LrEventMetaConfig,
lr_window: VecDeque<f64>,
anchor_window: VecDeque<usize>,
guard_window: VecDeque<ConvergenceAction>,
}
impl LrEventMeta {
pub fn new(mut config: LrEventMetaConfig) -> Self {
let min_cap = sustain_k_for(Phase::Warmup);
if config.guard_window_cap < min_cap {
crate::verbose!(
" ddp: LrEventMetaConfig.guard_window_cap {} < required {}; \
clamping",
config.guard_window_cap,
min_cap,
);
config.guard_window_cap = min_cap;
}
let lr_window = VecDeque::with_capacity(config.lr_window_cap);
let anchor_window = VecDeque::with_capacity(config.anchor_window_cap);
let guard_window = VecDeque::with_capacity(config.guard_window_cap);
Self { config, lr_window, anchor_window, guard_window }
}
pub fn with_default_config() -> Self {
Self::new(LrEventMetaConfig::default())
}
pub fn observe(
&mut self,
lr: f64,
anchor: usize,
verdict: ConvergenceAction,
phase: Phase,
) -> MetaAction {
push_capped(&mut self.lr_window, lr, self.config.lr_window_cap);
push_capped(&mut self.anchor_window, anchor, self.config.anchor_window_cap);
push_capped(&mut self.guard_window, verdict, self.config.guard_window_cap);
if phase < Phase::Warmup {
return MetaAction::Noop;
}
let trend = self.anchor_trend();
let lr_cliff_factor = if self.sharp_drop_detected() {
Some(self.effective_factor(LR_CLIFF_BASE_FACTOR, trend))
} else {
None
};
let conv_factor = if self.convergence_pattern_fires(phase) {
Some(self.effective_factor(base_factor_for(phase), trend))
} else {
None
};
match (lr_cliff_factor, conv_factor) {
(None, None) => MetaAction::Noop,
(Some(f), None) | (None, Some(f)) => MetaAction::NudgeDown { factor: f },
(Some(f1), Some(f2)) => {
let combined = (f1 * f2)
.clamp(self.config.effective_factor_min, self.config.effective_factor_max);
MetaAction::NudgeDown { factor: combined }
}
}
}
pub fn anchor_trend(&self) -> f64 {
if self.anchor_window.len() < 2 {
return 0.0;
}
let n = self.anchor_window.len() as f64;
let mean: f64 = self.anchor_window.iter().map(|&a| a as f64).sum::<f64>() / n;
if mean <= 0.0 {
return 0.0;
}
let current = *self.anchor_window.back().unwrap() as f64;
(current - mean) / mean
}
pub fn is_settled(&self) -> bool {
self.anchor_trend().abs() < SETTLED_EPSILON
}
fn sharp_drop_detected(&self) -> bool {
if self.lr_window.len() < 2 {
return false;
}
let mut iter = self.lr_window.iter().rev();
let lr_now = *iter.next().unwrap();
let lr_prev = *iter.next().unwrap();
if lr_prev <= 0.0 || lr_now >= lr_prev {
return false;
}
let drop_ratio = (lr_prev - lr_now) / lr_prev;
drop_ratio > self.config.sharp_drop_threshold
}
fn convergence_pattern_fires(&self, phase: Phase) -> bool {
let k = sustain_k_for(phase);
if k == 0 || self.guard_window.len() < k {
return false;
}
self.guard_window
.iter()
.rev()
.take(k)
.all(|v| {
matches!(
v,
ConvergenceAction::NudgeDown { .. } | ConvergenceAction::SuppressGrowth,
)
})
}
fn effective_factor(&self, base: f64, trend: f64) -> f64 {
let raw = base - (1.0 - base) * trend;
raw.clamp(self.config.effective_factor_min, self.config.effective_factor_max)
}
pub fn lr_window(&self) -> &VecDeque<f64> {
&self.lr_window
}
pub fn anchor_window(&self) -> &VecDeque<usize> {
&self.anchor_window
}
pub fn guard_window(&self) -> &VecDeque<ConvergenceAction> {
&self.guard_window
}
pub fn config(&self) -> &LrEventMetaConfig {
&self.config
}
}
const LR_CLIFF_BASE_FACTOR: f64 = 0.5;
const SETTLED_EPSILON: f64 = 0.10;
fn base_factor_for(phase: Phase) -> f64 {
match phase {
Phase::Probe | Phase::Warmup => 0.3,
Phase::Stable => 0.5,
Phase::Mature => 0.7,
}
}
fn sustain_k_for(phase: Phase) -> usize {
match phase {
Phase::Probe | Phase::Warmup => 5,
Phase::Stable => 3,
Phase::Mature => 2,
}
}
fn push_capped<T>(buf: &mut VecDeque<T>, val: T, cap: usize) {
if cap == 0 {
return;
}
while buf.len() >= cap {
buf.pop_front();
}
buf.push_back(val);
}
#[cfg(test)]
mod tests {
use super::*;
fn run_sequence(
meta: &mut LrEventMeta,
phase: Phase,
seq: &[(f64, usize, ConvergenceAction)],
) -> MetaAction {
let mut last = MetaAction::Noop;
for &(lr, anchor, v) in seq {
last = meta.observe(lr, anchor, v, phase);
}
last
}
#[test]
fn single_observation_emits_noop() {
let mut meta = LrEventMeta::with_default_config();
let action = meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Stable);
assert_eq!(action, MetaAction::Noop);
}
#[test]
fn observe_records_into_windows_with_cap() {
let cfg = LrEventMetaConfig {
lr_window_cap: 3,
anchor_window_cap: 3,
guard_window_cap: 3,
..LrEventMetaConfig::default()
};
let mut meta = LrEventMeta::new(cfg);
for i in 0..5 {
meta.observe(
0.1 * (i as f64),
100 + i,
ConvergenceAction::Stable,
Phase::Stable,
);
}
assert_eq!(meta.lr_window().len(), 3);
assert_eq!(meta.anchor_window().len(), 3);
assert_eq!(meta.guard_window().len(), 5);
assert!((meta.lr_window()[0] - 0.2).abs() < 1e-9);
assert_eq!(*meta.anchor_window().back().unwrap(), 104);
}
#[test]
fn push_capped_zero_capacity_is_noop() {
let mut buf: VecDeque<i32> = VecDeque::new();
push_capped(&mut buf, 1, 0);
assert!(buf.is_empty());
}
#[test]
fn probe_phase_silences_both_watchers() {
let mut meta = LrEventMeta::with_default_config();
meta.observe(0.1, 100, ConvergenceAction::NudgeDown { factor: 0.5 }, Phase::Probe);
let action = meta.observe(
0.01,
100,
ConvergenceAction::NudgeDown { factor: 0.5 },
Phase::Probe,
);
assert_eq!(action, MetaAction::Noop, "Probe must never act");
}
#[test]
fn lr_cliff_fires_on_multistep_drop() {
let mut meta = LrEventMeta::with_default_config();
meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Stable);
let action = meta.observe(0.01, 100, ConvergenceAction::Stable, Phase::Stable);
match action {
MetaAction::NudgeDown { factor } => {
assert!((factor - 0.5).abs() < 1e-9, "expected 0.5, got {factor}");
}
MetaAction::Noop => panic!("expected NudgeDown on 90% drop"),
}
}
#[test]
fn lr_cliff_silent_on_smooth_decay() {
let mut meta = LrEventMeta::with_default_config();
let mut lr = 0.1;
let mut last = MetaAction::Noop;
for _ in 0..10 {
last = meta.observe(lr, 100, ConvergenceAction::Stable, Phase::Stable);
lr *= 0.98; }
assert_eq!(last, MetaAction::Noop, "smooth decay must not fire LR cliff");
}
#[test]
fn lr_cliff_silent_on_lr_rise() {
let mut meta = LrEventMeta::with_default_config();
meta.observe(0.001, 100, ConvergenceAction::Stable, Phase::Stable);
let action = meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Stable);
assert_eq!(action, MetaAction::Noop, "rising LR must not fire LR cliff");
}
#[test]
fn lr_cliff_silent_with_too_short_window() {
let mut meta = LrEventMeta::with_default_config();
let action = meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Stable);
assert_eq!(action, MetaAction::Noop);
}
#[test]
fn lr_cliff_threshold_at_boundary() {
let cfg = LrEventMetaConfig {
sharp_drop_threshold: 0.3,
..LrEventMetaConfig::default()
};
let mut meta = LrEventMeta::new(cfg);
meta.observe(1.0, 100, ConvergenceAction::Stable, Phase::Stable);
let action = meta.observe(0.75, 100, ConvergenceAction::Stable, Phase::Stable);
assert_eq!(action, MetaAction::Noop, "25% drop is under threshold");
let mut meta = LrEventMeta::new(cfg);
meta.observe(1.0, 100, ConvergenceAction::Stable, Phase::Stable);
let action = meta.observe(0.65, 100, ConvergenceAction::Stable, Phase::Stable);
assert!(matches!(action, MetaAction::NudgeDown { .. }), "35% drop over threshold");
}
#[test]
fn convergence_fires_after_k_consecutive_at_stable() {
let mut meta = LrEventMeta::with_default_config();
let nudge = ConvergenceAction::NudgeDown { factor: 0.5 };
run_sequence(
&mut meta,
Phase::Stable,
&[(0.1, 100, nudge), (0.1, 100, nudge)],
);
let action = meta.observe(0.1, 100, nudge, Phase::Stable);
match action {
MetaAction::NudgeDown { factor } => {
assert!((factor - 0.5).abs() < 1e-9, "expected 0.5, got {factor}");
}
MetaAction::Noop => panic!("3 consecutive NudgeDown at Stable should fire"),
}
}
#[test]
fn convergence_silent_on_isolated_nudge() {
let mut meta = LrEventMeta::with_default_config();
let nudge = ConvergenceAction::NudgeDown { factor: 0.5 };
run_sequence(
&mut meta,
Phase::Stable,
&[
(0.1, 100, ConvergenceAction::Stable),
(0.1, 100, nudge),
(0.1, 100, ConvergenceAction::Stable),
],
);
let action = meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Stable);
assert_eq!(action, MetaAction::Noop, "isolated NudgeDown must not fire");
}
#[test]
fn convergence_silent_on_two_at_stable_needs_three() {
let mut meta = LrEventMeta::with_default_config();
let nudge = ConvergenceAction::NudgeDown { factor: 0.5 };
meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Stable);
meta.observe(0.1, 100, nudge, Phase::Stable);
let action = meta.observe(0.1, 100, nudge, Phase::Stable);
assert_eq!(action, MetaAction::Noop, "Stable phase needs K=3, only 2 sustained");
}
#[test]
fn convergence_at_mature_needs_only_two() {
let mut meta = LrEventMeta::with_default_config();
let sg = ConvergenceAction::SuppressGrowth;
meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Mature);
meta.observe(0.1, 100, sg, Phase::Mature);
let action = meta.observe(0.1, 100, sg, Phase::Mature);
match action {
MetaAction::NudgeDown { factor } => {
assert!((factor - 0.7).abs() < 1e-9);
}
MetaAction::Noop => panic!("Mature K=2 with 2 sustained must fire"),
}
}
#[test]
fn convergence_mixed_sustained_kinds_fire() {
let mut meta = LrEventMeta::with_default_config();
let nudge = ConvergenceAction::NudgeDown { factor: 0.5 };
let sg = ConvergenceAction::SuppressGrowth;
meta.observe(0.1, 100, sg, Phase::Stable);
meta.observe(0.1, 100, nudge, Phase::Stable);
let action = meta.observe(0.1, 100, sg, Phase::Stable);
assert!(matches!(action, MetaAction::NudgeDown { .. }));
}
#[test]
fn anchor_trend_rising() {
let mut meta = LrEventMeta::with_default_config();
for &a in &[100, 110, 120, 130, 140] {
meta.observe(0.1, a, ConvergenceAction::Stable, Phase::Stable);
}
let trend = meta.anchor_trend();
assert!(trend > 0.0, "rising anchor must have positive trend, got {trend}");
}
#[test]
fn anchor_trend_falling() {
let mut meta = LrEventMeta::with_default_config();
for &a in &[100, 80, 60, 40, 20] {
meta.observe(0.1, a, ConvergenceAction::Stable, Phase::Stable);
}
let trend = meta.anchor_trend();
assert!(trend < 0.0, "falling anchor must have negative trend, got {trend}");
}
#[test]
fn anchor_trend_stable_near_zero() {
let mut meta = LrEventMeta::with_default_config();
for _ in 0..5 {
meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Stable);
}
let trend = meta.anchor_trend();
assert!(trend.abs() < 1e-9, "constant anchor trend must be 0, got {trend}");
}
#[test]
fn anchor_trend_short_window_is_zero() {
let mut meta = LrEventMeta::with_default_config();
meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Stable);
assert_eq!(meta.anchor_trend(), 0.0);
}
#[test]
fn is_settled_within_epsilon() {
let mut meta = LrEventMeta::with_default_config();
for _ in 0..5 {
meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Stable);
}
assert!(meta.is_settled());
let mut meta = LrEventMeta::with_default_config();
for &a in &[100, 50, 150, 50, 150] {
meta.observe(0.1, a, ConvergenceAction::Stable, Phase::Stable);
}
assert!(!meta.is_settled() || meta.anchor_trend().abs() < SETTLED_EPSILON);
}
#[test]
fn effective_factor_symmetric_around_zero_trend() {
let meta = LrEventMeta::with_default_config();
let base = 0.5;
assert!((meta.effective_factor(base, 0.0) - 0.5).abs() < 1e-9);
assert!((meta.effective_factor(base, -0.5) - 0.75).abs() < 1e-9);
assert!((meta.effective_factor(base, 0.5) - 0.25).abs() < 1e-9);
}
#[test]
fn effective_factor_clamps_at_extremes() {
let cfg = LrEventMetaConfig {
effective_factor_min: 0.1,
effective_factor_max: 0.95,
..LrEventMetaConfig::default()
};
let meta = LrEventMeta::new(cfg);
let f = meta.effective_factor(0.5, 5.0);
assert!((f - 0.1).abs() < 1e-9, "expected clamp to 0.1, got {f}");
let f = meta.effective_factor(0.5, -5.0);
assert!((f - 0.95).abs() < 1e-9, "expected clamp to 0.95, got {f}");
}
#[test]
fn both_watchers_compound_factor_at_mature() {
let mut meta = LrEventMeta::with_default_config();
let nudge = ConvergenceAction::NudgeDown { factor: 0.5 };
meta.observe(0.1, 100, nudge, Phase::Mature);
meta.observe(0.1, 100, nudge, Phase::Mature);
let action = meta.observe(0.01, 100, nudge, Phase::Mature);
match action {
MetaAction::NudgeDown { factor } => {
assert!((factor - 0.35).abs() < 1e-9, "expected compound 0.35, got {factor}");
}
MetaAction::Noop => panic!("both watchers should fire and compound"),
}
}
#[test]
fn per_phase_base_factors_match_design() {
assert!((base_factor_for(Phase::Warmup) - 0.3).abs() < 1e-9);
assert!((base_factor_for(Phase::Stable) - 0.5).abs() < 1e-9);
assert!((base_factor_for(Phase::Mature) - 0.7).abs() < 1e-9);
}
#[test]
fn per_phase_sustain_k_match_design() {
assert_eq!(sustain_k_for(Phase::Warmup), 5);
assert_eq!(sustain_k_for(Phase::Stable), 3);
assert_eq!(sustain_k_for(Phase::Mature), 2);
}
#[test]
fn trend_dampens_subsequent_nudge_after_cliff() {
let mut meta = LrEventMeta::with_default_config();
for _ in 0..4 {
meta.observe(0.1, 100, ConvergenceAction::Stable, Phase::Stable);
}
let action = meta.observe(0.01, 50, ConvergenceAction::Stable, Phase::Stable);
match action {
MetaAction::NudgeDown { factor } => {
assert!(factor > 0.5, "falling trend must dampen factor, got {factor}");
assert!(factor < 0.95, "factor should not be fully clamped, got {factor}");
}
MetaAction::Noop => panic!("LR cliff still fires under dampening"),
}
}
#[test]
fn trend_amplifies_nudge_during_relax_up_climb() {
let mut meta = LrEventMeta::with_default_config();
for &a in &[100, 110, 120, 130, 140] {
meta.observe(0.1, a, ConvergenceAction::Stable, Phase::Stable);
}
let action = meta.observe(0.01, 140, ConvergenceAction::Stable, Phase::Stable);
match action {
MetaAction::NudgeDown { factor } => {
assert!(factor < 0.5, "rising trend must amplify nudge, got {factor}");
assert!(factor > 0.1, "factor should not be fully clamped, got {factor}");
}
MetaAction::Noop => panic!("LR cliff fires regardless of trend"),
}
}
}