#[derive(Debug, Clone, PartialEq, Default)]
pub struct SparseVector {
indices: Vec<u32>,
values: Vec<f32>,
}
impl SparseVector {
pub fn new(indices: Vec<u32>, values: Vec<f32>) -> Option<Self> {
if indices.len() != values.len() {
return None;
}
let mut pairs: Vec<(u32, f32)> = indices
.into_iter()
.zip(values)
.filter(|(_, v)| *v != 0.0)
.collect();
pairs.sort_unstable_by_key(|(i, _)| *i);
let (indices, values) = pairs.into_iter().unzip();
Some(Self { indices, values })
}
pub fn indices(&self) -> &[u32] {
&self.indices
}
pub fn values(&self) -> &[f32] {
&self.values
}
pub fn nnz(&self) -> usize {
self.indices.len()
}
pub fn is_empty(&self) -> bool {
self.indices.is_empty()
}
pub fn prune_top_k(&self, k: usize) -> SparseVector {
if k >= self.nnz() {
return self.clone();
}
if k == 0 {
return SparseVector::default();
}
let mut order: Vec<usize> = (0..self.nnz()).collect();
order.sort_unstable_by(|&a, &b| {
self.values[b]
.abs()
.total_cmp(&self.values[a].abs())
.then(self.indices[a].cmp(&self.indices[b]))
});
order.truncate(k);
order.sort_unstable_by_key(|&i| self.indices[i]);
SparseVector {
indices: order.iter().map(|&i| self.indices[i]).collect(),
values: order.iter().map(|&i| self.values[i]).collect(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn new_rejects_length_mismatch() {
assert!(SparseVector::new(vec![1, 2], vec![0.5]).is_none());
}
#[test]
fn new_drops_zeros_and_sorts_by_index() {
let sv = SparseVector::new(vec![5, 1, 3], vec![0.2, 0.9, 0.0]).unwrap();
assert_eq!(sv.indices(), &[1, 5]);
assert_eq!(sv.values(), &[0.9, 0.2]);
assert_eq!(sv.nnz(), 2);
assert!(!sv.is_empty());
}
#[test]
fn new_drops_signed_zero() {
let sv = SparseVector::new(vec![1, 2], vec![-0.0, 0.5]).unwrap();
assert_eq!(sv.indices(), &[2]);
assert_eq!(sv.values(), &[0.5]);
}
#[test]
fn empty_vector_is_empty() {
let sv = SparseVector::new(vec![], vec![]).unwrap();
assert!(sv.is_empty());
assert_eq!(sv.nnz(), 0);
}
#[test]
fn prune_keeps_highest_magnitude_and_index_order() {
let sv = SparseVector::new(vec![1, 2, 3, 4], vec![0.1, 0.9, -0.7, 0.3]).unwrap();
let p = sv.prune_top_k(2);
assert_eq!(p.indices(), &[2, 3]);
assert_eq!(p.values(), &[0.9, -0.7]);
}
#[test]
fn prune_is_noop_when_k_covers_everything() {
let sv = SparseVector::new(vec![1, 2], vec![0.5, 0.4]).unwrap();
assert_eq!(sv.prune_top_k(2), sv);
assert_eq!(sv.prune_top_k(99), sv);
}
#[test]
fn prune_to_zero_yields_empty() {
let sv = SparseVector::new(vec![1, 2], vec![0.5, 0.4]).unwrap();
assert!(sv.prune_top_k(0).is_empty());
}
}