MCPcopy Create free account
hub / github.com/davisking/dlib / test_rls

Function test_rls

dlib/test/rls.cpp:24–172  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

22
23
24 void test_rls()
25 {
26 dlib::rand rnd;
27
28 running_stats<double> rs1, rs2, rs3, rs4, rs5;
29
30 for (int k = 0; k < 2; ++k)
31 {
32 for (long num_vars = 1; num_vars < 4; ++num_vars)
33 {
34 print_spinner();
35 for (long size = 1; size < 300; ++size)
36 {
37 {
38 matrix<double> X = randm(size,num_vars,rnd);
39 matrix<double,0,1> Y = randm(size,1,rnd);
40
41
42 const double C = 1000;
43 const double forget_factor = 1.0;
44 rls r(forget_factor, C);
45 for (long i = 0; i < Y.size(); ++i)
46 {
47 r.train(trans(rowm(X,i)), Y(i));
48 }
49
50
51 matrix<double> w = pinv(1.0/C*identity_matrix<double>(X.nc()) + trans(X)*X)*trans(X)*Y;
52
53 rs1.add(length(r.get_w() - w));
54 }
55
56 {
57 matrix<double> X = randm(size,num_vars,rnd);
58 matrix<double,0,1> Y = randm(size,1,rnd);
59
60 matrix<double,0,1> G(size,1);
61
62 const double C = 10000;
63 const double forget_factor = 0.8;
64 rls r(forget_factor, C);
65 for (long i = 0; i < Y.size(); ++i)
66 {
67 r.train(trans(rowm(X,i)), Y(i));
68
69 G(i) = std::pow(forget_factor, i/2.0);
70 }
71
72 G = flipud(G);
73
74 X = diagm(G)*X;
75 Y = diagm(G)*Y;
76
77 matrix<double> w = pinv(1.0/C*identity_matrix<double>(X.nc()) + trans(X)*X)*trans(X)*Y;
78
79 rs5.add(length(r.get_w() - w));
80 }
81

Callers 1

perform_testMethod · 0.85

Calls 15

print_spinnerFunction · 0.85
pinvFunction · 0.85
flipudFunction · 0.85
diagmFunction · 0.85
join_rowsFunction · 0.85
absFunction · 0.85
randmFunction · 0.70
transFunction · 0.50
rowmFunction · 0.50
lengthFunction · 0.50
GFunction · 0.50
colmFunction · 0.50

Tested by

no test coverage detected