use std::{
iter::Sum,
ops::{AddAssign, DivAssign, MulAssign, SubAssign},
};
use num_traits::{Float, cast::AsPrimitive, float::FloatCore};
use rustfft::FftNum;
use rayon::{
iter::{IndexedParallelIterator, IntoParallelRefMutIterator, ParallelIterator},
slice::ParallelSliceMut,
};
use crate::tsne;
pub(crate) trait Repulsion<T, const D: usize> {
fn step(
&mut self,
y: &[T],
p_rows: &[usize],
p_columns: &[u32],
p_values: &[T],
positive: &mut [T],
negative: &mut [T],
) -> T;
fn error(
&self,
p_rows: &[usize],
p_columns: &[u32],
p_values: &[T],
y: &[T],
n_samples: usize,
) -> T;
}
pub(crate) struct BarnesHutRepulsion<T, W, const D: usize>
where
W: barnes_hut_tree::MortonWord,
{
arena: barnes_hut_tree::BarnesHutTree<T, W, D>,
q_sums: Vec<T>,
theta: T,
theta_sq: T,
}
impl<T, W, const D: usize> BarnesHutRepulsion<T, W, D>
where
T: Float + FloatCore + Send + Sync + AddAssign,
W: barnes_hut_tree::MortonWord,
{
pub(crate) fn new(theta: T) -> Self {
Self {
arena: barnes_hut_tree::BarnesHutTree::empty(),
q_sums: Vec::new(),
theta,
theta_sq: theta * theta,
}
}
}
impl<T, W, const D: usize> Repulsion<T, D> for BarnesHutRepulsion<T, W, D>
where
T: Float + FloatCore + Send + Sync + Sum + AddAssign + SubAssign + MulAssign + DivAssign,
W: barnes_hut_tree::MortonWord,
barnes_hut_tree::Dim<D>: barnes_hut_tree::Morton<D, Word = W>,
{
fn step(
&mut self,
y: &[T],
p_rows: &[usize],
p_columns: &[u32],
p_values: &[T],
positive: &mut [T],
negative: &mut [T],
) -> T {
let n_samples = y.len() / D;
self.arena.rebuild_uniform(y);
self.q_sums.resize(n_samples, T::zero());
let theta_sq = self.theta_sq;
let arena = &self.arena;
positive
.par_chunks_mut(D)
.zip(negative.par_chunks_mut(D))
.zip(self.q_sums.par_iter_mut())
.enumerate()
.for_each_init(
|| {
(
[T::zero(); D],
[T::zero(); D],
<barnes_hut_tree::Dim<D> as barnes_hut_tree::Morton<D>>::empty_stack(),
)
},
|(edge_row, nonedge_row, stack),
(index, ((positive_out, negative_out), q_sum_out))| {
*edge_row = [T::zero(); D];
*nonedge_row = [T::zero(); D];
let mut q_sum = T::zero();
tsne::compute_edge_forces::<T, D>(
index, y, p_rows, p_columns, p_values, edge_row,
);
arena.compute_non_edge_forces(
index,
theta_sq,
y,
nonedge_row,
&mut q_sum,
stack.as_mut(),
);
positive_out.copy_from_slice(&edge_row[..]);
negative_out.copy_from_slice(&nonedge_row[..]);
*q_sum_out = q_sum;
},
);
let q_sum: T = self.q_sums.iter().copied().sum();
Float::recip(q_sum)
}
fn error(
&self,
p_rows: &[usize],
p_columns: &[u32],
p_values: &[T],
y: &[T],
n_samples: usize,
) -> T {
tsne::evaluate_error_approximately::<T, D>(
p_rows, p_columns, p_values, y, n_samples, self.theta,
)
}
}
pub(crate) struct InterpolatedRepulsion<T: FftNum, const D: usize> {
interpolant: tsne::interpolation::Interpolant<T, D>,
}
impl<T, const D: usize> InterpolatedRepulsion<T, D>
where
T: Send + Sync + Float + FftNum + AsPrimitive<usize> + Sum,
{
pub(crate) fn new() -> Self {
Self {
interpolant: tsne::interpolation::Interpolant::new(),
}
}
}
impl<T, const D: usize> Repulsion<T, D> for InterpolatedRepulsion<T, D>
where
T: Float
+ FftNum
+ Send
+ Sync
+ Sum
+ AsPrimitive<usize>
+ AddAssign
+ SubAssign
+ MulAssign
+ DivAssign,
{
fn step(
&mut self,
y: &[T],
p_rows: &[usize],
p_columns: &[u32],
p_values: &[T],
positive: &mut [T],
negative: &mut [T],
) -> T {
let n_samples = y.len() / D;
positive
.par_chunks_mut(D)
.enumerate()
.for_each(|(index, row)| {
row.fill(T::zero());
tsne::compute_edge_forces::<T, D>(index, y, p_rows, p_columns, p_values, row);
});
let mut z = T::zero();
self.interpolant
.repulsive_forces(y, n_samples, negative, &mut z);
z.recip()
}
fn error(
&self,
p_rows: &[usize],
p_columns: &[u32],
p_values: &[T],
y: &[T],
n_samples: usize,
) -> T {
tsne::evaluate_error_interpolated::<T, D>(p_rows, p_columns, p_values, y, n_samples)
}
}