use chrono::{DateTime, Utc};
use rust_decimal::Decimal;
use serde::{Deserialize, Serialize};
use std::collections::VecDeque;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OnlineLearningConfig {
pub learning_rate: f64,
pub adaptive_learning_rate: bool,
pub learning_rate_decay: f64,
pub drift_window_size: usize,
pub drift_threshold: f64,
}
impl Default for OnlineLearningConfig {
fn default() -> Self {
Self {
learning_rate: 0.01,
adaptive_learning_rate: true,
learning_rate_decay: 0.99,
drift_window_size: 100,
drift_threshold: 0.1,
}
}
}
pub trait OnlineLearner {
fn update(&mut self, features: &[f64], target: f64) -> anyhow::Result<()>;
fn batch_update(&mut self, batch: &[(Vec<f64>, f64)]) -> anyhow::Result<()> {
for (features, target) in batch {
self.update(features, *target)?;
}
Ok(())
}
fn learning_rate(&self) -> f64;
fn set_learning_rate(&mut self, rate: f64);
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OnlineLinearRegression {
pub weights: Vec<f64>,
pub bias: f64,
pub learning_rate: f64,
pub update_count: usize,
pub config: OnlineLearningConfig,
}
impl OnlineLinearRegression {
pub fn new(num_features: usize, config: OnlineLearningConfig) -> Self {
Self {
weights: vec![0.0; num_features],
bias: 0.0,
learning_rate: config.learning_rate,
update_count: 0,
config,
}
}
pub fn predict(&self, features: &[f64]) -> anyhow::Result<f64> {
if features.len() != self.weights.len() {
anyhow::bail!("Feature dimension mismatch");
}
let prediction: f64 = features
.iter()
.zip(&self.weights)
.map(|(x, w)| x * w)
.sum::<f64>()
+ self.bias;
Ok(prediction)
}
fn apply_adaptive_learning_rate(&mut self) {
if self.config.adaptive_learning_rate {
self.learning_rate *= self.config.learning_rate_decay;
}
}
}
impl OnlineLearner for OnlineLinearRegression {
fn update(&mut self, features: &[f64], target: f64) -> anyhow::Result<()> {
if features.len() != self.weights.len() {
anyhow::bail!("Feature dimension mismatch");
}
let prediction = self.predict(features)?;
let error = target - prediction;
for (i, &feature) in features.iter().enumerate() {
self.weights[i] += self.learning_rate * error * feature;
}
self.bias += self.learning_rate * error;
self.update_count += 1;
self.apply_adaptive_learning_rate();
Ok(())
}
fn learning_rate(&self) -> f64 {
self.learning_rate
}
fn set_learning_rate(&mut self, rate: f64) {
self.learning_rate = rate;
}
}
#[derive(Debug, Clone)]
pub struct DriftDetector {
error_window: VecDeque<f64>,
window_size: usize,
threshold: f64,
baseline_mean: f64,
baseline_std: f64,
pub drift_count: usize,
}
impl DriftDetector {
pub fn new(window_size: usize, threshold: f64) -> Self {
Self {
error_window: VecDeque::with_capacity(window_size),
window_size,
threshold,
baseline_mean: 0.0,
baseline_std: 0.0,
drift_count: 0,
}
}
pub fn add_error(&mut self, error: f64) -> bool {
self.error_window.push_back(error.abs());
if self.error_window.len() > self.window_size {
self.error_window.pop_front();
}
if self.error_window.len() == self.window_size && self.baseline_mean == 0.0 {
self.update_baseline();
return false;
}
if self.baseline_mean > 0.0 && self.error_window.len() == self.window_size {
let current_mean = self.calculate_mean();
let drift_detected =
(current_mean - self.baseline_mean).abs() / self.baseline_std > self.threshold;
if drift_detected {
self.drift_count += 1;
self.update_baseline();
}
return drift_detected;
}
false
}
fn calculate_mean(&self) -> f64 {
self.error_window.iter().sum::<f64>() / self.error_window.len() as f64
}
fn calculate_std(&self) -> f64 {
let mean = self.calculate_mean();
let variance = self
.error_window
.iter()
.map(|&x| (x - mean).powi(2))
.sum::<f64>()
/ self.error_window.len() as f64;
variance.sqrt()
}
fn update_baseline(&mut self) {
self.baseline_mean = self.calculate_mean();
self.baseline_std = self.calculate_std().max(1e-6); }
pub fn reset(&mut self) {
self.error_window.clear();
self.baseline_mean = 0.0;
self.baseline_std = 0.0;
}
}
#[derive(Debug)]
pub struct AdaptiveOnlineModel<L: OnlineLearner> {
pub learner: L,
pub drift_detector: DriftDetector,
pub retrain_on_drift: bool,
data_buffer: VecDeque<(Vec<f64>, f64)>,
max_buffer_size: usize,
}
impl<L: OnlineLearner> AdaptiveOnlineModel<L> {
pub fn new(learner: L, drift_detector: DriftDetector, max_buffer_size: usize) -> Self {
Self {
learner,
drift_detector,
retrain_on_drift: true,
data_buffer: VecDeque::with_capacity(max_buffer_size),
max_buffer_size,
}
}
pub fn update(&mut self, features: &[f64], target: f64) -> anyhow::Result<DriftStatus> {
if self.data_buffer.len() >= self.max_buffer_size {
self.data_buffer.pop_front();
}
self.data_buffer.push_back((features.to_vec(), target));
let prediction = target;
self.learner.update(features, target)?;
let error = (target - prediction).abs();
let drift_detected = self.drift_detector.add_error(error);
if drift_detected && self.retrain_on_drift {
for (feats, tgt) in &self.data_buffer {
self.learner.update(feats, *tgt)?;
}
Ok(DriftStatus::DriftDetectedAndRetrained)
} else if drift_detected {
Ok(DriftStatus::DriftDetected)
} else {
Ok(DriftStatus::NoDrift)
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DriftStatus {
NoDrift,
DriftDetected,
DriftDetectedAndRetrained,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct IncrementalStats {
pub total_updates: usize,
pub drift_events: usize,
pub current_learning_rate: f64,
pub last_update: DateTime<Utc>,
pub average_error: f64,
}
impl IncrementalStats {
pub fn new() -> Self {
Self {
total_updates: 0,
drift_events: 0,
current_learning_rate: 0.01,
last_update: Utc::now(),
average_error: 0.0,
}
}
pub fn update(&mut self, learning_rate: f64, error: f64, drift_detected: bool) {
self.total_updates += 1;
if drift_detected {
self.drift_events += 1;
}
self.current_learning_rate = learning_rate;
self.last_update = Utc::now();
let alpha = 0.1;
self.average_error = alpha * error + (1.0 - alpha) * self.average_error;
}
}
impl Default for IncrementalStats {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct OnlineMovingAverage {
window: VecDeque<Decimal>,
window_size: usize,
sum: Decimal,
}
impl OnlineMovingAverage {
pub fn new(window_size: usize) -> Self {
Self {
window: VecDeque::with_capacity(window_size),
window_size,
sum: Decimal::ZERO,
}
}
pub fn add(&mut self, value: Decimal) -> Decimal {
if self.window.len() == self.window_size {
if let Some(old) = self.window.pop_front() {
self.sum -= old;
}
}
self.window.push_back(value);
self.sum += value;
self.average()
}
pub fn average(&self) -> Decimal {
if self.window.is_empty() {
Decimal::ZERO
} else {
self.sum / Decimal::from(self.window.len())
}
}
pub fn reset(&mut self) {
self.window.clear();
self.sum = Decimal::ZERO;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_online_linear_regression() {
let config = OnlineLearningConfig::default();
let mut model = OnlineLinearRegression::new(2, config);
for _ in 0..100 {
let x1 = 1.0;
let x2 = 2.0;
let y = 2.0 * x1 + 3.0 * x2;
model.update(&[x1, x2], y).unwrap();
}
let prediction = model.predict(&[1.0, 2.0]).unwrap();
assert!((prediction - 8.0).abs() < 1.0); }
#[test]
fn test_drift_detector() {
let mut detector = DriftDetector::new(10, 2.0);
for _ in 0..20 {
detector.add_error(0.1);
}
for _ in 0..10 {
detector.add_error(1.0);
}
assert!(detector.drift_count > 0);
}
#[test]
fn test_online_moving_average() {
let mut oma = OnlineMovingAverage::new(3);
assert_eq!(oma.add(Decimal::from(10)), Decimal::from(10));
assert_eq!(oma.add(Decimal::from(20)), Decimal::from(15));
assert_eq!(oma.add(Decimal::from(30)), Decimal::from(20));
assert_eq!(oma.add(Decimal::from(40)), Decimal::from(30)); }
#[test]
fn test_adaptive_learning_rate() {
let config = OnlineLearningConfig {
learning_rate: 0.01,
adaptive_learning_rate: true,
learning_rate_decay: 0.9,
drift_window_size: 100,
drift_threshold: 0.1,
};
let mut model = OnlineLinearRegression::new(2, config);
let initial_lr = model.learning_rate;
model.update(&[1.0, 2.0], 5.0).unwrap();
assert!(model.learning_rate < initial_lr);
}
#[test]
fn test_incremental_stats() {
let mut stats = IncrementalStats::new();
stats.update(0.01, 0.5, false);
assert_eq!(stats.total_updates, 1);
assert_eq!(stats.drift_events, 0);
stats.update(0.01, 0.3, true);
assert_eq!(stats.total_updates, 2);
assert_eq!(stats.drift_events, 1);
}
}