use crate::traits::{LogicalColumn, RawColumn, Scalar, VectorView, VectorViewMut};
#[derive(Clone, Copy, Debug)]
pub struct SparseColumnRef<'a, F> {
row_indices: &'a [usize],
values: &'a [F],
len: usize,
}
impl<'a, F> SparseColumnRef<'a, F> {
pub(crate) fn new(row_indices: &'a [usize], values: &'a [F], len: usize) -> Self {
Self {
row_indices,
values,
len,
}
}
pub fn row_indices(&self) -> &'a [usize] {
self.row_indices
}
pub fn values(&self) -> &'a [F] {
self.values
}
}
impl<F: Scalar> RawColumn<F> for SparseColumnRef<'_, F> {
fn len(&self) -> usize {
self.len
}
fn stored_len(&self) -> usize {
self.values.len()
}
fn for_each_stored(&self, mut f: impl FnMut(usize, F)) {
for (&row, &value) in self.row_indices.iter().zip(self.values) {
f(row, value);
}
}
}
#[derive(Clone, Copy, Debug)]
pub struct LazyColumn<C, F> {
raw: C,
center: F,
scale: F,
}
impl<C, F> LazyColumn<C, F> {
pub(crate) fn new(raw: C, center: F, scale: F) -> Self {
Self { raw, center, scale }
}
pub fn raw(&self) -> &C {
&self.raw
}
}
pub type LazySparseColumn<'a, F> = LazyColumn<SparseColumnRef<'a, F>, F>;
impl<'a, F: Scalar> LazySparseColumn<'a, F> {
pub fn row_indices(&self) -> &'a [usize] {
self.raw.row_indices()
}
pub fn values(&self) -> &'a [F] {
self.raw.values()
}
pub fn implicit_value(&self) -> F {
-self.center / self.scale
}
pub fn raw_sum(&self) -> F {
self.raw.raw_sum()
}
pub fn stored_corrections(&self) -> impl Iterator<Item = (usize, F)> + '_ {
self.row_indices()
.iter()
.copied()
.zip(self.values().iter().map(|&value| value / self.scale))
}
}
impl<C: RawColumn<F>, F: Scalar> LazyColumn<C, F> {
pub fn len(&self) -> usize {
self.raw.len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn center(&self) -> F {
self.center
}
pub fn scale(&self) -> F {
self.scale
}
pub fn sum(&self) -> F {
let len = F::from_usize(self.len()).unwrap();
(self.raw.raw_sum() - len * self.center) / self.scale
}
pub fn norm_squared(&self) -> F {
let mut stored_squared_deviations = F::zero();
self.raw.for_each_stored(|_, value| {
let deviation = value - self.center;
stored_squared_deviations = stored_squared_deviations + deviation * deviation;
});
let implicit_count = F::from_usize(self.len() - self.raw.stored_len()).unwrap();
let centered_norm_squared =
stored_squared_deviations + implicit_count * self.center * self.center;
centered_norm_squared / (self.scale * self.scale)
}
pub fn dot<V: VectorView<F> + ?Sized>(&self, vector: &V) -> F {
assert_eq!(
vector.len(),
self.len(),
"vector length must equal column length"
);
let vector_sum = vector.sum();
self.dot_with_sum(vector, vector_sum)
}
pub fn dot_with_sum<V: VectorView<F> + ?Sized>(&self, vector: &V, vector_sum: F) -> F {
assert_eq!(
vector.len(),
self.len(),
"vector length must equal column length"
);
let mut raw_dot = F::zero();
self.raw
.for_each_stored(|row, value| raw_dot = raw_dot + value * vector.get(row));
(raw_dot - self.center * vector_sum) / self.scale
}
pub fn weighted_dot<V, W>(&self, vector: &V, weights: &W) -> F
where
V: VectorView<F> + ?Sized,
W: VectorView<F> + ?Sized,
{
assert_eq!(
vector.len(),
self.len(),
"vector length must equal column length"
);
assert_eq!(
weights.len(),
self.len(),
"weights length must equal column length"
);
let weighted_vector_sum = (0..self.len())
.map(|i| vector.get(i) * weights.get(i))
.sum();
self.weighted_dot_with_sum(vector, weights, weighted_vector_sum)
}
pub fn weighted_dot_with_sum<V, W>(&self, vector: &V, weights: &W, weighted_vector_sum: F) -> F
where
V: VectorView<F> + ?Sized,
W: VectorView<F> + ?Sized,
{
assert_eq!(
vector.len(),
self.len(),
"vector length must equal column length"
);
assert_eq!(
weights.len(),
self.len(),
"weights length must equal column length"
);
let mut raw_weighted_dot = F::zero();
self.raw.for_each_stored(|row, value| {
raw_weighted_dot = raw_weighted_dot + value * weights.get(row) * vector.get(row);
});
(raw_weighted_dot - self.center * weighted_vector_sum) / self.scale
}
pub fn weighted_norm_squared<W: VectorView<F> + ?Sized>(&self, weights: &W) -> F {
assert_eq!(
weights.len(),
self.len(),
"weights length must equal column length"
);
self.weighted_norm_squared_with_sum(weights, weights.sum())
}
pub fn weighted_norm_squared_with_sum<W: VectorView<F> + ?Sized>(
&self,
weights: &W,
weight_sum: F,
) -> F {
assert_eq!(
weights.len(),
self.len(),
"weights length must equal column length"
);
let mut stored_squared_deviations = F::zero();
let mut stored_weight = F::zero();
self.raw.for_each_stored(|row, value| {
let weight = weights.get(row);
let deviation = value - self.center;
stored_squared_deviations = stored_squared_deviations + weight * deviation * deviation;
stored_weight = stored_weight + weight;
});
let implicit_weight = weight_sum - stored_weight;
let centered_norm_squared =
stored_squared_deviations + implicit_weight * self.center * self.center;
centered_norm_squared / (self.scale * self.scale)
}
pub fn scaled_add_to<V: VectorViewMut<F> + ?Sized>(&self, alpha: F, destination: &mut V) {
self.raw.affine_add_to(
alpha / self.scale,
-alpha * self.center / self.scale,
destination,
);
}
}
impl<C: RawColumn<F>, F: Scalar> LogicalColumn<F> for LazyColumn<C, F> {
fn len(&self) -> usize {
self.len()
}
fn center(&self) -> F {
self.center()
}
fn scale(&self) -> F {
self.scale()
}
fn sum(&self) -> F {
self.sum()
}
fn norm_squared(&self) -> F {
self.norm_squared()
}
fn dot<V: VectorView<F> + ?Sized>(&self, vector: &V) -> F {
self.dot(vector)
}
fn dot_with_sum<V: VectorView<F> + ?Sized>(&self, vector: &V, vector_sum: F) -> F {
self.dot_with_sum(vector, vector_sum)
}
fn weighted_dot<V, W>(&self, vector: &V, weights: &W) -> F
where
V: VectorView<F> + ?Sized,
W: VectorView<F> + ?Sized,
{
self.weighted_dot(vector, weights)
}
fn weighted_dot_with_sum<V, W>(&self, vector: &V, weights: &W, sum: F) -> F
where
V: VectorView<F> + ?Sized,
W: VectorView<F> + ?Sized,
{
self.weighted_dot_with_sum(vector, weights, sum)
}
fn weighted_norm_squared<W: VectorView<F> + ?Sized>(&self, weights: &W) -> F {
self.weighted_norm_squared(weights)
}
fn weighted_norm_squared_with_sum<W: VectorView<F> + ?Sized>(&self, weights: &W, sum: F) -> F {
self.weighted_norm_squared_with_sum(weights, sum)
}
fn scaled_add_to<V: VectorViewMut<F> + ?Sized>(&self, alpha: F, destination: &mut V) {
self.scaled_add_to(alpha, destination);
}
}