jacobian

Unnamed repository; edit this file 'description' to name the repository.
Log | Files | Refs | README

commit 4032d3c1402eedf8412842dc8c7a25936794a302
parent f3466f7dd07d9be57b724d9b09ae6adb228f89ea
Author: David Freifeld <freifeld.david@gmail.com>
Date:   Sun, 28 Jun 2020 16:39:10 -0700

Experiments with Adagrad

Diffstat:
Mbpnn.cpp | 8++++++--
Mbpnn.hpp | 2+-
Mexample.cpp | 3+--
3 files changed, 8 insertions(+), 5 deletions(-)

diff --git a/bpnn.cpp b/bpnn.cpp @@ -32,7 +32,6 @@ void Layer::init_weights(Layer next) std::mt19937 gen(rd()); (*weights)((int)i / nodes, i%nodes) = d(gen); } - prev_update = new Eigen::MatrixXd (contents->cols(), next.contents->cols()); } Network::Network(char* path, int batch_sz, float learn_rate, float bias_rate) @@ -173,7 +172,12 @@ void Network::backpropagate() } for (int i = 0; i < length-1; i++) { Eigen::MatrixXd gradient = gradients[i]; - *layers[length-2-i].weights -= learning_rate * deltas[i]; + Eigen::MatrixXd sum (layers[length-2-i].weights->rows(), layers[length-2-i].weights->cols()); + for (int j = 0; j < layers[length-2-i].prev_updates.size(); j++) { + sum = sum + layers[length-2-i].prev_updates[j].cwiseProduct(layers[length-2-i].prev_updates[j]); + } + layers[length-2-i].prev_updates.emplace_back(((1/learning_rate) * sum.cwiseSqrt()).transpose() * deltas[i]); + *layers[length-2-i].weights -= layers[length-2-i].prev_updates[layers[length-2-i].prev_updates.size()]; *layers[length-1-i].bias -= bias_lr * gradients[i]; } } diff --git a/bpnn.hpp b/bpnn.hpp @@ -20,7 +20,7 @@ public: Eigen::MatrixXd* weights; Eigen::MatrixXd* bias; Eigen::MatrixXd* dZ; - Eigen::MatrixXd* prev_update; + std::vector<Eigen::MatrixXd> prev_updates; std::function<double(double)> activation; std::function<double(double)> activation_deriv; diff --git a/example.cpp b/example.cpp @@ -5,14 +5,13 @@ int main() { - sleep(30); Network net ("./data_banknote_authentication.txt", 10, 0.01, 0.001); net.add_layer(4, "linear"); net.add_layer(5, "lecun_tanh"); net.add_layer(1, "resig"); net.initialize(); // net.list_net(); - net.train(500); + net.train(50); net.list_net(); //printf("%i\n", wc("./data_banknote_authentication.txt"));