commit ca9855fa11f74b94bf48328bcd1f8c1ed150f3ae
parent 4363c18d8d6229c5654605a0f7c5235253e74937
Author: David Freifeld <freifeld.david@gmail.com>
Date: Mon, 15 Jun 2020 18:35:50 -0700
New simpler backprop behaves strangely
Diffstat:
| M | test.cpp | | | 62 | ++++++++++++++++++++++++++++---------------------------------- |
1 file changed, 28 insertions(+), 34 deletions(-)
diff --git a/test.cpp b/test.cpp
@@ -153,35 +153,23 @@ float Network::cost()
void Network::backpropagate()
{
- int N = batch_size;
-
+ // std::cout << "\nROUND\n\n\n\n\n\n";
std::vector<Eigen::MatrixXd> gradients;
- std::vector<Eigen::MatrixXd> errors;
- gradients.push_back(((*layers[length-1].contents ) - (*labels)).cwiseProduct(*layers[length-1].dZ));
- // std::cout << D << "\n\nTHEN\n\n" << layers[length-2].contents->transpose() << "\n\nNEXT\n\n" << e << "\n\nSO\n\n" << gradients[0] << "\n\n\n\n\n";
- // int counter = 0;
- // for (int i = length-2; i >= 1; i--) {
- // Eigen::MatrixXd D_l (layers[i].contents->cols(), layers[i].contents->cols());
- // for (int j = 0; j < layers[i].contents->cols(); j++) {
- // for (int k = 0; k < layers[i].contents->cols(); k++) {
- // D_l(k, j) = 0;
- // }
- // }
- // for (int j = 0; j < layers[i].contents->cols(); j++) {
- // D_l(j, j) = (*layers[i].contents)(0,j) * (1 - (*layers[i].contents)(0,j));
- // }
- // // std::cout << D_l << "\n\nTHEN\n\n" << layers[i].weights->transpose() << "\n\nNEXT\n\n" << gradients[counter] << "\n\n";
-
- // Eigen::MatrixXd e_l = D_l * ( gradients[counter] * layers[i].weights->transpose());
- // // std::cout << "\n\nSO\n\n" << e_l << "\n\n\n\n\n";
- // gradients.push_back(e_l);
- // counter++;
- // }
- // for (int i = 1; i < gradients.size(); i++) {
- // Eigen::MatrixXd gradient = gradients[i];
- // // std::cout << *layers[length-2-i].weights << " \n\n and \n\n " << gradients[i] << "\n\n";
- // *layers[length-2-i].weights -= learning_rate * (1.0/N * gradients[i]);
- // }
+ std::vector<Eigen::MatrixXd> deltas;
+ gradients.push_back(((*layers[length-1].contents) - (*labels)).cwiseProduct(*layers[length-1].dZ));
+ deltas.push_back((*layers[length-2].contents).transpose() * gradients[0]);
+ int counter = 1;
+ for (int i = length-2; i >= 1; i--) {
+ // std::cout << gradients[counter-1] << "\n\nTHAT WAS GRADIENT\n\n" <<*layers[i].weights << "\n\nTHAT WAS WEIGHTS\n"
+ gradients.push_back((gradients[counter-1] * layers[i].weights->transpose()).cwiseProduct(*layers[i].dZ));
+ deltas.push_back((*layers[i-1].contents).transpose() * gradients[counter]);
+ // std::cout << gradients[counter] << "\n\nand\n\n" << *layers[i-1].weights << "\n\nweights\n\n" << deltas[counter] << "\n\ndelta above\n\n\n\n\n";
+ counter++;
+ }
+ for (int i = 1; i < gradients.size(); i++) {
+ Eigen::MatrixXd gradient = gradients[i];
+ *layers[length-2-i].weights -= learning_rate * (deltas[i]);
+ }
}
void Network::update_layer(float* vals, int datalen, int index)
@@ -277,16 +265,22 @@ void demo()
{
// std::cout << "\n\n\n";
int linecount = prep_file("./data_banknote_authentication.txt");
- Network net ("./shuffled.txt", 4, 2, 1, 5, 10, 1);
+ Network net ("./shuffled.txt", 4, 2, 1, 5, 10, 0.5);
float epoch_cost = 1000;
int epochs = 0;
net.batches= 1;
- while (epochs < 1) {
- int linecount = prep_file("./data_banknote_authentication.txt");
+ // for (int i = 0; i < 500; i++) {
+ // net.feedforward();
+ // net.backpropagate();
+ // std::cout << net.cost() << "\n";
+ // }
+
+ while (epochs < 50) {
+ // int linecount = prep_file("./data_banknote_authentication.txt");
float cost_sum = 0;
- for (int i = 0; i < linecount-net.batch_size; i++) {
+ for (int i = 0; i < linecount-net.batch_size; i+=net.batch_size) {
net.feedforward();
- // net.backpropagate();
+ net.backpropagate();
cost_sum += net.cost();
// std::cout << net.cost() << " as it is " << net.labels[0] << " vs " << *net.layers[net.length-1].contents << "\n";
net.batches++;
@@ -300,7 +294,7 @@ void demo()
printf("EPOCH %i: Cost is %f for %i instances.\n", epochs, epoch_cost, linecount);
epochs++;
}
- // net.list_net();
+ net.list_net();
net.test("./test.txt");
net.feedforward();
}