use scirs2_core::ndarray::Array2;
use scirs2_core::random::seeded_rng;
use serde::{Deserialize, Serialize};
use sklears_core::{
error::{Result, SklearsError},
types::Float,
};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OnlineNMFConfig {
pub n_components: usize,
pub learning_rate: Float,
pub decay_rate: Float,
pub momentum: Float,
pub l1_reg: Float,
pub l2_reg: Float,
pub batch_size: usize,
pub max_iter_per_batch: usize,
pub tolerance: Float,
pub random_state: Option<u64>,
pub init_method: String,
pub shuffle: bool,
}
impl Default for OnlineNMFConfig {
fn default() -> Self {
Self {
n_components: 10,
learning_rate: 0.01,
decay_rate: 0.99,
momentum: 0.9,
l1_reg: 0.0,
l2_reg: 0.01,
batch_size: 100,
max_iter_per_batch: 10,
tolerance: 1e-4,
random_state: None,
init_method: "random".to_string(),
shuffle: true,
}
}
}
#[derive(Debug, Clone)]
pub struct OnlineNMF {
config: OnlineNMFConfig,
w: Option<Array2<Float>>,
w_velocity: Option<Array2<Float>>,
n_samples_seen: usize,
current_learning_rate: Float,
loss_history: Vec<Float>,
fitted: bool,
}
impl OnlineNMF {
pub fn new(config: OnlineNMFConfig) -> Self {
Self {
current_learning_rate: config.learning_rate,
config,
w: None,
w_velocity: None,
n_samples_seen: 0,
loss_history: Vec::new(),
fitted: false,
}
}
pub fn with_defaults() -> Self {
Self::new(OnlineNMFConfig::default())
}
pub fn builder() -> OnlineNMFBuilder {
OnlineNMFBuilder::new()
}
fn initialize_dictionary(&mut self, n_features: usize) -> Result<()> {
let seed = self.config.random_state.unwrap_or(42);
let mut rng = seeded_rng(seed);
match self.config.init_method.as_str() {
"random" => {
let w = Array2::from_shape_fn((n_features, self.config.n_components), |_| {
rng.gen_range(0.0..1.0)
});
self.w = Some(w);
}
"nndsvd" => {
let w = Array2::from_shape_fn((n_features, self.config.n_components), |_| {
rng.gen_range(0.1..1.0)
});
self.w = Some(w);
}
_ => {
return Err(SklearsError::InvalidInput(format!(
"Unknown initialization method: {}",
self.config.init_method
)));
}
}
self.w_velocity = Some(Array2::zeros((n_features, self.config.n_components)));
Ok(())
}
pub fn partial_fit(&mut self, x: &Array2<Float>) -> Result<&mut Self> {
let (n_samples, n_features) = x.dim();
if x.iter().any(|&v| v < 0.0) {
return Err(SklearsError::InvalidInput(
"Input matrix must be non-negative for NMF".to_string(),
));
}
if self.w.is_none() {
self.initialize_dictionary(n_features)?;
}
{
let w = self.w.as_ref().expect("operation should succeed");
if w.nrows() != n_features {
return Err(SklearsError::InvalidInput(format!(
"Feature dimension mismatch: expected {}, got {}",
w.nrows(),
n_features
)));
}
}
let mut h = self.initialize_h(x)?;
for iter in 0..self.config.max_iter_per_batch {
let w_current = self.w.as_ref().expect("operation should succeed").clone();
let reconstruction = h.dot(&w_current.t());
let loss = self.compute_loss(x, &reconstruction, &h);
self.loss_history.push(loss);
if iter > 0 {
let prev_loss = self.loss_history[self.loss_history.len() - 2];
if (prev_loss - loss).abs() < self.config.tolerance {
break;
}
}
self.update_h(&mut h, x, &w_current)?;
self.update_w(x, &h)?;
}
self.n_samples_seen += n_samples;
self.update_learning_rate();
self.fitted = true;
Ok(self)
}
fn initialize_h(&self, x: &Array2<Float>) -> Result<Array2<Float>> {
let (n_samples, _) = x.dim();
let seed = self.config.random_state.unwrap_or(42) + self.n_samples_seen as u64;
let mut rng = seeded_rng(seed);
let h = Array2::from_shape_fn((n_samples, self.config.n_components), |_| {
rng.gen_range(0.0..1.0)
});
Ok(h)
}
fn update_h(&self, h: &mut Array2<Float>, x: &Array2<Float>, w: &Array2<Float>) -> Result<()> {
let reconstruction = h.dot(&w.t());
let residual = x - reconstruction;
let gradient_part = residual.dot(w);
let gradient = gradient_part.mapv(|v| -v)
+ self.config.l1_reg
+ h.mapv(|v| 2.0 * self.config.l2_reg * v);
let lr = self.current_learning_rate;
*h = (h.clone() - gradient.mapv(|v| lr * v)).mapv(|v| v.max(0.0));
Ok(())
}
fn update_w(&mut self, x: &Array2<Float>, h: &Array2<Float>) -> Result<()> {
let w = self.w.as_ref().expect("operation should succeed");
let w_velocity = self.w_velocity.as_ref().expect("operation should succeed");
let reconstruction = h.dot(&w.t());
let residual = x - reconstruction;
let residual_t = residual.t().to_owned();
let gradient_part = residual_t.dot(h);
let gradient = gradient_part.mapv(|v| -v)
+ self.config.l1_reg
+ w.mapv(|v| 2.0 * self.config.l2_reg * v);
let new_velocity = w_velocity.mapv(|v| v * self.config.momentum)
- gradient.mapv(|v| self.current_learning_rate * v);
let new_w = (w + &new_velocity).mapv(|v| v.max(0.0));
self.w = Some(new_w);
self.w_velocity = Some(new_velocity);
Ok(())
}
fn compute_loss(
&self,
x: &Array2<Float>,
reconstruction: &Array2<Float>,
h: &Array2<Float>,
) -> Float {
let residual = x - reconstruction;
let reconstruction_loss = residual.mapv(|v| v.powi(2)).sum();
let l1_term = if self.config.l1_reg > 0.0 {
self.config.l1_reg * h.mapv(|v| v.abs()).sum()
} else {
0.0
};
let l2_term = if self.config.l2_reg > 0.0 {
let w = self.w.as_ref().expect("operation should succeed");
self.config.l2_reg * (w.mapv(|v| v.powi(2)).sum() + h.mapv(|v| v.powi(2)).sum())
} else {
0.0
};
reconstruction_loss + l1_term + l2_term
}
fn update_learning_rate(&mut self) {
self.current_learning_rate *= self.config.decay_rate;
}
pub fn transform(&self, x: &Array2<Float>) -> Result<Array2<Float>> {
if !self.fitted {
return Err(SklearsError::InvalidInput(
"Model must be fitted before transform".to_string(),
));
}
let w = self.w.as_ref().expect("operation should succeed");
let (_n_samples, n_features) = x.dim();
if n_features != w.nrows() {
return Err(SklearsError::InvalidInput(format!(
"Feature dimension mismatch: expected {}, got {}",
w.nrows(),
n_features
)));
}
let mut h = self.initialize_h(x)?;
for _ in 0..self.config.max_iter_per_batch {
self.update_h(&mut h, x, w)?;
}
Ok(h)
}
pub fn fit_transform(&mut self, x: &Array2<Float>) -> Result<Array2<Float>> {
self.partial_fit(x)?;
self.transform(x)
}
pub fn get_components(&self) -> Result<Array2<Float>> {
self.w
.clone()
.ok_or_else(|| SklearsError::InvalidInput("Model not fitted".to_string()))
}
pub fn reconstruction_error(&self, x: &Array2<Float>) -> Result<Float> {
let h = self.transform(x)?;
let w = self.w.as_ref().expect("operation should succeed");
let reconstruction = h.dot(&w.t());
let residual = x - &reconstruction;
Ok(residual.mapv(|v| v.powi(2)).sum().sqrt())
}
pub fn get_loss_history(&self) -> &[Float] {
&self.loss_history
}
pub fn get_n_samples_seen(&self) -> usize {
self.n_samples_seen
}
pub fn reset(&mut self) {
self.w = None;
self.w_velocity = None;
self.n_samples_seen = 0;
self.current_learning_rate = self.config.learning_rate;
self.loss_history.clear();
self.fitted = false;
}
}
pub struct OnlineNMFBuilder {
config: OnlineNMFConfig,
}
impl OnlineNMFBuilder {
pub fn new() -> Self {
Self {
config: OnlineNMFConfig::default(),
}
}
pub fn n_components(mut self, n: usize) -> Self {
self.config.n_components = n;
self
}
pub fn learning_rate(mut self, lr: Float) -> Self {
self.config.learning_rate = lr;
self
}
pub fn decay_rate(mut self, rate: Float) -> Self {
self.config.decay_rate = rate;
self
}
pub fn momentum(mut self, m: Float) -> Self {
self.config.momentum = m;
self
}
pub fn l1_reg(mut self, reg: Float) -> Self {
self.config.l1_reg = reg;
self
}
pub fn l2_reg(mut self, reg: Float) -> Self {
self.config.l2_reg = reg;
self
}
pub fn batch_size(mut self, size: usize) -> Self {
self.config.batch_size = size;
self
}
pub fn max_iter_per_batch(mut self, iter: usize) -> Self {
self.config.max_iter_per_batch = iter;
self
}
pub fn tolerance(mut self, tol: Float) -> Self {
self.config.tolerance = tol;
self
}
pub fn random_state(mut self, seed: u64) -> Self {
self.config.random_state = Some(seed);
self
}
pub fn init_method(mut self, method: &str) -> Self {
self.config.init_method = method.to_string();
self
}
pub fn shuffle(mut self, shuffle: bool) -> Self {
self.config.shuffle = shuffle;
self
}
pub fn build(self) -> OnlineNMF {
OnlineNMF::new(self.config)
}
}
impl Default for OnlineNMFBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::random::thread_rng;
#[test]
fn test_online_nmf_creation() {
let model = OnlineNMF::builder()
.n_components(5)
.learning_rate(0.01)
.build();
assert_eq!(model.config.n_components, 5);
assert!(!model.fitted);
}
#[test]
fn test_online_nmf_partial_fit() {
let mut rng = thread_rng();
let data = Array2::from_shape_fn((20, 10), |_| rng.gen_range(0.0..1.0));
let mut model = OnlineNMF::builder()
.n_components(3)
.max_iter_per_batch(5)
.random_state(42)
.build();
let result = model.partial_fit(&data);
assert!(result.is_ok());
assert!(model.fitted);
assert_eq!(model.n_samples_seen, 20);
}
#[test]
fn test_online_nmf_transform() {
let mut rng = thread_rng();
let data = Array2::from_shape_fn((20, 10), |_| rng.gen_range(0.0..1.0));
let mut model = OnlineNMF::builder()
.n_components(3)
.random_state(42)
.build();
model.partial_fit(&data).expect("operation should succeed");
let transformed = model.transform(&data);
assert!(transformed.is_ok());
let h = transformed.expect("operation should succeed");
assert_eq!(h.dim(), (20, 3));
}
#[test]
fn test_online_nmf_incremental() {
let mut rng = thread_rng();
let mut model = OnlineNMF::builder()
.n_components(3)
.random_state(42)
.build();
for _ in 0..3 {
let batch = Array2::from_shape_fn((10, 8), |_| rng.gen_range(0.0..1.0));
model.partial_fit(&batch).expect("operation should succeed");
}
assert_eq!(model.n_samples_seen, 30);
assert!(model.fitted);
let w = model.get_components().expect("operation should succeed");
assert_eq!(w.dim(), (8, 3));
}
#[test]
fn test_online_nmf_negative_input() {
let data = Array2::from_shape_vec((2, 2), vec![1.0, -1.0, 2.0, 3.0])
.expect("shape and data length should match");
let mut model = OnlineNMF::builder().n_components(2).build();
let result = model.partial_fit(&data);
assert!(result.is_err());
}
#[test]
fn test_online_nmf_reconstruction() {
let mut rng = thread_rng();
let data = Array2::from_shape_fn((15, 8), |_| rng.gen_range(0.0..1.0));
let mut model = OnlineNMF::builder()
.n_components(4)
.max_iter_per_batch(20)
.random_state(42)
.build();
model.partial_fit(&data).expect("operation should succeed");
let error = model
.reconstruction_error(&data)
.expect("operation should succeed");
assert!(error >= 0.0);
assert!(error.is_finite());
}
#[test]
fn test_online_nmf_reset() {
let mut rng = thread_rng();
let data = Array2::from_shape_fn((10, 5), |_| rng.gen_range(0.0..1.0));
let mut model = OnlineNMF::builder()
.n_components(3)
.random_state(42)
.build();
model.partial_fit(&data).expect("operation should succeed");
assert!(model.fitted);
model.reset();
assert!(!model.fitted);
assert_eq!(model.n_samples_seen, 0);
}
#[test]
fn test_online_nmf_builder() {
let model = OnlineNMF::builder()
.n_components(5)
.learning_rate(0.02)
.momentum(0.95)
.l1_reg(0.01)
.l2_reg(0.001)
.batch_size(50)
.tolerance(1e-5)
.random_state(123)
.init_method("random")
.shuffle(true)
.build();
assert_eq!(model.config.n_components, 5);
assert_eq!(model.config.learning_rate, 0.02);
assert_eq!(model.config.momentum, 0.95);
assert_eq!(model.config.random_state, Some(123));
}
}