From c93abe19a6143be73df8666de46c0852bd32f568 Mon Sep 17 00:00:00 2001 From: LucaCappelletti94 Date: Wed, 1 Jul 2026 14:02:09 +0200 Subject: [PATCH] Bump bhtsne from 0.5.4 to 0.7 in linfa-tsne --- algorithms/linfa-tsne/Cargo.toml | 2 +- algorithms/linfa-tsne/src/lib.rs | 79 +++++++++++++++++++++----------- 2 files changed, 54 insertions(+), 27 deletions(-) diff --git a/algorithms/linfa-tsne/Cargo.toml b/algorithms/linfa-tsne/Cargo.toml index 1cf57e622..78c34eab0 100644 --- a/algorithms/linfa-tsne/Cargo.toml +++ b/algorithms/linfa-tsne/Cargo.toml @@ -16,7 +16,7 @@ categories = ["algorithms", "mathematics", "science"] [dependencies] thiserror = "2.0" ndarray = { version = "0.16" } -bhtsne = { version = "0.5.4", default-features = false } +bhtsne = { version = "0.7.9", default-features = false } linfa = { version = "0.8.1", path = "../.." } linfa-nn = { version = "0.8.1", path = "../linfa-nn" } diff --git a/algorithms/linfa-tsne/src/lib.rs b/algorithms/linfa-tsne/src/lib.rs index 2279e6d0c..67a89ec98 100644 --- a/algorithms/linfa-tsne/src/lib.rs +++ b/algorithms/linfa-tsne/src/lib.rs @@ -1,5 +1,4 @@ #![doc = include_str!("../README.md")] -use std::convert::TryFrom; use linfa_nn::distance::Distance; use ndarray::{Array2, ArrayView1}; @@ -21,10 +20,6 @@ impl> Transformer, Result>> for TSn return Err(TSneError::EmbeddingSizeTooLarge); } - let Ok(embedding_size) = u8::try_from(self.embedding_size()) else { - return Err(TSneError::EmbeddingSizeTooLarge); - }; - if F::cast(nsamples - 1) < F::cast(3) * self.perplexity() { return Err(TSneError::PerplexityTooLarge); } @@ -37,33 +32,65 @@ impl> Transformer, Result>> for TSn let data: Vec<_> = data.as_slice().unwrap().chunks(nfeatures).collect(); - let mut tsne = bhtsne::tSNE::new(&data); - let tsne = tsne - .embedding_dim(embedding_size) - .perplexity(self.perplexity()) + let embedding = if self.approx_threshold() <= F::zero() { + match self.embedding_size() { + 1 => self.run_exact::<1>(&data, preliminary_iter), + 2 => self.run_exact::<2>(&data, preliminary_iter), + 3 => self.run_exact::<3>(&data, preliminary_iter), + _ => { + return Err(TSneError::EmbeddingSizeTooLarge); + } + } + } else { + match self.embedding_size() { + 2 => self.run_barnes_hut::<2>(&data, preliminary_iter), + 3 => self.run_barnes_hut::<3>(&data, preliminary_iter), + _n => { + return Err(TSneError::EmbeddingSizeTooLarge); + } + } + }; + + Array2::from_shape_vec((nsamples, self.embedding_size()), embedding).map_err(|e| e.into()) + } +} + +impl> TSneValidParams { + #[inline] + fn run_exact(&self, data: &[&[F]], preliminary_iter: usize) -> Vec { + let mut tsne: bhtsne::tSNE = bhtsne::tSNE::new(data); + tsne.perplexity(self.perplexity()) .epochs(self.max_iter()) .stop_lying_epoch(preliminary_iter) .momentum_switch_epoch(preliminary_iter); - let tsne = if self.approx_threshold() <= F::zero() { - // compute exact t-SNE - tsne.exact(|a, b| { - let a = ArrayView1::from(a); - let b = ArrayView1::from(b); - self.metric().distance(a, b) - }) - } else { - // compute barnes-hut t-SNE - tsne.barnes_hut(self.approx_threshold(), |a, b| { - let a = ArrayView1::from(a); - let b = ArrayView1::from(b); - self.metric().distance(a, b) - }) - }; + tsne.exact(|a, b| { + let a = ArrayView1::from(a); + let b = ArrayView1::from(b); + self.metric().distance(a, b) + }); - let embedding = tsne.embedding(); + tsne.embedding() + } - Array2::from_shape_vec((nsamples, self.embedding_size()), embedding).map_err(|e| e.into()) + #[inline] + fn run_barnes_hut(&self, data: &[&[F]], preliminary_iter: usize) -> Vec + where + bhtsne::Dim: bhtsne::Morton, + { + let mut tsne: bhtsne::tSNE = bhtsne::tSNE::new(data); + tsne.perplexity(self.perplexity()) + .epochs(self.max_iter()) + .stop_lying_epoch(preliminary_iter) + .momentum_switch_epoch(preliminary_iter); + + tsne.barnes_hut(self.approx_threshold(), |a, b| { + let a = ArrayView1::from(a); + let b = ArrayView1::from(b); + self.metric().distance(a, b) + }); + + tsne.embedding() } }