use linfa::dataset::{DatasetBase, Labels, Records};
use linfa::metrics::SilhouetteScore;
use linfa::traits::Transformer;
use linfa_clustering::Dbscan;
use linfa_datasets::generate;
use ndarray::array;
use ndarray_npy::write_npy;
use ndarray_rand::rand::SeedableRng;
use rand_xoshiro::Xoshiro256Plus;
fn main() {
let mut rng = Xoshiro256Plus::seed_from_u64(42);
let expected_centroids = array![[10., 10.], [1., 12.], [20., 30.], [-20., 30.],];
let n = 100;
let dataset: DatasetBase<_, _> = generate::blobs(n, &expected_centroids, &mut rng).into();
let min_points = 3;
println!(
"Clustering #{} data points grouped in 4 clusters of {} points each",
dataset.nsamples(),
n
);
let cluster_memberships = Dbscan::params(min_points)
.tolerance(1.)
.transform(dataset)
.unwrap();
let label_count = cluster_memberships.label_count().remove(0);
println!();
println!("Result: ");
for (label, count) in label_count {
match label {
None => println!(" - {count} noise points"),
Some(i) => println!(" - {count} points in cluster {i}"),
}
}
println!();
let silhouette_score = cluster_memberships.silhouette_score().unwrap();
println!("Silhouette score: {silhouette_score}");
let (records, cluster_memberships) = (cluster_memberships.records, cluster_memberships.targets);
write_npy("clustered_dataset.npy", &records).expect("Failed to write .npy file");
write_npy(
"clustered_memberships.npy",
&cluster_memberships.map(|&x| x.map(|c| c as i64).unwrap_or(-1)),
)
.expect("Failed to write .npy file");
}