conspire/math/tensor/rank_2/sparse_vec_2d/
mod.rs1#[cfg(test)]
2mod test;
3
4use super::TensorRank2;
5use crate::math::{
6 Hessian, HessianAccumulate, HessianAccumulateGeneral, Rank2, Scalar, SquareMatrix, Tensor,
7 tensor::vec::TensorVector,
8};
9use std::ops::Mul;
10
11use super::sparse_vec::TensorRank2SparseVec;
12
13use crate::math::{TensorArray, TensorRank0, assert::FiniteDifference};
14
15pub type TensorRank2SparseVec2D<const D: usize, const I: usize, const J: usize> =
17 TensorVector<TensorRank2SparseVec<D, I, J>>;
18
19impl<const D: usize, const I: usize, const J: usize> TensorRank2SparseVec2D<D, I, J> {
20 pub fn zero(len: usize) -> Self {
21 (0..len).map(|_| TensorRank2SparseVec::default()).collect()
22 }
23}
24
25impl<const D: usize, const I: usize> HessianAccumulate<D, I> for TensorRank2SparseVec2D<D, I, I> {
26 fn accumulate(&mut self, a: usize, b: usize, block: TensorRank2<D, I, I>) {
27 if a == b {
28 self[a][b] += block;
29 } else {
30 self[b][a] += block.transpose();
31 self[a][b] += block;
32 }
33 }
34}
35
36impl<const D: usize, const I: usize> HessianAccumulateGeneral<D, I>
37 for TensorRank2SparseVec2D<D, I, I>
38{
39 fn accumulate_general(&mut self, a: usize, b: usize, block: TensorRank2<D, I, I>) {
40 self[a][b] += block;
41 }
42}
43
44impl<const D: usize, const I: usize, const J: usize, const K: usize>
45 Mul<TensorRank2SparseVec2D<D, J, K>> for TensorRank2<D, I, J>
46{
47 type Output = TensorRank2SparseVec2D<D, I, K>;
48 fn mul(self, tensor_rank_2_sparse_vec_2d: TensorRank2SparseVec2D<D, J, K>) -> Self::Output {
49 tensor_rank_2_sparse_vec_2d
50 .into_iter()
51 .map(|row| {
52 TensorRank2SparseVec(
53 row.0
54 .into_iter()
55 .map(|(column, block)| (column, &self * block))
56 .collect(),
57 )
58 })
59 .collect()
60 }
61}
62
63impl<const D: usize, const I: usize, const J: usize, const K: usize> Mul<TensorRank2<D, J, K>>
64 for TensorRank2SparseVec2D<D, I, J>
65{
66 type Output = TensorRank2SparseVec2D<D, I, K>;
67 fn mul(self, tensor_rank_2: TensorRank2<D, J, K>) -> Self::Output {
68 self.into_iter()
69 .map(|row| {
70 TensorRank2SparseVec(
71 row.0
72 .into_iter()
73 .map(|(column, block)| (column, block * &tensor_rank_2))
74 .collect(),
75 )
76 })
77 .collect()
78 }
79}
80
81impl<const D: usize, const I: usize, const J: usize> Hessian for TensorRank2SparseVec2D<D, I, J> {
82 fn entry(&self, row: usize, column: usize) -> Scalar {
83 match self[row / D]
84 .0
85 .binary_search_by_key(&(column / D), |&(b, _)| b)
86 {
87 Ok(k) => self[row / D].0[k].1[row % D][column % D],
88 Err(_) => 0.0,
89 }
90 }
91 fn fill_into(self, square_matrix: &mut SquareMatrix) {
92 self.iter().enumerate().for_each(|(a, row)| {
93 row.entries().for_each(|(b, block)| {
94 block.iter().enumerate().for_each(|(i, block_i)| {
95 block_i
96 .iter()
97 .enumerate()
98 .for_each(|(j, block_ij)| square_matrix[D * a + i][D * b + j] = *block_ij)
99 })
100 })
101 });
102 }
103 fn retain_from(self, retained: &[bool]) -> SquareMatrix {
104 let mut remap = vec![0; retained.len()];
105 let mut count = 0;
106 retained.iter().enumerate().for_each(|(p, &keep)| {
107 if keep {
108 remap[p] = count;
109 count += 1;
110 }
111 });
112 let mut square_matrix = SquareMatrix::zero(count);
113 self.iter().enumerate().for_each(|(a, row)| {
114 row.entries().for_each(|(b, block)| {
115 block.iter().enumerate().for_each(|(i, block_i)| {
116 block_i.iter().enumerate().for_each(|(j, block_ij)| {
117 if retained[D * a + i] && retained[D * b + j] {
118 square_matrix[remap[D * a + i]][remap[D * b + j]] = *block_ij
119 }
120 })
121 })
122 })
123 });
124 square_matrix
125 }
126}
127
128impl<const D: usize, const I: usize, const J: usize> FiniteDifference
129 for TensorRank2SparseVec2D<D, I, J>
130{
131 fn error_fd(&self, comparator: &Self, epsilon: TensorRank0) -> Option<(bool, usize)> {
132 let zero = TensorRank2::zero();
133 let block_errors =
134 |self_ab: &TensorRank2<D, I, J>, comparator_ab: &TensorRank2<D, I, J>| {
135 let mut errors = (0, 0);
136 self_ab.iter().zip(comparator_ab.iter()).for_each(
137 |(self_ab_i, comparator_ab_i)| {
138 self_ab_i.iter().zip(comparator_ab_i.iter()).for_each(
139 |(&self_ab_ij, &comparator_ab_ij)| {
140 if (self_ab_ij / comparator_ab_ij - 1.0).abs() >= epsilon
141 && (self_ab_ij.abs() >= epsilon
142 || comparator_ab_ij.abs() >= epsilon)
143 {
144 errors.0 += 1;
145 if (self_ab_ij - comparator_ab_ij).abs() >= epsilon {
146 errors.1 += 1;
147 }
148 }
149 },
150 )
151 },
152 );
153 errors
154 };
155 let (error_count, severe_count) = self
156 .iter()
157 .zip(comparator.iter())
158 .map(|(self_a, comparator_a)| {
159 let mut errors = (0, 0);
160 let (mut p, mut q) = (0, 0);
161 while p < self_a.0.len() || q < comparator_a.0.len() {
162 let b = self_a.0.get(p).map(|&(b, _)| b);
163 let c = comparator_a.0.get(q).map(|&(c, _)| c);
164 let block = match (b, c) {
165 (Some(b), Some(c)) if b == c => {
166 p += 1;
167 q += 1;
168 block_errors(&self_a.0[p - 1].1, &comparator_a.0[q - 1].1)
169 }
170 (Some(b), Some(c)) if b < c => {
171 p += 1;
172 block_errors(&self_a.0[p - 1].1, &zero)
173 }
174 (Some(_), None) => {
175 p += 1;
176 block_errors(&self_a.0[p - 1].1, &zero)
177 }
178 _ => {
179 q += 1;
180 block_errors(&zero, &comparator_a.0[q - 1].1)
181 }
182 };
183 errors.0 += block.0;
184 errors.1 += block.1;
185 }
186 errors
187 })
188 .fold((0, 0), |sum, errors| (sum.0 + errors.0, sum.1 + errors.1));
189 if error_count > 0 {
190 Some((severe_count > 0, error_count))
191 } else {
192 None
193 }
194 }
195}