jacobian

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

commit af578071cf826a2cd819e101000342a9ae9ccbf9
parent 5600f7366e406638d11c850f480434803fd60b6e
Author: David Freifeld <freifeld.david@gmail.com>
Date:   Sat,  1 Aug 2020 12:15:12 -0700

Added AdaMax + found adam bugs

Diffstat:
Msrc/bpnn.cpp | 19++++++++++++++++---
1 file changed, 16 insertions(+), 3 deletions(-)

diff --git a/src/bpnn.cpp b/src/bpnn.cpp @@ -108,9 +108,22 @@ void Network::init_optimizer(char* name, ...) va_end(args); // TODO: Add bias correction (requires figuring out measuring t) update = [this, beta1, beta2, epsilon](std::vector<Eigen::MatrixXf> deltas, int i) { - *layers[length-2-i].weights -= (layers[length-2-i].m * learning_rate * (layers[length-2-i].v->cwiseSqrt() + epsilon).array().pow(-1).matrix() - *layers[length-2-i].m = (beta1 * m) + (1-beta1)*deltas[i]; - *layers[length-2-i].v = (beta2 * v) + (1-beta2)*(deltas[i].cwiseProduct(deltas[i])); + *layers[length-2-i].weights -= *layers[length-2-i].m * learning_rate * ((layers[length-2-i].v->cwiseSqrt()).array()+epsilon).pow(-1).matrix(); + *layers[length-2-i].m = (beta1 * *layers[length-2-i].m) + ((1-beta1)*deltas[i]); + *layers[length-2-i].v = (beta2 * *layers[length-2-i].v) + (1-beta2)*(deltas[i].cwiseProduct(deltas[i])); + }; + } + if (strcmp(name, "adamax") == 0) { + float beta1 = va_arg(args, double); + float beta2 = va_arg(args, double); + float epsilon = va_arg(args, double); + va_end(args); + // TODO: Add bias correction for m (requires figuring out measuring t) + update = [this, beta1, beta2, epsilon](std::vector<Eigen::MatrixXf> deltas, int i) { + *layers[length-2-i].weights -= *layers[length-2-i].m * learning_rate * layers[length-2-i].v->array().pow(-1).matrix(); + *layers[length-2-i].m = (beta1 * *layers[length-2-i].m) + ((1-beta1)*deltas[i]); + if ((beta2 * *layers[length-2-i].v) > deltas[i].array().abs().matrix()) *layers[length-2-i].v = (beta2 * *layers[length-2-i].v); + else *layers[length-2-i].v = deltas[i].array().abs().matrix(); }; } }