pub mod dist;
pub mod init;
use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use ndarray_rand::rand;
pub use crate::kmeans::{dist::KMeansDist, init::KMeansInit};
pub struct KMeans {
k_clusters: usize,
max_iter: usize,
tolerance: f64,
init_fn: KMeansInit,
dist_fn: KMeansDist,
}
impl KMeans {
pub fn new_random(k_clusters: usize) -> Self {
assert_ne!(k_clusters, 0);
Self {
k_clusters,
init_fn: KMeansInit::Forgy,
dist_fn: KMeansDist::Euclidean,
tolerance: 1e-4,
max_iter: 300,
}
}
pub fn new_plusplus(k_clusters: usize) -> Self {
assert_ne!(k_clusters, 0);
Self {
init_fn: KMeansInit::PlusPlus,
..Self::new_random(k_clusters)
}
}
pub fn with_dist_metric(mut self, dist_metric: KMeansDist) -> Self {
self.dist_fn = dist_metric;
self
}
pub fn with_tolerance(mut self, tolerance: f64) -> Self {
assert!(tolerance > 0.0);
self.tolerance = tolerance;
self
}
pub fn with_max_iter(mut self, max_iter: usize) -> Self {
assert_ne!(max_iter, 0);
self.max_iter = max_iter;
self
}
pub fn fit(&self, data: ArrayView2<f64>) -> KMeansFitted {
let mut rng = rand::thread_rng();
let mut centroids = self
.init_fn
.run(self.k_clusters, data, &mut rng, self.dist_fn);
let mut memberships = Array1::zeros(data.nrows());
for _ in 0..self.max_iter {
assign_clusters(data, centroids.view(), &mut memberships, self.dist_fn);
let new_centroids = {
let mut counts = Array1::<f64>::zeros(self.k_clusters);
let mut new_centroids = Array2::<f64>::zeros((self.k_clusters, data.ncols()));
for (point, &membership) in data.outer_iter().zip(&memberships) {
let mut centroid = new_centroids.row_mut(membership);
centroid += &point;
counts[membership] += 1.0;
}
for (mut new_centroid, count) in new_centroids.outer_iter_mut().zip(counts) {
if count > 0.0 {
new_centroid /= count;
}
}
new_centroids
};
let distance = self.dist_fn.run(centroids.view(), new_centroids.view());
centroids = new_centroids;
if distance < self.tolerance {
break;
}
}
KMeansFitted {
centroids,
dist_method: self.dist_fn,
}
}
pub fn fit_predict(&self, data: ArrayView2<f64>) -> Array1<usize> {
self.fit(data).predict(data)
}
}
pub struct KMeansFitted {
centroids: Array2<f64>,
dist_method: KMeansDist,
}
impl KMeansFitted {
pub fn centroids(&self) -> ArrayView2<f64> {
self.centroids.view()
}
pub fn predict_inplace(&self, data: ArrayView2<f64>, memberships: &mut Array1<usize>) {
assert_eq!(data.nrows(), memberships.len());
assign_clusters(data, self.centroids(), memberships, self.dist_method);
}
pub fn predict(&self, data: ArrayView2<f64>) -> Array1<usize> {
let mut memberships = Array1::zeros(data.nrows());
assign_clusters(data, self.centroids(), &mut memberships, self.dist_method);
memberships
}
}
fn assign_clusters(
data: ArrayView2<f64>,
centroids: ArrayView2<f64>,
memberships: &mut Array1<usize>,
dist_fn: KMeansDist,
) {
for (point, membership) in data.outer_iter().zip(memberships) {
let (cluster_assignment, _) = closest_centroid(point, centroids, dist_fn);
*membership = cluster_assignment;
}
}
fn closest_centroid(
point: ArrayView1<f64>,
centroids: ArrayView2<f64>,
dist_fn: KMeansDist,
) -> (usize, f64) {
if point.is_empty() || centroids.is_empty() {
unreachable!()
}
let mut cluster_assignment = 0;
let mut min_dist = f64::INFINITY;
for (c_idx, centroid) in centroids.outer_iter().enumerate() {
let dist = dist_fn.run(point, centroid);
if dist < min_dist {
min_dist = dist;
cluster_assignment = c_idx;
}
}
(cluster_assignment, min_dist)
}