use std::{
iter::Sum,
ops::{AddAssign, DivAssign, MulAssign, SubAssign},
};
use num_traits::{Float, cast::AsPrimitive};
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, const D: usize> {
arena: tsne::arena::Arena<T, D>,
q_sums: Vec<T>,
theta: T,
theta_sq: T,
}
impl<T, const D: usize> BarnesHutRepulsion<T, D>
where
T: Float + Send + Sync + AddAssign,
{
pub(crate) fn new(theta: T) -> Self {
Self {
arena: tsne::arena::Arena::empty(),
q_sums: Vec::new(),
theta,
theta_sq: theta * theta,
}
}
}
impl<T, const D: usize> Repulsion<T, D> for BarnesHutRepulsion<T, D>
where
T: Float + Send + Sync + Sum + AddAssign + SubAssign + MulAssign + DivAssign,
tsne::morton::Dim<D>: tsne::morton::Morton<D>,
{
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(y, n_samples);
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],
[0u32; tsne::arena::STACK_CAP],
)
},
|(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::arena::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,
);
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();
q_sum.recip()
}
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::arena::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)
}
}