commit e2e47b932dbe0e1450a56a0bd1c6fe77d674e2a6
parent 43e2b13bd0546c2b75a1955422bece2c4b21ec06
Author: David Freifeld <freifeld.david@gmail.com>
Date: Sun, 12 Jul 2020 10:54:54 -0700
Loss+accuracy tentatively work now
Diffstat:
2 files changed, 15 insertions(+), 7 deletions(-)
diff --git a/example.cpp b/example.cpp
@@ -13,7 +13,9 @@ double bench(int batch_sz)
net.init_decay("exp", 1, 10);
net.initialize();
// checks(net);
+ net.next_batch();
net.feedforward();
+ std::cout << net.cost() << " " << net.accuracy() << "\n";
//for (int i = 0; i < 50; i++) {
// net.train();
//}
@@ -24,7 +26,7 @@ double bench(int batch_sz)
int main()
{
- std::cout << bench(4) << "\n";
+ std::cout << bench(16) << "\n";
// bench(50);
// bench(50);
// bench(50);
diff --git a/src/bpnn.cpp b/src/bpnn.cpp
@@ -190,11 +190,12 @@ float Network::cost()
float tempsum = 0;
for (int j = 0; j < layers[length-1].contents->cols(); j++) {
float truth;
- if (j==(*labels)(i,0)) truth = (*labels)(i,0);
+ if (j==(*labels)(i,0)) truth = 1;
else truth = 0;
- tempsum += truth * log((*layers[length-1].contents)(i,j))
+ // std::cout << truth << " VS " << (*layers[length-1].contents)(i,j) << " SO " << truth * log((*layers[length-1].contents)(i,j)) << "\n";
+ tempsum += truth * log((*layers[length-1].contents)(i,j));
}
- sum+=tempsum;
+ sum-=tempsum;
}
for (int i = 0; i < layers.size()-1; i++) {
reg += (layers[i].weights->cwiseProduct(*layers[i].weights)).sum();
@@ -207,11 +208,16 @@ float Network::accuracy()
float correct = 0;
for (int i = 0; i < layers[length-1].contents->rows(); i++) {
float ans = -INFINITY;
+ float index = -1;
for (int j = 0; j < layers[length-1].contents->cols(); j++) {
- if ((*layers[length-1].contents)(i, j) > ans) ans = j;
+ if ((*layers[length-1].contents)(i, j) > ans) {
+ ans = (*layers[length-1].contents)(i, j);
+ index = j;
+ //std::cout << "UPDATE ANS: " << index << " as "<< (*layers[length-1].contents)(i, j) << " so " << ans << "\n";
+ }
}
-
- if ((*labels)(i, 0) == ans) correct += 1;
+ //std::cout << (*labels)(i, 0) << " " << index << "\n";
+ if ((*labels)(i, 0) == index) correct += 1;
}
return (1.0/batch_size) * correct;
}