1#[cfg(test)]
2mod test;
3
4use super::TensorRank2;
5use crate::math::{Tensor, TensorArray, TensorRank0};
6use std::{
7 fmt::{Display, Formatter, Result},
8 iter::Sum,
9 ops::{Add, AddAssign, Div, DivAssign, Index, IndexMut, Mul, MulAssign, Sub, SubAssign},
10};
11
12#[derive(Clone, Debug, Default, PartialEq)]
14pub struct TensorRank2SparseVec<const D: usize, const I: usize, const J: usize>(
15 pub(super) Vec<(usize, TensorRank2<D, I, J>)>,
16);
17
18impl<const D: usize, const I: usize, const J: usize> TensorRank2SparseVec<D, I, J> {
19 pub fn entries(&self) -> impl Iterator<Item = (usize, &TensorRank2<D, I, J>)> {
20 self.0.iter().map(|(column, entry)| (*column, entry))
21 }
22}
23
24impl<const D: usize, const I: usize, const J: usize> FromIterator<TensorRank2<D, I, J>>
25 for TensorRank2SparseVec<D, I, J>
26{
27 fn from_iter<T>(into_iterator: T) -> Self
28 where
29 T: IntoIterator<Item = TensorRank2<D, I, J>>,
30 {
31 Self(into_iterator.into_iter().enumerate().collect())
32 }
33}
34
35impl<const D: usize, const I: usize, const J: usize> Index<usize>
36 for TensorRank2SparseVec<D, I, J>
37{
38 type Output = TensorRank2<D, I, J>;
39 fn index(&self, index: usize) -> &Self::Output {
40 match self.0.binary_search_by_key(&index, |&(column, _)| column) {
41 Ok(k) => &self.0[k].1,
42 Err(_) => panic!("Entry ({index}) not present."),
43 }
44 }
45}
46
47impl<const D: usize, const I: usize, const J: usize> IndexMut<usize>
48 for TensorRank2SparseVec<D, I, J>
49{
50 fn index_mut(&mut self, index: usize) -> &mut Self::Output {
51 let k = match self.0.binary_search_by_key(&index, |&(column, _)| column) {
52 Ok(k) => k,
53 Err(k) => {
54 self.0.insert(k, (index, TensorRank2::zero()));
55 k
56 }
57 };
58 &mut self.0[k].1
59 }
60}
61
62impl<const D: usize, const I: usize, const J: usize> Display for TensorRank2SparseVec<D, I, J> {
63 fn fmt(&self, f: &mut Formatter) -> Result {
64 write!(f, "Need to implement Display")
65 }
66}
67
68impl<const D: usize, const I: usize, const J: usize> Tensor for TensorRank2SparseVec<D, I, J> {
69 type Item = TensorRank2<D, I, J>;
70 fn iter(&self) -> impl Iterator<Item = &Self::Item> {
71 self.0.iter().map(|(_, entry)| entry)
72 }
73 fn iter_mut(&mut self) -> impl Iterator<Item = &mut Self::Item> {
74 self.0.iter_mut().map(|(_, entry)| entry)
75 }
76 fn len(&self) -> usize {
77 self.0.len()
78 }
79 fn size(&self) -> usize {
80 self.0.len() * D * D
81 }
82}
83
84fn merge<const D: usize, const I: usize, const J: usize>(
85 a: TensorRank2SparseVec<D, I, J>,
86 b: &TensorRank2SparseVec<D, I, J>,
87 sign: TensorRank0,
88) -> TensorRank2SparseVec<D, I, J> {
89 let mut merged = a;
90 b.0.iter()
91 .for_each(|(column, entry)| merged[*column] += entry * sign);
92 merged
93}
94
95impl<const D: usize, const I: usize, const J: usize> Add for TensorRank2SparseVec<D, I, J> {
96 type Output = Self;
97 fn add(self, other: Self) -> Self {
98 merge(self, &other, 1.0)
99 }
100}
101
102impl<const D: usize, const I: usize, const J: usize> Add<&Self> for TensorRank2SparseVec<D, I, J> {
103 type Output = Self;
104 fn add(self, other: &Self) -> Self {
105 merge(self, other, 1.0)
106 }
107}
108
109impl<const D: usize, const I: usize, const J: usize> AddAssign for TensorRank2SparseVec<D, I, J> {
110 fn add_assign(&mut self, other: Self) {
111 other
112 .0
113 .into_iter()
114 .for_each(|(column, entry)| self[column] += entry);
115 }
116}
117
118impl<const D: usize, const I: usize, const J: usize> AddAssign<&Self>
119 for TensorRank2SparseVec<D, I, J>
120{
121 fn add_assign(&mut self, other: &Self) {
122 other
123 .0
124 .iter()
125 .for_each(|(column, entry)| self[*column] += entry);
126 }
127}
128
129impl<const D: usize, const I: usize, const J: usize> Sub for TensorRank2SparseVec<D, I, J> {
130 type Output = Self;
131 fn sub(self, other: Self) -> Self {
132 merge(self, &other, -1.0)
133 }
134}
135
136impl<const D: usize, const I: usize, const J: usize> Sub<&Self> for TensorRank2SparseVec<D, I, J> {
137 type Output = Self;
138 fn sub(self, other: &Self) -> Self {
139 merge(self, other, -1.0)
140 }
141}
142
143impl<const D: usize, const I: usize, const J: usize> SubAssign for TensorRank2SparseVec<D, I, J> {
144 fn sub_assign(&mut self, other: Self) {
145 other
146 .0
147 .into_iter()
148 .for_each(|(column, entry)| self[column] -= entry);
149 }
150}
151
152impl<const D: usize, const I: usize, const J: usize> SubAssign<&Self>
153 for TensorRank2SparseVec<D, I, J>
154{
155 fn sub_assign(&mut self, other: &Self) {
156 other
157 .0
158 .iter()
159 .for_each(|(column, entry)| self[*column] -= entry);
160 }
161}
162
163impl<const D: usize, const I: usize, const J: usize> Mul<TensorRank0>
164 for TensorRank2SparseVec<D, I, J>
165{
166 type Output = Self;
167 fn mul(mut self, scalar: TensorRank0) -> Self {
168 self *= &scalar;
169 self
170 }
171}
172
173impl<const D: usize, const I: usize, const J: usize> MulAssign<TensorRank0>
174 for TensorRank2SparseVec<D, I, J>
175{
176 fn mul_assign(&mut self, scalar: TensorRank0) {
177 self.0.iter_mut().for_each(|(_, entry)| *entry *= &scalar);
178 }
179}
180
181impl<const D: usize, const I: usize, const J: usize> MulAssign<&TensorRank0>
182 for TensorRank2SparseVec<D, I, J>
183{
184 fn mul_assign(&mut self, scalar: &TensorRank0) {
185 self.0.iter_mut().for_each(|(_, entry)| *entry *= scalar);
186 }
187}
188
189impl<const D: usize, const I: usize, const J: usize> Div<TensorRank0>
190 for TensorRank2SparseVec<D, I, J>
191{
192 type Output = Self;
193 fn div(mut self, scalar: TensorRank0) -> Self {
194 self /= &scalar;
195 self
196 }
197}
198
199impl<const D: usize, const I: usize, const J: usize> DivAssign<TensorRank0>
200 for TensorRank2SparseVec<D, I, J>
201{
202 fn div_assign(&mut self, scalar: TensorRank0) {
203 self.0.iter_mut().for_each(|(_, entry)| *entry /= &scalar);
204 }
205}
206
207impl<const D: usize, const I: usize, const J: usize> DivAssign<&TensorRank0>
208 for TensorRank2SparseVec<D, I, J>
209{
210 fn div_assign(&mut self, scalar: &TensorRank0) {
211 self.0.iter_mut().for_each(|(_, entry)| *entry /= scalar);
212 }
213}
214
215impl<const D: usize, const I: usize, const J: usize> Sum for TensorRank2SparseVec<D, I, J> {
216 fn sum<T>(iter: T) -> Self
217 where
218 T: Iterator<Item = Self>,
219 {
220 iter.fold(Self::default(), |sum, entry| sum + entry)
221 }
222}