Skip to main content

yui_matrix/sparse/
sp_vec.rs

1//! [`SpVec<R>`]: a sparse column vector — a [`SpMat`](super::SpMat) with one column.
2
3use std::ops::{Add, AddAssign, Neg, Sub, SubAssign, Mul, Range};
4use std::fmt::{Display, Debug};
5use nalgebra_sparse::CscMatrix;
6use nalgebra_sparse::na::{Scalar, ClosedAddAssign, ClosedSubAssign, ClosedMulAssign};
7use num_traits::{Zero, One};
8use auto_impl_ops::auto_ops;
9use yui_core::abst::{Ring, RingOps, AddGrpOps, AddGrp};
10use super::sp_mat::SpMat;
11use crate::Perm;
12
13/// Sparse column vector, stored as a single-column [`SpMat`] (CSC with
14/// `n_cols == 1`).
15#[derive(Clone, Debug)]
16pub struct SpVec<R> {
17    inner: CscMatrix<R> // ncols == 1
18}
19
20impl<R> SpVec<R> {
21    fn new(inner: CscMatrix<R>) -> Self {
22        assert_eq!(inner.ncols(), 1);
23        Self { inner }
24    }
25
26    #[allow(unused)]
27    pub(crate) fn inner(&self) -> &CscMatrix<R> {
28        &self.inner
29    }
30
31    pub(crate) fn into_inner(self) -> CscMatrix<R> {
32        self.inner
33    }
34
35    pub fn data(&self) -> (&[usize], &[R]) {
36        let (_, indices, values) = self.inner.csc_data();
37        (indices, values)
38    }
39
40    pub fn zero(dim: usize) -> Self {
41        let inner = CscMatrix::zeros(dim, 1);
42        Self::new(inner)
43    }
44
45    pub fn is_zero(&self) -> bool
46    where R: Zero {
47        self.inner.values().iter().all(|a| a.is_zero())
48    }
49
50    pub fn unit(n: usize, i: usize) -> Self
51    where R: One {
52        let inner = CscMatrix::try_from_csc_data(
53            n, 1,
54            vec![0, 1],
55            vec![i],
56            vec![R::one()]
57        ).unwrap();
58
59        Self::new(inner)
60    }
61
62    pub fn dim(&self) -> usize {
63        self.inner.nrows()
64    }
65
66    pub fn iter(&self) -> impl Iterator<Item = (usize, &R)> {
67        self.inner.triplet_iter().map(|(i, _, a)| (i, a))
68    }
69
70    pub fn iter_nz(&self) -> impl Iterator<Item = (usize, &R)>
71    where R: Zero {
72        self.iter().filter(|(_, a)| !a.is_zero())
73    }
74
75    pub fn into_dense(self) -> Vec<R>
76    where R: Clone + Zero {
77        self.into()
78    }
79
80    pub fn into_mat(self) -> SpMat<R> {
81        self.into()
82    }
83}
84
85impl<R> From<Vec<R>> for SpVec<R>
86where R: Scalar + Zero + ClosedAddAssign {
87    fn from(vec: Vec<R>) -> Self {
88        Self::from_entries(vec.len(), vec.into_iter().enumerate())
89    }
90}
91
92impl<R> From<SpVec<R>> for Vec<R>
93where R: Clone + Zero {
94    fn from(value: SpVec<R>) -> Self {
95        let mut res = vec![R::zero(); value.dim()];
96        let (_, rows, vals) = value.inner.disassemble();
97        for (i, a) in rows.into_iter().zip(vals) {
98            res[i] = a;
99        }
100        res
101    }
102}
103
104// SpVec(n) as SpMat(n, 1)
105impl<R> From<SpVec<R>> for SpMat<R> {
106    fn from(vec: SpVec<R>) -> Self {
107        SpMat::from(vec.into_inner())
108    }
109}
110
111impl<R> SpMat<R> {
112    fn into_spvec(self) -> SpVec<R> {
113        assert_eq!(self.inner().ncols(), 1);
114        SpVec::new(self.into_inner())
115    }
116}
117
118impl<R> SpVec<R>
119where R: Scalar + Zero + ClosedAddAssign {
120    pub fn try_from_csc_data(dim: usize, row_indices: Vec<usize>, values: Vec<R>) -> Option<SpVec<R>> {
121        let col_offsets = vec![0, row_indices.len()];
122        let csc = CscMatrix::try_from_csc_data(dim, 1, col_offsets, row_indices, values).ok()?;
123        Some(SpMat::from(csc).into_spvec())
124    }
125
126    pub fn from_entries<T>(dim: usize, entries: T) -> Self
127    where T: IntoIterator<Item = (usize, R)> {
128        SpMat::from_entries(
129            (dim, 1),
130            entries.into_iter().map(|(i, a)| (i, 0, a))
131        ).into_spvec()
132    }
133
134    pub fn from_sorted_entries<T>(dim: usize, entries: T) -> Self
135    where T: IntoIterator<Item = (usize, R)> {
136        let init = (vec![], vec![]);
137        let (row_indices, values) = entries.into_iter().fold(init, |mut res, (i, a)| {
138            assert!(i < dim);
139            res.0.push(i);
140            res.1.push(a);
141            res
142        });
143        Self::try_from_csc_data(dim, row_indices, values).unwrap()
144    }
145
146    pub fn extract<F>(&self, dim: usize, f: F) -> SpVec<R>
147    where F: Fn(usize) -> Option<usize> {
148        SpVec::from_entries(dim, self.iter().filter_map(|(i, a)|
149            f(i).map(|i| (i, a.clone()))
150        ))
151    }
152
153    // drops explicitly-stored zeros (CSC arithmetic keeps cancellation zeros).
154    pub fn drop_zeros(self) -> Self {
155        if self.iter().all(|(_, a)| !a.is_zero()) {
156            return self;
157        }
158
159        let n = self.dim();
160        let (_, rows, vals) = self.inner.disassemble();
161        let (rows, vals) = rows.into_iter().zip(vals).filter(|(_, a)| !a.is_zero()).unzip();
162        Self::try_from_csc_data(n, rows, vals).unwrap()
163    }
164
165    pub fn permute(&self, p: &Perm) -> SpVec<R> {
166        self.extract(self.dim(), |i| Some(p.at(i)))
167    }
168
169    pub fn subvec(&self, range: Range<usize>) -> SpVec<R> {
170        self.extract(
171            range.end - range.start,
172            |i| range.contains(&i).then(|| i - range.start)
173        )
174    }
175
176    pub fn stack(top: Self, bot: Self) -> SpVec<R> {
177        let (n1, n2) = (top.dim(), bot.dim());
178        let (_, mut rows, mut vals) = top.inner.disassemble();
179        let (_, bot_rows, bot_vals) = bot.inner.disassemble();
180
181        rows.extend(bot_rows.into_iter().map(|i| i + n1));
182        vals.extend(bot_vals);
183
184        Self::try_from_csc_data(n1 + n2, rows, vals).unwrap()
185    }
186
187    pub fn split(self, at: usize) -> (SpVec<R>, SpVec<R>) {
188        let n = self.dim();
189        assert!(at <= n);
190
191        let (_, mut rows, mut vals) = self.inner.disassemble();
192        let split_idx = rows.partition_point(|&i| i < at);
193
194        let bot_rows: Vec<usize> = rows.split_off(split_idx).into_iter().map(|i| i - at).collect();
195        let bot_vals = vals.split_off(split_idx);
196
197        let top = Self::try_from_csc_data(at, rows, vals).unwrap();
198        let bot = Self::try_from_csc_data(n - at, bot_rows, bot_vals).unwrap();
199        (top, bot)
200    }
201
202}
203
204impl<R> Default for SpVec<R> {
205    fn default() -> Self {
206        Self::zero(0)
207    }
208}
209
210impl<R: PartialEq + Zero> PartialEq for SpVec<R> {
211    fn eq(&self, other: &Self) -> bool {
212        self.dim() == other.dim() && self.iter_nz().eq(other.iter_nz())
213    }
214}
215
216impl<R: Eq + Zero> Eq for SpVec<R> {}
217
218impl<R> Neg for SpVec<R>
219where R: AddGrp, for<'a> &'a R: AddGrpOps<R> {
220    type Output = Self;
221    fn neg(self) -> Self::Output {
222        SpVec { inner: -self.inner }
223    }
224}
225
226impl<R> Neg for &SpVec<R>
227where R: Scalar + Neg<Output = R> {
228    type Output = SpVec<R>;
229    fn neg(self) -> Self::Output {
230        SpVec { inner: -&self.inner }
231    }
232}
233
234macro_rules! impl_binop {
235    ($trait:ident, $method:ident) => {
236        #[auto_ops]
237        impl<R> $trait<&SpVec<R>> for &SpVec<R>
238        where R: Scalar + ClosedAddAssign + ClosedSubAssign + ClosedMulAssign + Zero + One + Neg<Output = R> {
239            type Output = SpVec<R>;
240            fn $method(self, rhs: &SpVec<R>) -> Self::Output {
241                let res = (&self.inner).$method(&rhs.inner);
242                SpVec::new(res)
243            }
244        }
245    };
246}
247
248impl_binop!(Add, add);
249impl_binop!(Sub, sub);
250
251// SpMat * SpVec
252#[auto_ops(val_val, val_ref, ref_val)]
253impl<R> Mul<&SpVec<R>> for &SpMat<R>
254where R: Ring, for<'x> &'x R: RingOps<R> {
255    type Output = SpVec<R>;
256    fn mul(self, rhs: &SpVec<R>) -> Self::Output {
257        let res = self.inner() * &rhs.inner;
258        SpVec::new(res)
259    }
260}
261
262impl<R> Display for SpVec<R>
263where R: Ring, for<'a> &'a R: RingOps<R> {
264    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
265        self.inner.fmt(f)
266    }
267}
268
269#[cfg(test)]
270mod tests {
271    use itertools::Itertools;
272    use super::*;
273
274    #[test]
275    fn from_vec() {
276        let v = SpVec::from(vec![1,0,3,5,0]);
277        assert_eq!(v.inner.disassemble(), (vec![0, 3], vec![0, 2, 3], vec![1, 3, 5]));
278    }
279
280    #[test]
281    fn from_entries() {
282        let v = SpVec::from_entries(5, vec![(0, 1), (4, 5), (2, 3)]);
283        assert_eq!(v.inner.disassemble(), (vec![0, 3], vec![0, 2, 4], vec![1, 3, 5]));
284    }
285
286    #[test]
287    fn to_dense() {
288        let v = SpVec::from(vec![1,0,3,5,0]);
289        assert_eq!(v.into_dense(), vec![1,0,3,5,0]);
290    }
291
292    #[test]
293    fn add() {
294        let v = SpVec::from(vec![1,0,3,5,0]);
295        let w = SpVec::from(vec![2,1,-1,3,2]);
296        assert_eq!(v + w, SpVec::from(vec![3,1,2,8,2]));
297    }
298
299    #[test]
300    fn sub() {
301        let v = SpVec::from(vec![1,0,3,5,0]);
302        let w = SpVec::from(vec![2,1,-1,3,2]);
303        assert_eq!(v - w, SpVec::from(vec![-1,-1,4,2,-2]));
304    }
305
306    #[test]
307    fn neg() {
308        let v = SpVec::from(vec![1,0,3,5,0]);
309        assert_eq!(-v, SpVec::from(vec![-1,0,-3,-5,0]));
310    }
311
312    #[test]
313    fn subvec() {
314        let v = SpVec::from((0..10).collect_vec());
315        let w = v.subvec(3..7);
316        assert_eq!(w, SpVec::from(vec![3,4,5,6]))
317    }
318
319    #[test]
320    fn subvec2() {
321        let v = SpVec::from((0..10).collect_vec());
322        let w = v.subvec(1..9);
323        let w = w.subvec(1..4);
324        assert_eq!(w, SpVec::from(vec![2,3,4]))
325    }
326
327    #[test]
328    fn permute() {
329        let p = Perm::new(vec![1,3,0,2]);
330        let v = SpVec::from(vec![0,1,2,3]);
331        let w = v.permute(&p);
332        assert_eq!(w, SpVec::from(vec![2,0,3,1]));
333    }
334
335    #[test]
336    fn stack() {
337        let v1 = SpVec::from((0..3).collect_vec());
338        let v2 = SpVec::from((5..8).collect_vec());
339        let w = SpVec::stack(v1, v2);
340        assert_eq!(w, SpVec::from(vec![0,1,2,5,6,7]));
341    }
342
343    #[test]
344    fn split() {
345        let v = SpVec::from((0..10).collect_vec());
346        let (x, y) = v.split(4);
347        assert_eq!(x, SpVec::from((0..4).collect_vec()));
348        assert_eq!(y, SpVec::from((4..10).collect_vec()));
349    }
350}