feat(ml): add online learning with EWC regularization

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 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-03 04:25:14 +01:00
parent c2e31e2c40
commit 845f8707b9
2 changed files with 980 additions and 0 deletions

View File

@@ -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

View File

@@ -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<f32>,
/// Discrete action taken.
pub action: usize,
/// Scalar reward received.
pub reward: f32,
/// State vector after transition.
pub next_state: Vec<f32>,
/// 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<f64>,
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<String, Tensor>,
/// Optimal parameter snapshot (theta-star).
optimal_params: HashMap<String, Tensor>,
/// 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<String, Tensor> = 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<String> = data.keys().cloned().collect();
let param_tensors: Vec<Tensor> = 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<Tensor, MLError> {
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<Tensor> = 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<Experience>,
config: OnlineLearningConfig,
ewc: EWCRegularizer,
sharpe_monitor: RollingSharpe,
trade_count: usize,
baseline_sharpe: Option<f64>,
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<VarMap, MLError> {
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());
}
}