Skip to main content

conspire/math/tensor/rank_2/sparse_vec_2d/
mod.rs

1#[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
15/// A vector of sparse vectors of rank-2 tensors, storing only inserted entries.
16pub 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}