use scirs2_core::ndarray::{Array1, Array2, ArrayView2};
use sklears_core::{
error::{Result as SklResult, SklearsError},
traits::{Estimator, Fit, Predict, Untrained},
types::Float,
};
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct EarlyStoppingConfig {
pub min_delta: Float,
pub patience: usize,
pub monitor: String,
pub mode_max: bool,
pub restore_best_weights: bool,
}
impl Default for EarlyStoppingConfig {
fn default() -> Self {
Self {
min_delta: 1e-4,
patience: 10,
monitor: "loss".to_string(),
mode_max: false,
restore_best_weights: true,
}
}
}
#[derive(Debug, Clone)]
pub struct EarlyStopping {
config: EarlyStoppingConfig,
best_value: Option<Float>,
best_iteration: usize,
wait_count: usize,
should_stop: bool,
}
impl EarlyStopping {
pub fn new(config: EarlyStoppingConfig) -> Self {
Self {
config,
best_value: None,
best_iteration: 0,
wait_count: 0,
should_stop: false,
}
}
pub fn update(&mut self, value: Float, iteration: usize) -> bool {
match self.best_value {
None => {
self.best_value = Some(value);
self.best_iteration = iteration;
false
}
Some(best) => {
let is_improvement = if self.config.mode_max {
value > best + self.config.min_delta
} else {
value < best - self.config.min_delta
};
if is_improvement {
self.best_value = Some(value);
self.best_iteration = iteration;
self.wait_count = 0;
false
} else {
self.wait_count += 1;
if self.wait_count >= self.config.patience {
self.should_stop = true;
true
} else {
false
}
}
}
}
}
pub fn should_stop(&self) -> bool {
self.should_stop
}
pub fn best_value(&self) -> Option<Float> {
self.best_value
}
pub fn best_iteration(&self) -> usize {
self.best_iteration
}
}
#[derive(Debug, Clone)]
pub struct WarmStartRegressorConfig {
pub max_iter: usize,
pub learning_rate: Float,
pub alpha: Float,
pub tol: Float,
pub early_stopping: Option<EarlyStoppingConfig>,
pub verbose: bool,
}
impl Default for WarmStartRegressorConfig {
fn default() -> Self {
Self {
max_iter: 1000,
learning_rate: 0.01,
alpha: 0.0001,
tol: 1e-4,
early_stopping: Some(EarlyStoppingConfig::default()),
verbose: false,
}
}
}
#[derive(Debug, Clone)]
pub struct WarmStartRegressor<S = Untrained> {
state: S,
config: WarmStartRegressorConfig,
}
#[derive(Debug, Clone)]
pub struct WarmStartRegressorTrained {
pub coef: Array2<Float>,
pub intercept: Array1<Float>,
pub n_features: usize,
pub n_outputs: usize,
pub n_iter: usize,
pub loss_history: Vec<Float>,
pub best_loss: Float,
pub best_iter: usize,
pub best_coef: Option<Array2<Float>>,
pub best_intercept: Option<Array1<Float>>,
pub converged: bool,
pub config: WarmStartRegressorConfig,
}
impl WarmStartRegressor<Untrained> {
pub fn new() -> Self {
Self {
state: Untrained,
config: WarmStartRegressorConfig::default(),
}
}
pub fn config(mut self, config: WarmStartRegressorConfig) -> Self {
self.config = config;
self
}
pub fn max_iter(mut self, max_iter: usize) -> Self {
self.config.max_iter = max_iter;
self
}
pub fn learning_rate(mut self, lr: Float) -> Self {
self.config.learning_rate = lr;
self
}
pub fn early_stopping(mut self, config: EarlyStoppingConfig) -> Self {
self.config.early_stopping = Some(config);
self
}
}
impl Default for WarmStartRegressor<Untrained> {
fn default() -> Self {
Self::new()
}
}
impl Fit<ArrayView2<'_, Float>, ArrayView2<'_, Float>> for WarmStartRegressor<Untrained> {
type Fitted = WarmStartRegressor<WarmStartRegressorTrained>;
#[allow(non_snake_case)] fn fit(self, X: &ArrayView2<Float>, y: &ArrayView2<Float>) -> SklResult<Self::Fitted> {
if X.nrows() != y.nrows() {
return Err(SklearsError::InvalidInput(
"Number of samples in X and y must match".to_string(),
));
}
let n_samples = X.nrows();
let n_features = X.ncols();
let n_outputs = y.ncols();
let mut coef = Array2::zeros((n_features, n_outputs));
let mut intercept = Array1::zeros(n_outputs);
let mut loss_history = Vec::new();
let mut best_loss = Float::INFINITY;
let mut best_iter = 0;
let mut best_coef = None;
let mut best_intercept = None;
let mut early_stopping = self
.config
.early_stopping
.as_ref()
.map(|cfg| EarlyStopping::new(cfg.clone()));
let mut converged = false;
for iter in 0..self.config.max_iter {
let mut total_loss = 0.0;
for i in 0..n_samples {
let x_i = X.row(i);
let y_i = y.row(i);
let pred = coef.t().dot(&x_i) + &intercept;
let error = &y_i - &pred;
total_loss += error.mapv(|x| x.powi(2)).sum();
for j in 0..n_features {
for k in 0..n_outputs {
let gradient = -error[k] * x_i[j] + self.config.alpha * coef[[j, k]];
coef[[j, k]] -= self.config.learning_rate * gradient;
}
}
for k in 0..n_outputs {
intercept[k] += self.config.learning_rate * error[k];
}
}
let avg_loss = total_loss / (n_samples as Float * n_outputs as Float);
loss_history.push(avg_loss);
if avg_loss < best_loss {
best_loss = avg_loss;
best_iter = iter;
if self.config.early_stopping.is_some() {
best_coef = Some(coef.clone());
best_intercept = Some(intercept.clone());
}
}
if iter > 0 && (loss_history[iter - 1] - avg_loss).abs() < self.config.tol {
converged = true;
if self.config.verbose {
println!("Converged at iteration {}", iter);
}
break;
}
if let Some(ref mut es) = early_stopping {
if es.update(avg_loss, iter) {
if self.config.verbose {
println!("Early stopping at iteration {}", iter);
}
break;
}
}
if self.config.verbose && iter % 100 == 0 {
println!("Iteration {}: loss = {:.6}", iter, avg_loss);
}
}
if let Some(cfg) = &self.config.early_stopping {
if cfg.restore_best_weights {
if let Some(ref best_c) = best_coef {
coef = best_c.clone();
}
if let Some(ref best_i) = best_intercept {
intercept = best_i.clone();
}
}
}
Ok(WarmStartRegressor {
state: WarmStartRegressorTrained {
coef,
intercept,
n_features,
n_outputs,
n_iter: loss_history.len(),
loss_history,
best_loss,
best_iter,
best_coef,
best_intercept,
converged,
config: self.config,
},
config: WarmStartRegressorConfig::default(),
})
}
}
impl WarmStartRegressor<WarmStartRegressorTrained> {
#[allow(non_snake_case)] pub fn continue_training(
mut self,
X: &ArrayView2<Float>,
y: &ArrayView2<Float>,
additional_iterations: usize,
) -> SklResult<Self> {
if X.nrows() != y.nrows() {
return Err(SklearsError::InvalidInput(
"Number of samples in X and y must match".to_string(),
));
}
if X.ncols() != self.state.n_features || y.ncols() != self.state.n_outputs {
return Err(SklearsError::InvalidInput(
"Feature or output dimensions do not match".to_string(),
));
}
let n_samples = X.nrows();
let mut early_stopping = self
.state
.config
.early_stopping
.as_ref()
.map(|cfg| EarlyStopping::new(cfg.clone()));
for iter in 0..additional_iterations {
let mut total_loss = 0.0;
for i in 0..n_samples {
let x_i = X.row(i);
let y_i = y.row(i);
let pred = self.state.coef.t().dot(&x_i) + &self.state.intercept;
let error = &y_i - &pred;
total_loss += error.mapv(|x| x.powi(2)).sum();
for j in 0..self.state.n_features {
for k in 0..self.state.n_outputs {
let gradient =
-error[k] * x_i[j] + self.state.config.alpha * self.state.coef[[j, k]];
self.state.coef[[j, k]] -= self.state.config.learning_rate * gradient;
}
}
for k in 0..self.state.n_outputs {
self.state.intercept[k] += self.state.config.learning_rate * error[k];
}
}
let avg_loss = total_loss / (n_samples as Float * self.state.n_outputs as Float);
self.state.loss_history.push(avg_loss);
if avg_loss < self.state.best_loss {
self.state.best_loss = avg_loss;
self.state.best_iter = self.state.n_iter + iter;
if self.state.config.early_stopping.is_some() {
self.state.best_coef = Some(self.state.coef.clone());
self.state.best_intercept = Some(self.state.intercept.clone());
}
}
let loss_len = self.state.loss_history.len();
if loss_len > 1 {
let prev_loss = self.state.loss_history[loss_len - 2];
if (prev_loss - avg_loss).abs() < self.state.config.tol {
self.state.converged = true;
break;
}
}
if let Some(ref mut es) = early_stopping {
if es.update(avg_loss, self.state.n_iter + iter) {
break;
}
}
}
self.state.n_iter += additional_iterations;
Ok(self)
}
pub fn loss_history(&self) -> &[Float] {
&self.state.loss_history
}
pub fn best_loss(&self) -> Float {
self.state.best_loss
}
pub fn converged(&self) -> bool {
self.state.converged
}
pub fn coef(&self) -> &Array2<Float> {
&self.state.coef
}
pub fn n_iter(&self) -> usize {
self.state.n_iter
}
}
impl Predict<ArrayView2<'_, Float>, Array2<Float>>
for WarmStartRegressor<WarmStartRegressorTrained>
{
#[allow(non_snake_case)] fn predict(&self, X: &ArrayView2<Float>) -> SklResult<Array2<Float>> {
if X.ncols() != self.state.n_features {
return Err(SklearsError::InvalidInput(format!(
"Expected {} features, got {}",
self.state.n_features,
X.ncols()
)));
}
let n_samples = X.nrows();
let mut predictions = Array2::zeros((n_samples, self.state.n_outputs));
for i in 0..n_samples {
let x_i = X.row(i);
let pred = self.state.coef.t().dot(&x_i) + &self.state.intercept;
predictions.row_mut(i).assign(&pred);
}
Ok(predictions)
}
}
impl Estimator for WarmStartRegressor<Untrained> {
type Config = WarmStartRegressorConfig;
type Error = SklearsError;
type Float = Float;
fn config(&self) -> &Self::Config {
&self.config
}
}
impl Estimator for WarmStartRegressor<WarmStartRegressorTrained> {
type Config = WarmStartRegressorConfig;
type Error = SklearsError;
type Float = Float;
fn config(&self) -> &Self::Config {
&self.state.config
}
}
#[derive(Debug, Clone)]
pub struct PredictionCache {
cache: HashMap<u64, Array2<Float>>,
max_size: usize,
hits: usize,
misses: usize,
}
impl PredictionCache {
pub fn new(max_size: usize) -> Self {
Self {
cache: HashMap::new(),
max_size,
hits: 0,
misses: 0,
}
}
#[allow(non_snake_case)] pub fn get(&mut self, X: &ArrayView2<Float>) -> Option<Array2<Float>> {
let hash = self.hash_input(X);
if let Some(pred) = self.cache.get(&hash) {
self.hits += 1;
Some(pred.clone())
} else {
self.misses += 1;
None
}
}
#[allow(non_snake_case)] pub fn put(&mut self, X: &ArrayView2<Float>, prediction: Array2<Float>) {
if self.cache.len() >= self.max_size {
if let Some(first_key) = self.cache.keys().next().copied() {
self.cache.remove(&first_key);
}
}
let hash = self.hash_input(X);
self.cache.insert(hash, prediction);
}
pub fn clear(&mut self) {
self.cache.clear();
}
pub fn stats(&self) -> (usize, usize, Float) {
let total = self.hits + self.misses;
let hit_rate = if total > 0 {
self.hits as Float / total as Float
} else {
0.0
};
(self.hits, self.misses, hit_rate)
}
#[allow(non_snake_case)] fn hash_input(&self, X: &ArrayView2<Float>) -> u64 {
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
let mut hasher = DefaultHasher::new();
for &val in X.iter() {
val.to_bits().hash(&mut hasher);
}
hasher.finish()
}
}
#[cfg(test)]
#[allow(non_snake_case)] mod tests {
use super::*;
use approx::assert_abs_diff_eq;
use scirs2_core::ndarray::array;
#[test]
fn test_early_stopping_basic() {
let config = EarlyStoppingConfig {
min_delta: 0.1,
patience: 3,
mode_max: false,
..Default::default()
};
let mut es = EarlyStopping::new(config);
assert!(!es.update(1.0, 0));
assert!(!es.update(0.8, 1)); assert!(!es.update(0.79, 2)); assert!(!es.update(0.78, 3)); assert!(es.update(0.77, 4)); }
#[test]
fn test_early_stopping_mode_max() {
let config = EarlyStoppingConfig {
min_delta: 0.01,
patience: 2,
mode_max: true,
..Default::default()
};
let mut es = EarlyStopping::new(config);
assert!(!es.update(0.5, 0));
assert!(!es.update(0.6, 1)); assert!(!es.update(0.59, 2)); assert!(es.update(0.58, 3)); }
#[test]
#[allow(non_snake_case)]
fn test_warm_start_regressor_basic() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 4.0]];
let y = array![[1.0, 2.0], [2.0, 3.0], [3.0, 4.0]];
let model = WarmStartRegressor::new().max_iter(100).learning_rate(0.1);
let trained = model
.fit(&X.view(), &y.view())
.expect("model fitting should succeed");
let predictions = trained
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (3, 2));
assert!(trained.n_iter() > 0);
}
#[test]
#[allow(non_snake_case)]
fn test_warm_start_continue_training() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1.0, 2.0], [2.0, 3.0]];
let model = WarmStartRegressor::new().max_iter(10).learning_rate(0.1);
let trained = model
.fit(&X.view(), &y.view())
.expect("model fitting should succeed");
let initial_iter = trained.n_iter();
let initial_loss = trained
.loss_history()
.last()
.copied()
.expect("collection should not be empty");
let continued = trained
.continue_training(&X.view(), &y.view(), 20)
.expect("operation should succeed");
let final_loss = continued
.loss_history()
.last()
.copied()
.expect("collection should not be empty");
assert!(continued.n_iter() > initial_iter);
assert!(final_loss <= initial_loss + 1.0); }
#[test]
#[allow(non_snake_case)]
fn test_warm_start_with_early_stopping() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 4.0]];
let y = array![[1.0, 2.0], [2.0, 3.0], [3.0, 4.0]];
let es_config = EarlyStoppingConfig {
patience: 5,
min_delta: 1e-6,
..Default::default()
};
let model = WarmStartRegressor::new()
.max_iter(1000)
.early_stopping(es_config)
.learning_rate(0.1);
let trained = model
.fit(&X.view(), &y.view())
.expect("model fitting should succeed");
assert!(trained.n_iter() < 1000);
assert!(trained.best_loss() < Float::INFINITY);
}
#[test]
fn test_prediction_cache_basic() {
let mut cache = PredictionCache::new(10);
let X = array![[1.0, 2.0], [2.0, 3.0]];
let pred = array![[1.0, 2.0], [2.0, 3.0]];
assert!(cache.get(&X.view()).is_none());
cache.put(&X.view(), pred.clone());
let cached = cache.get(&X.view()).expect("index should be valid");
assert_eq!(cached.dim(), pred.dim());
assert_eq!(cache.stats().0, 1); assert_eq!(cache.stats().1, 1); }
#[test]
fn test_prediction_cache_eviction() {
let mut cache = PredictionCache::new(2);
let X1 = array![[1.0, 2.0]];
let X2 = array![[2.0, 3.0]];
let X3 = array![[3.0, 4.0]];
let pred = array![[1.0, 2.0]];
cache.put(&X1.view(), pred.clone());
cache.put(&X2.view(), pred.clone());
cache.put(&X3.view(), pred.clone());
assert_eq!(cache.cache.len(), 2);
}
#[test]
fn test_cache_stats() {
let mut cache = PredictionCache::new(10);
let X = array![[1.0, 2.0]];
let pred = array![[1.0, 2.0]];
cache.get(&X.view()); cache.put(&X.view(), pred);
cache.get(&X.view()); cache.get(&X.view());
let (hits, misses, hit_rate) = cache.stats();
assert_eq!(hits, 2);
assert_eq!(misses, 1);
assert_abs_diff_eq!(hit_rate, 2.0 / 3.0, epsilon = 1e-6);
}
}