commit 6ec82c19469932fd5edd96eab98888708331ed70
parent 547b48dba0d910049c1ebb78aceed69029b96d29
Author: David Freifeld <freifeld.david@gmail.com>
Date: Fri, 10 Jul 2020 15:09:58 -0700
Working L2 regularization in backprop?
Diffstat:
3 files changed, 8 insertions(+), 7 deletions(-)
diff --git a/bpnn.cpp b/bpnn.cpp
@@ -40,8 +40,9 @@ void Layer::init_weights(Layer next)
}
}
-Network::Network(char* path, int batch_sz, float learn_rate, float bias_rate, float ratio)
+Network::Network(char* path, int batch_sz, float learn_rate, float bias_rate, float l, float ratio)
{
+ lambda = l;
learning_rate = learn_rate;
bias_lr = bias_rate;
int total_instances = prep_file(path, SHUFFLED_PATH);
@@ -181,7 +182,7 @@ void Network::backpropagate()
counter++;
}
for (int i = 0; i < length-1; i++) {
- *layers[length-2-i].weights -= (learning_rate * deltas[i]) + (learning_rate * deltas[i]);
+ *layers[length-2-i].weights -= (learning_rate * deltas[i]) + (learning_rate * (lambda/batch_size) * *layers[length-2-i].weights);
*layers[length-1-i].bias -= bias_lr * gradients[i];
}
}
@@ -322,7 +323,7 @@ void Network::train()
epoch_acc = 1.0/((float) instances/batch_size) * acc_sum;
epoch_cost = 1.0/((float) instances/batch_size) * cost_sum;
test(TEST_PATH);
- // printf("Epoch complete - cost %f - acc %f - val_cost %f - val_acc %f\n", epoch_cost, epoch_acc, val_cost, val_acc);
+ printf("Epoch complete - cost %f - acc %f - val_cost %f - val_acc %f\n", epoch_cost, epoch_acc, val_cost, val_acc);
batches=1;
rewind(data);
}
diff --git a/bpnn.hpp b/bpnn.hpp
@@ -52,7 +52,7 @@ public:
Eigen::MatrixXf* labels;
- Network(char* path, int batch_sz, float learn_rate, float bias_rate, float ratio);
+ Network(char* path, int batch_sz, float learn_rate, float bias_rate, float l, float ratio);
void add_layer(int nodes, char* activation);
void initialize();
void update_layer(float* vals, int datalen, int index);
diff --git a/example.cpp b/example.cpp
@@ -6,7 +6,7 @@
double bench(int batch_sz)
{
auto start = std::chrono::high_resolution_clock::now();
- Network net ("./data_banknote_authentication.txt", batch_sz, 0.0155, 0.03, 0.9);
+ Network net ("./data_banknote_authentication.txt", batch_sz, 0.0155, 0.03, 0.0, 0.9);
net.add_layer(4, "linear");
net.add_layer(5, "relu");
net.add_layer(1, "resig");
@@ -16,13 +16,13 @@ double bench(int batch_sz)
net.train();
}
auto end = std::chrono::high_resolution_clock::now();
- // net.list_net();
+ net.list_net();
return std::chrono::duration_cast<std::chrono::nanoseconds>(end - start).count() / pow(10,9);
}
int main()
{
- std::cout << bench(50) << "\n";
+ std::cout << bench(16) << "\n";
// bench(50);
// bench(50);
// bench(50);