1#[cfg(test)]
2mod test;
3
4use super::TensorRank2;
5use crate::math::{Hessian, HessianAccumulate, Rank2, Scalar, SquareMatrix, Tensor, TensorRank0};
6use std::{
7 fmt::{self, Display, Formatter},
8 iter::Sum,
9 ops::{Add, AddAssign, Div, DivAssign, Mul, MulAssign, Sub, SubAssign},
10};
11
12use super::sparse_vec::TensorRank2SparseVec;
13use super::sparse_vec_2d::TensorRank2SparseVec2D;
14
15use crate::math::{TensorArray, assert::FiniteDifference};
16
17#[derive(Clone, Debug, Default, PartialEq)]
25pub struct TensorRank2SparseVec2DSymmetric<const D: usize, const I: usize, const J: usize>(
26 TensorRank2SparseVec2D<D, I, J>,
27);
28
29impl<const D: usize, const I: usize, const J: usize> TensorRank2SparseVec2DSymmetric<D, I, J> {
30 pub fn zero(len: usize) -> Self {
31 Self(TensorRank2SparseVec2D::zero(len))
32 }
33}
34
35impl<const D: usize, const I: usize, const J: usize> Display
36 for TensorRank2SparseVec2DSymmetric<D, I, J>
37{
38 fn fmt(&self, f: &mut Formatter) -> fmt::Result {
39 write!(f, "Need to implement Display")
40 }
41}
42
43impl<const D: usize, const I: usize, const J: usize> Tensor
44 for TensorRank2SparseVec2DSymmetric<D, I, J>
45{
46 type Item = TensorRank2SparseVec<D, I, J>;
47 fn iter(&self) -> impl Iterator<Item = &Self::Item> {
48 self.0.iter()
49 }
50 fn iter_mut(&mut self) -> impl Iterator<Item = &mut Self::Item> {
51 self.0.iter_mut()
52 }
53 fn len(&self) -> usize {
54 self.0.len()
55 }
56 fn size(&self) -> usize {
57 self.0.size()
58 }
59}
60
61impl<const D: usize, const I: usize, const J: usize> Add
62 for TensorRank2SparseVec2DSymmetric<D, I, J>
63{
64 type Output = Self;
65 fn add(self, other: Self) -> Self {
66 Self(self.0 + other.0)
67 }
68}
69
70impl<const D: usize, const I: usize, const J: usize> Add<&Self>
71 for TensorRank2SparseVec2DSymmetric<D, I, J>
72{
73 type Output = Self;
74 fn add(self, other: &Self) -> Self {
75 Self(self.0 + &other.0)
76 }
77}
78
79impl<const D: usize, const I: usize, const J: usize> AddAssign
80 for TensorRank2SparseVec2DSymmetric<D, I, J>
81{
82 fn add_assign(&mut self, other: Self) {
83 self.0 += other.0;
84 }
85}
86
87impl<const D: usize, const I: usize, const J: usize> AddAssign<&Self>
88 for TensorRank2SparseVec2DSymmetric<D, I, J>
89{
90 fn add_assign(&mut self, other: &Self) {
91 self.0 += &other.0;
92 }
93}
94
95impl<const D: usize, const I: usize, const J: usize> Sub
96 for TensorRank2SparseVec2DSymmetric<D, I, J>
97{
98 type Output = Self;
99 fn sub(self, other: Self) -> Self {
100 Self(self.0 - other.0)
101 }
102}
103
104impl<const D: usize, const I: usize, const J: usize> Sub<&Self>
105 for TensorRank2SparseVec2DSymmetric<D, I, J>
106{
107 type Output = Self;
108 fn sub(self, other: &Self) -> Self {
109 Self(self.0 - &other.0)
110 }
111}
112
113impl<const D: usize, const I: usize, const J: usize> SubAssign
114 for TensorRank2SparseVec2DSymmetric<D, I, J>
115{
116 fn sub_assign(&mut self, other: Self) {
117 self.0 -= other.0;
118 }
119}
120
121impl<const D: usize, const I: usize, const J: usize> SubAssign<&Self>
122 for TensorRank2SparseVec2DSymmetric<D, I, J>
123{
124 fn sub_assign(&mut self, other: &Self) {
125 self.0 -= &other.0;
126 }
127}
128
129impl<const D: usize, const I: usize, const J: usize> Mul<TensorRank0>
130 for TensorRank2SparseVec2DSymmetric<D, I, J>
131{
132 type Output = Self;
133 fn mul(self, scalar: TensorRank0) -> Self {
134 Self(self.0 * scalar)
135 }
136}
137
138impl<const D: usize, const I: usize, const J: usize> MulAssign<TensorRank0>
139 for TensorRank2SparseVec2DSymmetric<D, I, J>
140{
141 fn mul_assign(&mut self, scalar: TensorRank0) {
142 self.0 *= scalar;
143 }
144}
145
146impl<const D: usize, const I: usize, const J: usize> MulAssign<&TensorRank0>
147 for TensorRank2SparseVec2DSymmetric<D, I, J>
148{
149 fn mul_assign(&mut self, scalar: &TensorRank0) {
150 self.0 *= scalar;
151 }
152}
153
154impl<const D: usize, const I: usize, const J: usize> Div<TensorRank0>
155 for TensorRank2SparseVec2DSymmetric<D, I, J>
156{
157 type Output = Self;
158 fn div(self, scalar: TensorRank0) -> Self {
159 Self(self.0 / scalar)
160 }
161}
162
163impl<const D: usize, const I: usize, const J: usize> DivAssign<TensorRank0>
164 for TensorRank2SparseVec2DSymmetric<D, I, J>
165{
166 fn div_assign(&mut self, scalar: TensorRank0) {
167 self.0 /= scalar;
168 }
169}
170
171impl<const D: usize, const I: usize, const J: usize> DivAssign<&TensorRank0>
172 for TensorRank2SparseVec2DSymmetric<D, I, J>
173{
174 fn div_assign(&mut self, scalar: &TensorRank0) {
175 self.0 /= scalar;
176 }
177}
178
179impl<const D: usize, const I: usize, const J: usize> Sum
180 for TensorRank2SparseVec2DSymmetric<D, I, J>
181{
182 fn sum<T>(iter: T) -> Self
183 where
184 T: Iterator<Item = Self>,
185 {
186 iter.fold(Self::default(), |sum, entry| sum + entry)
187 }
188}
189
190impl<const D: usize, const I: usize> HessianAccumulate<D, I>
191 for TensorRank2SparseVec2DSymmetric<D, I, I>
192{
193 fn accumulate(&mut self, a: usize, b: usize, block: TensorRank2<D, I, I>) {
194 if a <= b {
195 self.0[a][b] += block;
196 } else {
197 self.0[b][a] += block.transpose();
198 }
199 }
200}
201
202impl<const D: usize, const I: usize, const J: usize> Hessian
203 for TensorRank2SparseVec2DSymmetric<D, I, J>
204{
205 fn entry(&self, row: usize, column: usize) -> Scalar {
206 let (a, b, i, j) = if row / D <= column / D {
207 (row / D, column / D, row % D, column % D)
208 } else {
209 (column / D, row / D, column % D, row % D)
210 };
211 match self.0[a].0.binary_search_by_key(&b, |&(c, _)| c) {
212 Ok(k) => self.0[a].0[k].1[i][j],
213 Err(_) => 0.0,
214 }
215 }
216 fn fill_into(self, square_matrix: &mut SquareMatrix) {
217 self.0.iter().enumerate().for_each(|(a, row)| {
218 row.entries().for_each(|(b, block)| {
219 block.iter().enumerate().for_each(|(i, block_i)| {
220 block_i.iter().enumerate().for_each(|(j, block_ij)| {
221 square_matrix[D * a + i][D * b + j] = *block_ij;
222 if a != b {
223 square_matrix[D * b + j][D * a + i] = *block_ij;
224 }
225 })
226 })
227 })
228 });
229 }
230 fn retain_from(self, retained: &[bool]) -> SquareMatrix {
231 let mut remap = vec![0; retained.len()];
232 let mut count = 0;
233 retained.iter().enumerate().for_each(|(p, &keep)| {
234 if keep {
235 remap[p] = count;
236 count += 1;
237 }
238 });
239 let mut square_matrix = SquareMatrix::zero(count);
240 self.0.iter().enumerate().for_each(|(a, row)| {
241 row.entries().for_each(|(b, block)| {
242 block.iter().enumerate().for_each(|(i, block_i)| {
243 block_i.iter().enumerate().for_each(|(j, block_ij)| {
244 if retained[D * a + i] && retained[D * b + j] {
245 square_matrix[remap[D * a + i]][remap[D * b + j]] = *block_ij;
246 if a != b {
247 square_matrix[remap[D * b + j]][remap[D * a + i]] = *block_ij;
248 }
249 }
250 })
251 })
252 })
253 });
254 square_matrix
255 }
256}
257
258impl<const D: usize, const I: usize, const J: usize> FiniteDifference
259 for TensorRank2SparseVec2DSymmetric<D, I, J>
260{
261 fn error_fd(&self, comparator: &Self, epsilon: TensorRank0) -> Option<(bool, usize)> {
262 let zero = TensorRank2::zero();
263 let block_errors =
264 |self_ab: &TensorRank2<D, I, J>, comparator_ab: &TensorRank2<D, I, J>| {
265 let mut errors = (0, 0);
266 self_ab.iter().zip(comparator_ab.iter()).for_each(
267 |(self_ab_i, comparator_ab_i)| {
268 self_ab_i.iter().zip(comparator_ab_i.iter()).for_each(
269 |(&self_ab_ij, &comparator_ab_ij)| {
270 if (self_ab_ij / comparator_ab_ij - 1.0).abs() >= epsilon
271 && (self_ab_ij.abs() >= epsilon
272 || comparator_ab_ij.abs() >= epsilon)
273 {
274 errors.0 += 1;
275 if (self_ab_ij - comparator_ab_ij).abs() >= epsilon {
276 errors.1 += 1;
277 }
278 }
279 },
280 )
281 },
282 );
283 errors
284 };
285 let (error_count, severe_count) = self
286 .0
287 .iter()
288 .zip(comparator.0.iter())
289 .map(|(self_a, comparator_a)| {
290 let mut errors = (0, 0);
291 let (mut p, mut q) = (0, 0);
292 while p < self_a.0.len() || q < comparator_a.0.len() {
293 let b = self_a.0.get(p).map(|&(b, _)| b);
294 let c = comparator_a.0.get(q).map(|&(c, _)| c);
295 let block = match (b, c) {
296 (Some(b), Some(c)) if b == c => {
297 p += 1;
298 q += 1;
299 block_errors(&self_a.0[p - 1].1, &comparator_a.0[q - 1].1)
300 }
301 (Some(b), Some(c)) if b < c => {
302 p += 1;
303 block_errors(&self_a.0[p - 1].1, &zero)
304 }
305 (Some(_), None) => {
306 p += 1;
307 block_errors(&self_a.0[p - 1].1, &zero)
308 }
309 _ => {
310 q += 1;
311 block_errors(&zero, &comparator_a.0[q - 1].1)
312 }
313 };
314 errors.0 += block.0;
315 errors.1 += block.1;
316 }
317 errors
318 })
319 .fold((0, 0), |sum, errors| (sum.0 + errors.0, sum.1 + errors.1));
320 if error_count > 0 {
321 Some((severe_count > 0, error_count))
322 } else {
323 None
324 }
325 }
326}