1use 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#[derive(Clone, Debug)]
16pub struct SpVec<R> {
17 inner: CscMatrix<R> }
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
104impl<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 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#[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}