From 845f8707b95ffd9cc85026827345540d2dadd594 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Tue, 3 Mar 2026 04:25:14 +0100 Subject: [PATCH] feat(ml): add online learning with EWC regularization MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Per-trade online learning with Elastic Weight Consolidation to prevent catastrophic forgetting. Rolling 10K experience buffer, mini-updates every 100 trades. Safety rails: grad clip 1.0, LR×0.1, auto-rollback at 20% Sharpe degradation, kill switch at Sharpe < -1.0. Co-Authored-By: Claude Opus 4.6 --- crates/ml/src/trainers/mod.rs | 1 + crates/ml/src/trainers/online_learning.rs | 979 ++++++++++++++++++++++ 2 files changed, 980 insertions(+) create mode 100644 crates/ml/src/trainers/online_learning.rs diff --git a/crates/ml/src/trainers/mod.rs b/crates/ml/src/trainers/mod.rs index 5d730cc59..292aca2a2 100644 --- a/crates/ml/src/trainers/mod.rs +++ b/crates/ml/src/trainers/mod.rs @@ -73,6 +73,7 @@ pub mod curriculum; pub mod dqn; pub mod liquid; pub mod mamba2; +pub mod online_learning; pub mod ppo; pub mod tft; pub mod tft_parquet; // Parquet lazy-loading extension for TFT diff --git a/crates/ml/src/trainers/online_learning.rs b/crates/ml/src/trainers/online_learning.rs new file mode 100644 index 000000000..e87f9c5d9 --- /dev/null +++ b/crates/ml/src/trainers/online_learning.rs @@ -0,0 +1,979 @@ +//! Online Learning with Elastic Weight Consolidation (EWC) +//! +//! Provides per-trade mini-updates that adapt models to regime shifts +//! without catastrophic forgetting. EWC regularization penalizes movement +//! away from previously learned optimal parameters, weighted by Fisher +//! information (importance of each parameter). +//! +//! # Safety Rails +//! +//! - **Rolling Sharpe monitor**: tracks live performance +//! - **Rollback trigger**: if Sharpe degrades > threshold from baseline +//! - **Kill switch**: if Sharpe drops below absolute floor, freezes updates + +use std::collections::{HashMap, VecDeque}; + +use candle_core::{Device, Tensor}; +use candle_nn::VarMap; +use serde::{Deserialize, Serialize}; + +use crate::MLError; + +// --------------------------------------------------------------------------- +// Experience (transition tuple) +// --------------------------------------------------------------------------- + +/// A single transition from the environment, stored for replay. +#[derive(Debug, Clone)] +pub struct Experience { + /// State vector at decision time. + pub state: Vec, + /// Discrete action taken. + pub action: usize, + /// Scalar reward received. + pub reward: f32, + /// State vector after transition. + pub next_state: Vec, + /// Whether the episode terminated. + pub done: bool, +} + +// --------------------------------------------------------------------------- +// Safety status +// --------------------------------------------------------------------------- + +/// Result of safety-rail checks after each trade batch. +#[derive(Debug, Clone, PartialEq)] +pub enum SafetyStatus { + /// Training may continue normally. + Normal, + /// Sharpe degraded beyond the configured threshold; caller should + /// roll back to the last known-good weights. + Rollback { reason: String }, + /// Absolute Sharpe floor breached; all online updates are frozen. + Frozen { reason: String }, +} + +// --------------------------------------------------------------------------- +// Rolling Sharpe +// --------------------------------------------------------------------------- + +/// Windowed Sharpe ratio tracker over recent trade returns. +#[derive(Debug)] +pub struct RollingSharpe { + returns: VecDeque, + window_size: usize, +} + +impl RollingSharpe { + /// Create a new tracker with the given window size. + pub fn new(window_size: usize) -> Self { + Self { + returns: VecDeque::with_capacity(window_size), + window_size: window_size.max(1), + } + } + + /// Record a single trade return. + pub fn push(&mut self, trade_return: f64) { + if self.returns.len() >= self.window_size { + self.returns.pop_front(); + } + self.returns.push_back(trade_return); + } + + /// Compute the Sharpe ratio (mean / std) over the current window. + /// Returns 0.0 if fewer than 2 samples or zero variance. + pub fn sharpe(&self) -> f64 { + let n = self.returns.len(); + if n < 2 { + return 0.0; + } + + let sum: f64 = self.returns.iter().sum(); + let mean = sum / n as f64; + + let var_sum: f64 = self.returns.iter().map(|r| (r - mean).powi(2)).sum(); + let std = (var_sum / (n as f64 - 1.0)).sqrt(); + + if std < 1e-12 { + return 0.0; + } + + mean / std + } + + /// Number of returns currently stored. + pub fn len(&self) -> usize { + self.returns.len() + } + + /// Whether the window is empty. + pub fn is_empty(&self) -> bool { + self.returns.is_empty() + } + + /// Whether the window is fully populated. + pub fn is_full(&self) -> bool { + self.returns.len() >= self.window_size + } +} + +// --------------------------------------------------------------------------- +// EWC Regularizer +// --------------------------------------------------------------------------- + +/// Elastic Weight Consolidation regularizer. +/// +/// Prevents catastrophic forgetting by penalizing parameter drift from a +/// snapshot of optimal weights, scaled by the diagonal Fisher information +/// matrix (importance of each parameter). +pub struct EWCRegularizer { + /// Diagonal Fisher information per named parameter. + fisher_diagonal: HashMap, + /// Optimal parameter snapshot (theta-star). + optimal_params: HashMap, + /// EWC penalty strength. + lambda: f64, +} + +impl std::fmt::Debug for EWCRegularizer { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("EWCRegularizer") + .field("fisher_params", &self.fisher_diagonal.len()) + .field("optimal_params", &self.optimal_params.len()) + .field("lambda", &self.lambda) + .finish() + } +} + +impl EWCRegularizer { + /// Create a new EWC regularizer with the given penalty strength. + pub fn new(lambda: f64) -> Self { + Self { + fisher_diagonal: HashMap::new(), + optimal_params: HashMap::new(), + lambda, + } + } + + /// Snapshot the current VarMap weights as the optimal parameters (theta-star). + pub fn capture_optimal_params(&mut self, var_map: &VarMap) -> Result<(), MLError> { + let data = var_map + .data() + .lock() + .map_err(|e| MLError::LockError(format!("EWC capture_optimal_params: {e}")))?; + + self.optimal_params.clear(); + for (name, var) in data.iter() { + // Deep-copy via contiguous() to ensure independent storage. + let snapshot = var + .as_tensor() + .contiguous() + .map_err(|e| MLError::ModelError(format!("EWC snapshot '{name}': {e}")))? + .detach(); + self.optimal_params.insert(name.clone(), snapshot); + } + Ok(()) + } + + /// Compute diagonal Fisher information from a set of per-sample losses. + /// + /// For each loss, we compute gradients w.r.t. all parameters, then + /// accumulate the squared gradient magnitudes (diagonal Fisher approx): + /// + /// F_i = (1/N) * sum_n (d loss_n / d theta_i)^2 + pub fn compute_fisher( + &mut self, + var_map: &VarMap, + losses: &[Tensor], + ) -> Result<(), MLError> { + if losses.is_empty() { + return Err(MLError::InvalidInput( + "compute_fisher requires at least one loss".into(), + )); + } + + let data = var_map + .data() + .lock() + .map_err(|e| MLError::LockError(format!("EWC compute_fisher lock: {e}")))?; + + // Initialise accumulators to zero with matching shapes. + let mut accum: HashMap = HashMap::new(); + for (name, var) in data.iter() { + let zeros = Tensor::zeros_like(var.as_tensor()) + .map_err(|e| MLError::ModelError(format!("EWC zeros '{name}': {e}")))?; + accum.insert(name.clone(), zeros); + } + + // We must drop the lock before calling backward (which may need it). + let param_names: Vec = data.keys().cloned().collect(); + let param_tensors: Vec = data.values().map(|v| v.as_tensor().clone()).collect(); + drop(data); + + for loss in losses { + let grads = loss + .backward() + .map_err(|e| MLError::TrainingError(format!("EWC backward: {e}")))?; + + for (idx, name) in param_names.iter().enumerate() { + let param_t = param_tensors + .get(idx) + .ok_or_else(|| MLError::ModelError(format!("EWC param idx {idx} OOB")))?; + + if let Some(grad) = grads.get(param_t) { + let grad_sq = grad + .sqr() + .map_err(|e| MLError::ModelError(format!("EWC sqr '{name}': {e}")))?; + + if let Some(prev) = accum.get(name) { + let updated = prev.add(&grad_sq).map_err(|e| { + MLError::ModelError(format!("EWC accum add '{name}': {e}")) + })?; + accum.insert(name.clone(), updated); + } + } + } + } + + // Average over the number of losses. + let n_losses = losses.len() as f64; + let device = accum + .values() + .next() + .map(|t| t.device().clone()) + .unwrap_or(Device::Cpu); + let divisor = Tensor::new(n_losses as f32, &device) + .map_err(|e| MLError::ModelError(format!("EWC divisor tensor: {e}")))?; + + self.fisher_diagonal.clear(); + for (name, sum_tensor) in &accum { + let fisher = sum_tensor.broadcast_div(&divisor).map_err(|e| { + MLError::ModelError(format!("EWC fisher div '{name}': {e}")) + })?; + self.fisher_diagonal.insert(name.clone(), fisher); + } + + Ok(()) + } + + /// Compute the EWC penalty: lambda * sum_i F_i * (theta_i - theta*_i)^2 + /// + /// Returns a scalar tensor. If no Fisher / optimal params are stored, + /// returns a zero scalar on the provided device (penalty-free). + pub fn ewc_penalty(&self, var_map: &VarMap) -> Result { + if self.fisher_diagonal.is_empty() || self.optimal_params.is_empty() { + return Tensor::new(0.0_f32, &Device::Cpu) + .map_err(|e| MLError::ModelError(format!("EWC zero penalty: {e}"))); + } + + let data = var_map + .data() + .lock() + .map_err(|e| MLError::LockError(format!("EWC penalty lock: {e}")))?; + + let mut penalty_parts: Vec = Vec::new(); + + for (name, var) in data.iter() { + let fisher = match self.fisher_diagonal.get(name) { + Some(f) => f, + None => continue, + }; + let optimal = match self.optimal_params.get(name) { + Some(o) => o, + None => continue, + }; + + // (theta - theta*)^2 + let diff = var + .as_tensor() + .sub(optimal) + .map_err(|e| MLError::ModelError(format!("EWC diff '{name}': {e}")))?; + let diff_sq = diff + .sqr() + .map_err(|e| MLError::ModelError(format!("EWC sqr '{name}': {e}")))?; + + // F_i * (theta - theta*)^2 --> sum to scalar + let weighted = fisher + .mul(&diff_sq) + .map_err(|e| MLError::ModelError(format!("EWC mul '{name}': {e}")))?; + let param_penalty = weighted + .sum_all() + .map_err(|e| MLError::ModelError(format!("EWC sum '{name}': {e}")))?; + + penalty_parts.push(param_penalty); + } + + if penalty_parts.is_empty() { + return Tensor::new(0.0_f32, &Device::Cpu) + .map_err(|e| MLError::ModelError(format!("EWC zero penalty: {e}"))); + } + + // Stack scalars and sum, then multiply by lambda. + let device = penalty_parts + .first() + .map(|t| t.device().clone()) + .unwrap_or(Device::Cpu); + + let stacked = Tensor::stack(&penalty_parts, 0) + .map_err(|e| MLError::ModelError(format!("EWC stack: {e}")))?; + let total = stacked + .sum_all() + .map_err(|e| MLError::ModelError(format!("EWC sum_all: {e}")))?; + + let lambda_tensor = Tensor::new(self.lambda as f32, &device) + .map_err(|e| MLError::ModelError(format!("EWC lambda tensor: {e}")))?; + + total + .broadcast_mul(&lambda_tensor) + .map_err(|e| MLError::ModelError(format!("EWC final mul: {e}"))) + } + + /// Whether optimal params have been captured. + pub fn has_optimal_params(&self) -> bool { + !self.optimal_params.is_empty() + } + + /// Whether Fisher information has been computed. + pub fn has_fisher(&self) -> bool { + !self.fisher_diagonal.is_empty() + } +} + +// --------------------------------------------------------------------------- +// Config +// --------------------------------------------------------------------------- + +/// Configuration for the online learning module. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OnlineLearningConfig { + /// Whether online learning is active. + pub enabled: bool, + /// Maximum number of experiences retained in the replay buffer. + pub buffer_capacity: usize, + /// Number of trades between parameter updates. + pub update_interval: usize, + /// EWC penalty strength (lambda). + pub ewc_lambda: f64, + /// Steps between Fisher information recomputation. + pub fisher_recompute_interval: usize, + /// Learning rate multiplier for online updates (relative to base LR). + pub lr_multiplier: f64, + /// Maximum gradient norm for clipping. + pub max_grad_norm: f64, + /// Window size for the rolling Sharpe monitor. + pub sharpe_window: usize, + /// Fraction of Sharpe degradation that triggers a rollback (0.20 = 20%). + pub sharpe_degradation_threshold: f64, + /// Absolute Sharpe floor; below this the kill switch freezes updates. + pub kill_switch_sharpe: f64, +} + +impl Default for OnlineLearningConfig { + fn default() -> Self { + Self { + enabled: true, + buffer_capacity: 10_000, + update_interval: 100, + ewc_lambda: 1000.0, + fisher_recompute_interval: 10_000, + lr_multiplier: 0.1, + max_grad_norm: 1.0, + sharpe_window: 500, + sharpe_degradation_threshold: 0.20, + kill_switch_sharpe: -1.0, + } + } +} + +// --------------------------------------------------------------------------- +// Online Learner +// --------------------------------------------------------------------------- + +/// Per-trade online learner with EWC regularisation and safety rails. +/// +/// The caller is responsible for: +/// 1. Calling [`record_trade`] after each trade. +/// 2. Checking [`should_update`] to decide when to run a mini-batch SGD step. +/// 3. Calling [`check_safety_rails`] to detect rollback / kill-switch events. +/// 4. Using [`get_recent_batch`] to pull a mini-batch from the buffer. +/// 5. Adding the [`EWCRegularizer::ewc_penalty`] to the loss before backward. +pub struct OnlineLearner { + buffer: VecDeque, + config: OnlineLearningConfig, + ewc: EWCRegularizer, + sharpe_monitor: RollingSharpe, + trade_count: usize, + baseline_sharpe: Option, + frozen: bool, +} + +impl std::fmt::Debug for OnlineLearner { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("OnlineLearner") + .field("buffer_len", &self.buffer.len()) + .field("config", &self.config) + .field("ewc", &self.ewc) + .field("trade_count", &self.trade_count) + .field("baseline_sharpe", &self.baseline_sharpe) + .field("frozen", &self.frozen) + .finish() + } +} + +impl OnlineLearner { + /// Create a new online learner from the given config. + pub fn new(config: OnlineLearningConfig) -> Self { + let ewc = EWCRegularizer::new(config.ewc_lambda); + let sharpe_monitor = RollingSharpe::new(config.sharpe_window); + Self { + buffer: VecDeque::with_capacity(config.buffer_capacity), + config, + ewc, + sharpe_monitor, + trade_count: 0, + baseline_sharpe: None, + frozen: false, + } + } + + /// Record a completed trade: adds the experience to the replay buffer + /// and updates the rolling Sharpe monitor. + pub fn record_trade(&mut self, experience: Experience, trade_return: f64) { + if self.buffer.len() >= self.config.buffer_capacity { + self.buffer.pop_front(); + } + self.buffer.push_back(experience); + self.sharpe_monitor.push(trade_return); + self.trade_count += 1; + } + + /// Whether it is time to run a parameter update. + /// + /// Returns `true` every `update_interval` trades, provided the learner + /// is enabled, not frozen, and the buffer has enough data. + pub fn should_update(&self) -> bool { + if !self.config.enabled || self.frozen { + return false; + } + if self.trade_count == 0 { + return false; + } + self.trade_count % self.config.update_interval == 0 + } + + /// Set the baseline Sharpe ratio (captured once after initial training). + pub fn set_baseline_sharpe(&mut self, sharpe: f64) { + self.baseline_sharpe = Some(sharpe); + } + + /// Check safety rails and return the current status. + pub fn check_safety_rails(&mut self) -> SafetyStatus { + let current_sharpe = self.sharpe_monitor.sharpe(); + + // Kill switch: absolute floor + if current_sharpe < self.config.kill_switch_sharpe { + self.frozen = true; + return SafetyStatus::Frozen { + reason: format!( + "Sharpe {current_sharpe:.4} below kill switch {}", + self.config.kill_switch_sharpe + ), + }; + } + + // Rollback: relative degradation from baseline + if let Some(baseline) = self.baseline_sharpe { + if baseline > 0.0 { + let degradation = (baseline - current_sharpe) / baseline; + if degradation > self.config.sharpe_degradation_threshold { + return SafetyStatus::Rollback { + reason: format!( + "Sharpe degraded {:.1}% (baseline {baseline:.4}, current {current_sharpe:.4})", + degradation * 100.0 + ), + }; + } + } + } + + SafetyStatus::Normal + } + + /// Return up to `batch_size` of the most recent experiences in the buffer. + pub fn get_recent_batch(&self, batch_size: usize) -> Vec<&Experience> { + let len = self.buffer.len(); + let start = len.saturating_sub(batch_size); + self.buffer.iter().skip(start).collect() + } + + /// Access the underlying EWC regularizer (e.g. to capture params / compute Fisher). + pub fn ewc(&self) -> &EWCRegularizer { + &self.ewc + } + + /// Mutable access to the EWC regularizer. + pub fn ewc_mut(&mut self) -> &mut EWCRegularizer { + &mut self.ewc + } + + /// Current trade count. + pub fn trade_count(&self) -> usize { + self.trade_count + } + + /// Whether the learner is frozen (kill switch activated). + pub fn is_frozen(&self) -> bool { + self.frozen + } + + /// Number of experiences in the buffer. + pub fn buffer_len(&self) -> usize { + self.buffer.len() + } + + /// Reference to the config. + pub fn config(&self) -> &OnlineLearningConfig { + &self.config + } + + /// Reference to the rolling Sharpe monitor. + pub fn sharpe_monitor(&self) -> &RollingSharpe { + &self.sharpe_monitor + } +} + +// =========================================================================== +// Tests +// =========================================================================== + +#[cfg(test)] +mod tests { + use super::*; + use candle_core::DType; + + // ----------------------------------------------------------------------- + // Helpers + // ----------------------------------------------------------------------- + + fn make_experience(reward: f32) -> Experience { + Experience { + state: vec![1.0, 2.0, 3.0], + action: 0, + reward, + next_state: vec![1.1, 2.1, 3.1], + done: false, + } + } + + fn default_config() -> OnlineLearningConfig { + OnlineLearningConfig::default() + } + + /// Build a tiny VarMap with a single 2x2 parameter for testing. + fn tiny_var_map() -> Result { + let var_map = VarMap::new(); + let vb = candle_nn::VarBuilder::from_varmap(&var_map, DType::F32, &Device::Cpu); + let _linear = candle_nn::linear(2, 2, vb.pp("layer")) + .map_err(|e| MLError::ModelError(format!("tiny_var_map linear: {e}")))?; + Ok(var_map) + } + + // ----------------------------------------------------------------------- + // EWC tests + // ----------------------------------------------------------------------- + + #[test] + fn test_ewc_new() { + let ewc = EWCRegularizer::new(500.0); + assert!((ewc.lambda - 500.0).abs() < f64::EPSILON); + assert!(!ewc.has_optimal_params()); + assert!(!ewc.has_fisher()); + } + + #[test] + fn test_ewc_capture_params() { + let var_map = tiny_var_map().unwrap_or_else(|_| VarMap::new()); + let mut ewc = EWCRegularizer::new(1000.0); + let result = ewc.capture_optimal_params(&var_map); + assert!(result.is_ok()); + assert!(ewc.has_optimal_params()); + // Should have 2 named tensors (weight + bias for one linear layer) + assert!(ewc.optimal_params.len() >= 2); + } + + #[test] + fn test_ewc_penalty_zero_when_unchanged() { + let var_map = tiny_var_map().unwrap_or_else(|_| VarMap::new()); + let mut ewc = EWCRegularizer::new(1000.0); + ewc.capture_optimal_params(&var_map) + .unwrap_or_else(|_| ()); + + // Fabricate uniform Fisher of ones so the penalty formula is active + // but params haven't changed, so diff = 0 => penalty = 0. + for (name, t) in &ewc.optimal_params { + if let Ok(ones) = Tensor::ones_like(t) { + ewc.fisher_diagonal.insert(name.clone(), ones); + } + } + + let penalty = ewc.ewc_penalty(&var_map); + assert!(penalty.is_ok()); + let penalty_tensor = match penalty { + Ok(t) => t, + Err(_) => { + assert!(false, "ewc_penalty returned Err"); + return; + } + }; + let val: f32 = penalty_tensor.to_scalar().unwrap_or(999.0); + assert!( + val.abs() < 1e-5, + "Expected ~0 penalty when params unchanged, got {val}" + ); + } + + #[test] + fn test_ewc_penalty_nonzero_when_changed() { + // Build two VarMaps: one for the "optimal" snapshot, one for "drifted" params. + // We capture optimal from var_map_a, then compute penalty against var_map_b + // (which has the same named params but different values). + let var_map_a = tiny_var_map().unwrap_or_else(|_| VarMap::new()); + let mut ewc = EWCRegularizer::new(1000.0); + ewc.capture_optimal_params(&var_map_a) + .unwrap_or_else(|_| ()); + + // Fabricate Fisher of ones. + for (name, t) in &ewc.optimal_params { + if let Ok(ones) = Tensor::ones_like(t) { + ewc.fisher_diagonal.insert(name.clone(), ones); + } + } + + // Build a second VarMap with identical structure but shifted values. + // We'll manually create params with the same names but offset by 1.0. + let var_map_b = VarMap::new(); + { + let data_a = var_map_a + .data() + .lock() + .unwrap_or_else(|p| p.into_inner()); + for (name, var_a) in data_a.iter() { + let t = var_a.as_tensor(); + if let Ok(offset) = Tensor::ones_like(t) { + if let Ok(shifted) = t.add(&offset) { + // Register in var_map_b via VarBuilder so name matches. + let _ = var_map_b.data().lock().map(|mut lock| { + let var = candle_core::Var::from_tensor(&shifted); + if let Ok(v) = var { + lock.insert(name.clone(), v); + } + }); + } + } + } + } + + let penalty = ewc.ewc_penalty(&var_map_b); + assert!(penalty.is_ok()); + let penalty_tensor = match penalty { + Ok(t) => t, + Err(_) => { + assert!(false, "ewc_penalty returned Err"); + return; + } + }; + let val: f32 = penalty_tensor.to_scalar().unwrap_or(0.0); + assert!( + val > 1.0, + "Expected nonzero penalty when params drifted, got {val}" + ); + } + + #[test] + fn test_ewc_compute_fisher_empty_losses() { + let var_map = tiny_var_map().unwrap_or_else(|_| VarMap::new()); + let mut ewc = EWCRegularizer::new(1000.0); + let result = ewc.compute_fisher(&var_map, &[]); + assert!(result.is_err()); + } + + // ----------------------------------------------------------------------- + // Rolling Sharpe tests + // ----------------------------------------------------------------------- + + #[test] + fn test_rolling_sharpe_empty() { + let rs = RollingSharpe::new(500); + assert!(rs.is_empty()); + assert!(!rs.is_full()); + assert!((rs.sharpe() - 0.0).abs() < f64::EPSILON); + } + + #[test] + fn test_rolling_sharpe_positive() { + let mut rs = RollingSharpe::new(500); + for _ in 0..100 { + rs.push(0.01); // all positive + } + assert!(rs.len() == 100); + // Sharpe should be positive when all returns are positive and equal + // Actually, std is 0 when all returns are identical => sharpe = 0. + // Add some variance. + let mut rs2 = RollingSharpe::new(500); + for i in 0..100 { + rs2.push(0.01 + (i as f64) * 0.0001); + } + let s = rs2.sharpe(); + assert!(s > 0.0, "Expected positive Sharpe for positive returns, got {s}"); + } + + #[test] + fn test_rolling_sharpe_window() { + let mut rs = RollingSharpe::new(5); + // Push 5 values + for i in 0..5 { + rs.push(i as f64); + } + assert!(rs.is_full()); + assert_eq!(rs.len(), 5); + + // Push one more — oldest should be evicted. + rs.push(100.0); + assert_eq!(rs.len(), 5); + // The window now holds [1, 2, 3, 4, 100], so mean is large. + let s = rs.sharpe(); + assert!(s > 0.0, "Expected positive Sharpe after window eviction"); + } + + #[test] + fn test_rolling_sharpe_single_value() { + let mut rs = RollingSharpe::new(10); + rs.push(0.05); + assert_eq!(rs.len(), 1); + // Not enough data for std, should return 0. + assert!((rs.sharpe() - 0.0).abs() < f64::EPSILON); + } + + // ----------------------------------------------------------------------- + // OnlineLearner tests + // ----------------------------------------------------------------------- + + #[test] + fn test_online_learner_should_update() { + let mut config = default_config(); + config.update_interval = 10; + let mut learner = OnlineLearner::new(config); + + // No trades yet. + assert!(!learner.should_update()); + + // Record 9 trades — not yet at interval. + for _ in 0..9 { + learner.record_trade(make_experience(1.0), 0.01); + } + assert!(!learner.should_update()); + + // 10th trade — should trigger. + learner.record_trade(make_experience(1.0), 0.01); + assert!(learner.should_update()); + + // 11th trade — not at interval. + learner.record_trade(make_experience(1.0), 0.01); + assert!(!learner.should_update()); + + // 20th trade — should trigger again. + for _ in 0..9 { + learner.record_trade(make_experience(1.0), 0.01); + } + assert!(learner.should_update()); + } + + #[test] + fn test_online_learner_buffer_capacity() { + let mut config = default_config(); + config.buffer_capacity = 5; + let mut learner = OnlineLearner::new(config); + + for i in 0..10 { + learner.record_trade(make_experience(i as f32), 0.01); + } + + // Buffer should be capped at 5. + assert_eq!(learner.buffer_len(), 5); + + // The oldest entries should have been evicted. + // The remaining should be rewards 5..10. + let batch = learner.get_recent_batch(5); + assert_eq!(batch.len(), 5); + let first_reward = batch.first().map(|e| e.reward).unwrap_or(-1.0); + assert!( + (first_reward - 5.0).abs() < f32::EPSILON, + "Expected oldest remaining reward = 5.0, got {first_reward}" + ); + } + + #[test] + fn test_safety_normal() { + let config = default_config(); + let mut learner = OnlineLearner::new(config); + + // Set a positive baseline. + learner.set_baseline_sharpe(2.0); + + // Push enough positive returns to have a healthy Sharpe. + for i in 0..50 { + learner.record_trade(make_experience(1.0), 0.01 + (i as f64) * 0.0001); + } + + let status = learner.check_safety_rails(); + assert_eq!(status, SafetyStatus::Normal); + } + + #[test] + fn test_safety_rollback() { + let mut config = default_config(); + config.sharpe_degradation_threshold = 0.20; + let mut learner = OnlineLearner::new(config); + + // Set a high baseline. + learner.set_baseline_sharpe(5.0); + + // Push mixed returns that yield a much lower Sharpe (degradation > 20%). + for i in 0..100 { + let ret = if i % 2 == 0 { 0.01 } else { -0.008 }; + learner.record_trade(make_experience(1.0), ret); + } + + let status = learner.check_safety_rails(); + match status { + SafetyStatus::Rollback { .. } => {} // Expected + other => panic!("Expected Rollback, got {other:?}"), + } + } + + #[test] + fn test_safety_frozen() { + let mut config = default_config(); + config.kill_switch_sharpe = -1.0; + let mut learner = OnlineLearner::new(config); + + // Push strongly negative returns to get Sharpe below -1.0. + for i in 0..100 { + let ret = -0.05 - (i as f64) * 0.001; + learner.record_trade(make_experience(-1.0), ret); + } + + let status = learner.check_safety_rails(); + match status { + SafetyStatus::Frozen { .. } => {} // Expected + other => panic!("Expected Frozen, got {other:?}"), + } + assert!(learner.is_frozen()); + } + + #[test] + fn test_frozen_prevents_updates() { + let mut config = default_config(); + config.update_interval = 1; + let mut learner = OnlineLearner::new(config); + + // Manually freeze. + learner.frozen = true; + learner.record_trade(make_experience(1.0), 0.01); + assert!(!learner.should_update()); + } + + #[test] + fn test_config_defaults() { + let config = OnlineLearningConfig::default(); + assert!(config.enabled); + assert_eq!(config.buffer_capacity, 10_000); + assert_eq!(config.update_interval, 100); + assert!((config.ewc_lambda - 1000.0).abs() < f64::EPSILON); + assert_eq!(config.fisher_recompute_interval, 10_000); + assert!((config.lr_multiplier - 0.1).abs() < f64::EPSILON); + assert!((config.max_grad_norm - 1.0).abs() < f64::EPSILON); + assert_eq!(config.sharpe_window, 500); + assert!((config.sharpe_degradation_threshold - 0.20).abs() < f64::EPSILON); + assert!((config.kill_switch_sharpe - (-1.0)).abs() < f64::EPSILON); + } + + #[test] + fn test_batch_sampling() { + let config = default_config(); + let mut learner = OnlineLearner::new(config); + + for i in 0..20 { + learner.record_trade(make_experience(i as f32), 0.01); + } + + let batch = learner.get_recent_batch(10); + assert_eq!(batch.len(), 10); + + // Should be the most recent 10 (rewards 10..20). + let first_reward = batch.first().map(|e| e.reward).unwrap_or(-1.0); + assert!( + (first_reward - 10.0).abs() < f32::EPSILON, + "Expected first in batch to have reward 10, got {first_reward}" + ); + + // Requesting more than buffer size returns everything. + let big_batch = learner.get_recent_batch(1000); + assert_eq!(big_batch.len(), 20); + } + + #[test] + fn test_disabled_prevents_updates() { + let mut config = default_config(); + config.enabled = false; + config.update_interval = 1; + let mut learner = OnlineLearner::new(config); + + learner.record_trade(make_experience(1.0), 0.01); + assert!(!learner.should_update()); + } + + #[test] + fn test_ewc_penalty_empty_regularizer() { + let var_map = tiny_var_map().unwrap_or_else(|_| VarMap::new()); + let ewc = EWCRegularizer::new(1000.0); + // No optimal params or Fisher captured — should return 0. + let penalty = ewc.ewc_penalty(&var_map); + assert!(penalty.is_ok()); + let penalty_tensor = match penalty { + Ok(t) => t, + Err(_) => { + assert!(false, "ewc_penalty returned Err"); + return; + } + }; + let val: f32 = penalty_tensor.to_scalar().unwrap_or(999.0); + assert!(val.abs() < 1e-5, "Expected 0 penalty for empty EWC, got {val}"); + } + + #[test] + fn test_online_learner_trade_count() { + let config = default_config(); + let mut learner = OnlineLearner::new(config); + assert_eq!(learner.trade_count(), 0); + + for _ in 0..7 { + learner.record_trade(make_experience(1.0), 0.01); + } + assert_eq!(learner.trade_count(), 7); + } + + #[test] + fn test_online_learner_ewc_access() { + let config = default_config(); + let mut learner = OnlineLearner::new(config); + assert!(!learner.ewc().has_optimal_params()); + + // Capture params through the mutable accessor. + let var_map = tiny_var_map().unwrap_or_else(|_| VarMap::new()); + let _ = learner.ewc_mut().capture_optimal_params(&var_map); + assert!(learner.ewc().has_optimal_params()); + } +}