-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathaccumulator.hpp
More file actions
74 lines (57 loc) · 1.46 KB
/
Copy pathaccumulator.hpp
File metadata and controls
74 lines (57 loc) · 1.46 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
#ifndef ACCUMULATOR_HPP
#define ACCUMULATOR_HPP
#include "svd.hpp"
#include <iostream>
class Accumulator {
using SVDMatrix = SVDHelper;
SVDMatrix svd;
double total_logdet;
double current_logdet;
double dist;
typedef typename SVDMatrix::Matrix Matrix;
public:
class AssertionFailed {};
void start (const Matrix& M) {
svd.inPlaceSVD(M);
total_logdet = 0.0;
current_logdet = 0.0;
dist = 0.0;
}
void reset (size_t V) {
svd.setIdentity(V);
total_logdet = 0.0;
current_logdet = 0.0;
dist = 0.0;
}
Matrix &matrixU () { return svd.U; }
Matrix &matrixVt () { return svd.Vt; }
double logdet () const { return svd.S.array().log().sum(); }
double distance () const { return dist; }
void increase_logdet (double x) {
total_logdet += x;
current_logdet += x;
}
void increase_distance (double d) {
dist += d;
}
void decomposeU () {
svd.absorbU();
current_logdet = 0.0;
dist = 0.0;
}
bool testLogDet (double prec = 1.0e-6) const {
return std::fabs(svd.S.array().log().sum()-total_logdet)<prec;
}
double logDetError () const {
return std::fabs(svd.S.array().log().sum()-total_logdet);
}
void assertLogDet (double prec = 1.0e-6) const {
if (!testLogDet(prec)) {
std::cerr << svd.S.array().log().sum() << '-' << total_logdet << '=' << svd.S.array().log().sum()-total_logdet << std::endl;
throw AssertionFailed();
}
}
const SVDMatrix& SVD () const { return svd; }
SVDMatrix& SVD () { return svd; }
};
#endif // ACCUMULATOR_HPP