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:
@@ -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
|
||||
|
||||
979
crates/ml/src/trainers/online_learning.rs
Normal file
979
crates/ml/src/trainers/online_learning.rs
Normal 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());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user