jacobian

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

commit 787493f9d876e941eec624369e7ea61cd037160d
parent 634b760f035671ba2da44f15a608729585b01967
Author: David Freifeld <freifeld.david@gmail.com>
Date:   Fri, 26 Jun 2020 19:41:18 -0700

Experiments with SGD w/ momentum

Diffstat:
Mbpnn.cpp | 3++-
Mbpnn.hpp | 1+
Mexample.cpp | 17++++++++---------
3 files changed, 11 insertions(+), 10 deletions(-)

diff --git a/bpnn.cpp b/bpnn.cpp @@ -32,6 +32,7 @@ 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) @@ -184,7 +185,7 @@ 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]; + *layers[length-2-i].weights -= 0.5 * *layers[length-2-i].prev_update + learning_rate * deltas[i]; *layers[length-1-i].bias -= bias_lr * gradients[i]; } } diff --git a/bpnn.hpp b/bpnn.hpp @@ -20,6 +20,7 @@ public: Eigen::MatrixXd* weights; Eigen::MatrixXd* bias; Eigen::MatrixXd* dZ; + Eigen::MatrixXd* prev_update; std::function<double(double)> activation; std::function<double(double)> activation_deriv; diff --git a/example.cpp b/example.cpp @@ -3,14 +3,13 @@ int main() { - // Network net ("./data_banknote_authentication.txt", 10, 0.01, 0.001); - // net.add_layer(4, "linear"); - // net.add_layer(5, "sigmoid"); - // net.set_activation(1, lecun_tanh, lecun_tanh_deriv); - // net.add_layer(1, "resig"); - // net.initialize(); - // net.list_net(); - // net.train(50); - printf("%d\n", wc("./bpnn.cpp")); + Network net ("./data_banknote_authentication.txt", 10, 0.01, 0.001); + net.add_layer(4, "linear"); + net.add_layer(5, "sigmoid"); + net.set_activation(1, lecun_tanh, lecun_tanh_deriv); + net.add_layer(1, "resig"); + net.initialize(); + // net.list_net(); + net.train(50); // net.list_net(); }