use scirs2_core::ndarray::{Array2, ArrayView2};
use scirs2_core::random::rngs::StdRng;
use scirs2_core::random::thread_rng;
use scirs2_core::random::{seq::SliceRandom, SeedableRng};
use scirs2_core::RngExt;
use sklears_core::{
error::{Result as SklResult, SklearsError},
traits::{Estimator, Fit, Transform, Untrained},
types::Float,
};
#[derive(Debug, Clone)]
pub struct MiniBatchUMAP<S = Untrained> {
state: S,
n_components: usize,
n_neighbors: usize,
batch_size: usize,
min_dist: f64,
learning_rate: f64,
n_epochs: usize,
random_state: Option<u64>,
}
impl MiniBatchUMAP<Untrained> {
pub fn new() -> Self {
Self {
state: Untrained,
n_components: 2,
n_neighbors: 15,
batch_size: 32,
min_dist: 0.1,
learning_rate: 1.0,
n_epochs: 200,
random_state: None,
}
}
pub fn n_components(mut self, n_components: usize) -> Self {
self.n_components = n_components;
self
}
pub fn n_neighbors(mut self, n_neighbors: usize) -> Self {
self.n_neighbors = n_neighbors;
self
}
pub fn batch_size(mut self, batch_size: usize) -> Self {
self.batch_size = batch_size;
self
}
pub fn min_dist(mut self, min_dist: f64) -> Self {
self.min_dist = min_dist;
self
}
pub fn learning_rate(mut self, learning_rate: f64) -> Self {
self.learning_rate = learning_rate;
self
}
pub fn n_epochs(mut self, n_epochs: usize) -> Self {
self.n_epochs = n_epochs;
self
}
pub fn random_state(mut self, random_state: Option<u64>) -> Self {
self.random_state = random_state;
self
}
}
impl Default for MiniBatchUMAP<Untrained> {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct MBUMAPTrained {
embedding: Array2<f64>,
}
impl Estimator for MiniBatchUMAP<Untrained> {
type Config = ();
type Error = SklearsError;
type Float = Float;
fn config(&self) -> &Self::Config {
&()
}
}
impl Fit<ArrayView2<'_, Float>, ()> for MiniBatchUMAP<Untrained> {
type Fitted = MiniBatchUMAP<MBUMAPTrained>;
fn fit(self, x: &ArrayView2<'_, Float>, _y: &()) -> SklResult<Self::Fitted> {
let (n_samples, _) = x.dim();
if n_samples < 2 {
return Err(SklearsError::InvalidParameter {
name: "n_samples".to_string(),
reason: "Mini-batch UMAP requires at least 2 samples".to_string(),
});
}
if self.n_neighbors >= n_samples {
return Err(SklearsError::InvalidParameter {
name: "n_neighbors".to_string(),
reason: format!(
"must be less than n_samples ({}), got {}",
n_samples, self.n_neighbors
),
});
}
let x_f64 = x.mapv(|v| v);
let knn_graph = self.build_knn_graph(&x_f64)?;
let mut embedding = self.initialize_embedding(n_samples)?;
for epoch in 0..self.n_epochs {
let mut rng = if let Some(seed) = self.random_state {
StdRng::seed_from_u64(seed + epoch as u64)
} else {
StdRng::seed_from_u64(thread_rng().random::<u64>())
};
let mut indices: Vec<usize> = (0..n_samples).collect();
indices.shuffle(&mut rng);
for chunk in indices.chunks(self.batch_size) {
let batch_indices = chunk.to_vec();
self.process_minibatch(&knn_graph, &mut embedding, &batch_indices, epoch)?;
}
}
Ok(MiniBatchUMAP {
state: MBUMAPTrained { embedding },
n_components: self.n_components,
n_neighbors: self.n_neighbors,
batch_size: self.batch_size,
min_dist: self.min_dist,
learning_rate: self.learning_rate,
n_epochs: self.n_epochs,
random_state: self.random_state,
})
}
}
impl Transform<ArrayView2<'_, Float>, Array2<f64>> for MiniBatchUMAP<MBUMAPTrained> {
fn transform(&self, _x: &ArrayView2<'_, Float>) -> SklResult<Array2<f64>> {
Ok(self.state.embedding.clone())
}
}
impl MiniBatchUMAP<Untrained> {
fn build_knn_graph(&self, x: &Array2<f64>) -> SklResult<Array2<f64>> {
let n_samples = x.nrows();
let mut knn_graph = Array2::zeros((n_samples, n_samples));
for i in 0..n_samples {
let mut distances: Vec<(usize, f64)> = Vec::new();
for j in 0..n_samples {
if i != j {
let dist = (&x.row(i) - &x.row(j)).mapv(|v| v * v).sum().sqrt();
distances.push((j, dist));
}
}
distances.sort_by(|a, b| a.1.partial_cmp(&b.1).expect("operation should succeed"));
for &(j, dist) in distances.iter().take(self.n_neighbors) {
let weight = (-dist).exp(); knn_graph[(i, j)] = weight;
}
}
for i in 0..n_samples {
for j in 0..n_samples {
knn_graph[(i, j)] = knn_graph[(i, j)].max(knn_graph[(j, i)]);
knn_graph[(j, i)] = knn_graph[(i, j)];
}
}
Ok(knn_graph)
}
fn initialize_embedding(&self, n_samples: usize) -> SklResult<Array2<f64>> {
let mut rng = if let Some(seed) = self.random_state {
StdRng::seed_from_u64(seed)
} else {
StdRng::seed_from_u64(thread_rng().random::<u64>())
};
let mut embedding = Array2::zeros((n_samples, self.n_components));
let scale = 10.0;
for i in 0..n_samples {
for j in 0..self.n_components {
embedding[[i, j]] = rng.sample::<f64, _>(scirs2_core::StandardNormal) * scale;
}
}
Ok(embedding)
}
fn process_minibatch(
&self,
knn_graph: &Array2<f64>,
embedding: &mut Array2<f64>,
batch_indices: &[usize],
epoch: usize,
) -> SklResult<()> {
let mut rng = if let Some(seed) = self.random_state {
StdRng::seed_from_u64(seed + epoch as u64)
} else {
StdRng::seed_from_u64(thread_rng().random::<u64>())
};
for &i in batch_indices {
for j in 0..knn_graph.ncols() {
let weight = knn_graph[(i, j)];
if weight > 0.0 && i != j {
self.apply_attractive_force(embedding, i, j, weight);
}
}
for _ in 0..5 {
let neg_j = rng.random_range(0..embedding.nrows());
if neg_j != i {
self.apply_repulsive_force(embedding, i, neg_j);
}
}
}
Ok(())
}
fn apply_attractive_force(&self, embedding: &mut Array2<f64>, i: usize, j: usize, weight: f64) {
let dist_sq = (&embedding.row(i) - &embedding.row(j))
.mapv(|x| x * x)
.sum();
let grad_coeff = -2.0 * weight * self.learning_rate / (1.0 + dist_sq);
for d in 0..self.n_components {
let diff = embedding[[i, d]] - embedding[[j, d]];
let update = grad_coeff * diff;
embedding[[i, d]] += update;
embedding[[j, d]] -= update;
}
}
fn apply_repulsive_force(&self, embedding: &mut Array2<f64>, i: usize, j: usize) {
let dist_sq = (&embedding.row(i) - &embedding.row(j))
.mapv(|x| x * x)
.sum();
if dist_sq > 0.0 {
let grad_coeff = 2.0 * self.learning_rate
/ ((self.min_dist * self.min_dist + dist_sq) * (1.0 + dist_sq));
for d in 0..self.n_components {
let diff = embedding[[i, d]] - embedding[[j, d]];
let update = grad_coeff * diff;
embedding[[i, d]] += update;
embedding[[j, d]] -= update;
}
}
}
}
impl MiniBatchUMAP<MBUMAPTrained> {
pub fn embedding(&self) -> &Array2<f64> {
&self.state.embedding
}
}