use linalg::matrix::Matrix;
use linalg::vector::Vector;
use learning::UnSupModel;
use rand::{Rng, thread_rng};
use libnum::abs;
pub enum InitAlgorithm {
Forgy,
RandomPartition,
KPlusPlus,
}
pub struct KMeansClassifier {
pub iters: usize,
pub k: usize,
pub centroids: Option<Matrix<f64>>,
pub init_algorithm: InitAlgorithm,
}
impl UnSupModel<Matrix<f64>, Vector<usize>> for KMeansClassifier {
fn predict(&self, inputs: &Matrix<f64>) -> Vector<usize> {
if let Some(ref centroids) = self.centroids {
return KMeansClassifier::find_closest_centroids(centroids, inputs).0;
} else {
panic!("Model has not been trained.");
}
}
fn train(&mut self, inputs: &Matrix<f64>) {
self.init_centroids(inputs);
let mut cost = 0.0;
let eps = 1e-14;
for _i in 0..self.iters {
let (idx, distances) = self.get_closest_centroids(inputs);
self.update_centroids(inputs, idx);
let cost_i = distances.sum();
if abs(cost - cost_i) < eps {
break;
}
cost = cost_i;
}
}
}
impl KMeansClassifier {
pub fn new(k: usize) -> KMeansClassifier {
KMeansClassifier {
iters: 100,
k: k,
centroids: None,
init_algorithm: InitAlgorithm::KPlusPlus,
}
}
fn init_centroids(&mut self, inputs: &Matrix<f64>) {
match self.init_algorithm {
InitAlgorithm::Forgy => {
self.centroids = Some(KMeansClassifier::forgy_init(self.k, inputs))
}
InitAlgorithm::RandomPartition => {
self.centroids = Some(KMeansClassifier::ran_partition_init(self.k, inputs))
}
InitAlgorithm::KPlusPlus => {
self.centroids = Some(KMeansClassifier::plusplus_init(self.k, inputs))
}
}
}
fn update_centroids(&mut self, inputs: &Matrix<f64>, classes: Vector<usize>) {
let mut new_centroids = Vec::with_capacity(self.k * inputs.cols());
for i in 0..self.k {
let mut vec_i = Vec::new();
for j in classes.data().iter() {
if *j == i {
vec_i.push(*j);
}
}
let mat_i = inputs.select_rows(&vec_i);
new_centroids.extend(mat_i.mean(0).data());
}
self.centroids = Some(Matrix::new(self.k, inputs.cols(), new_centroids));
}
fn get_closest_centroids(&self, inputs: &Matrix<f64>) -> (Vector<usize>, Vector<f64>) {
if let Some(ref c) = self.centroids {
return KMeansClassifier::find_closest_centroids(&c, inputs);
} else {
panic!("Centroids not correctly initialized.");
}
}
fn find_closest_centroids(centroids: &Matrix<f64>,
inputs: &Matrix<f64>)
-> (Vector<usize>, Vector<f64>) {
let mut idx = Vec::with_capacity(inputs.rows());
let mut distances = Vec::with_capacity(inputs.rows());
for i in 0..inputs.rows() {
let centroid_diff = centroids - inputs.select_rows(&vec![i; centroids.rows()]);
let dist = ¢roid_diff.elemul(¢roid_diff).sum_cols();
let (min_idx, min_dist) = dist.argmin();
idx.push(min_idx);
distances.push(min_dist);
}
(Vector::new(idx), Vector::new(distances))
}
fn forgy_init(k: usize, inputs: &Matrix<f64>) -> Matrix<f64> {
assert!(k <= inputs.rows());
let mut random_choices = Vec::with_capacity(k);
let mut rng = thread_rng();
while random_choices.len() < k {
let r = rng.gen_range(0, inputs.rows());
if !random_choices.contains(&r) {
random_choices.push(r);
}
}
inputs.select_rows(&random_choices)
}
fn ran_partition_init(k: usize, inputs: &Matrix<f64>) -> Matrix<f64> {
assert!(k <= inputs.rows());
let mut random_assignments = Vec::with_capacity(inputs.rows());
for i in 0..k {
random_assignments.push(i);
}
let mut rng = thread_rng();
for _ in k..inputs.rows() {
random_assignments.push(rng.gen_range(0, k));
}
let mut init_centroids = Vec::with_capacity(k * inputs.cols());
for i in 0..k {
let mut vec_i = Vec::new();
for j in &random_assignments {
if *j == i {
vec_i.push(*j);
}
}
let mat_i = inputs.select_rows(&vec_i);
init_centroids.extend(mat_i.mean(0).into_vec());
}
Matrix::new(k, inputs.cols(), init_centroids)
}
fn plusplus_init(k: usize, inputs: &Matrix<f64>) -> Matrix<f64> {
assert!(k <= inputs.rows());
let mut rng = thread_rng();
let mut init_centroids = Vec::with_capacity(k * inputs.cols());
let first_cen = rng.gen_range(0usize, inputs.rows());
init_centroids.append(&mut inputs.select_rows(&vec![first_cen]).into_vec());
for i in 1..k {
let temp_centroids = Matrix::new(i, inputs.cols(), init_centroids.clone());
let (_, dist) = KMeansClassifier::find_closest_centroids(&temp_centroids, &inputs);
let next_cen = sample_discretely(dist);
init_centroids.append(&mut inputs.select_rows(&vec![next_cen]).into_vec())
}
Matrix::new(k, inputs.cols(), init_centroids)
}
}
fn sample_discretely(unnorm_dist: Vector<f64>) -> usize {
assert!(unnorm_dist.size() > 0);
let sum = unnorm_dist.sum();
let rand = thread_rng().gen_range(0.0f64, sum);
let mut tempsum = 0.0;
for (i, p) in unnorm_dist.data().iter().enumerate() {
tempsum += *p;
if rand < tempsum {
return i;
}
}
panic!("No random value was sampled! There may be more clusters than unique data points.");
}