summaryrefslogtreecommitdiff
path: root/src/application.c
diff options
context:
space:
mode:
Diffstat (limited to 'src/application.c')
-rw-r--r--src/application.c82
1 files changed, 57 insertions, 25 deletions
diff --git a/src/application.c b/src/application.c
index b2105ce..6a868b3 100644
--- a/src/application.c
+++ b/src/application.c
@@ -5,6 +5,7 @@
#include <spires.h>
#include <stdio.h>
#include <stdlib.h>
+#include <string.h>
// spires reservoir parameters
#define NUM_NEURONS 400
@@ -20,12 +21,37 @@
#define PI 3.14159265358979323846
#define LAMBDA 1.0e-4
-#define NUM_TRAINING_STEPS 50
+#define NUM_TRAINING_STEPS 500
#define NUM_CROSSBAR_COLUMNS (NUM_OUTPUTS * 2)
// #define NUM_STEPS 2000
-int main(void)
+typedef enum {
+ BENCHMARK_ONLINE,
+ BENCHMARK_OFFLINE,
+} Benchmark_Mode;
+
+static int parse_mode(int argc, char **argv, Benchmark_Mode *mode)
+{
+ if (argc == 1 || (argc == 2 && strcmp(argv[1], "--online") == 0)) {
+ *mode = BENCHMARK_ONLINE;
+ return 0;
+ }
+ if (argc == 2 && strcmp(argv[1], "--offline") == 0) {
+ *mode = BENCHMARK_OFFLINE;
+ return 0;
+ }
+ fprintf(stderr, "Usage: %s [--online|--offline]\n", argv[0]);
+ return -1;
+}
+
+int main(int argc, char **argv)
{
+ Benchmark_Mode mode;
+ if (parse_mode(argc, argv, &mode) != 0)
+ return -1;
+ printf("Benchmark mode: %s\n",
+ mode == BENCHMARK_ONLINE ? "online" : "offline");
+
/* ---------- LIST ALL MODELS HERE ----------*/
MemModel models[] = {
{
@@ -110,25 +136,23 @@ int main(void)
return -1;
}
- /* ---------- Collect reservoir states ----------*/
Reservoir_State_Matrix state_matrix = {0};
- if (collect_reservoir_states(reservoir, training_inputs,
- NUM_TRAINING_STEPS, &state_matrix) != 0) {
- fprintf(stderr, "Failed to collect reservoir states");
- spires_reservoir_destroy(reservoir);
- return -1;
- }
- printf("collected state matrix size: %zu x %zu\n",
- state_matrix.num_samples, state_matrix.num_features);
-
- /* ---------- Generate raster plot ----------*/
- if (plot_raster(&state_matrix, NUM_NEURONS, 0.5) != 0) {
- fprintf(stderr, "Failed to plot raster");
+ if (mode == BENCHMARK_OFFLINE) {
+ if (collect_reservoir_states(reservoir, training_inputs,
+ NUM_TRAINING_STEPS,
+ &state_matrix) != 0) {
+ fprintf(stderr, "Failed to collect reservoir states");
+ spires_reservoir_destroy(reservoir);
+ return -1;
+ }
+ printf("Collected offline state matrix: %zu x %zu\n",
+ state_matrix.num_samples, state_matrix.num_features);
+ if (plot_raster(&state_matrix, NUM_NEURONS, 0.5) != 0)
+ fprintf(stderr, "Failed to plot raster\n");
}
/* ---------- Run Benchmark on each model ----------*/
- size_t predictions_per_model =
- state_matrix.num_samples * config.num_outputs;
+ size_t predictions_per_model = NUM_TRAINING_STEPS * config.num_outputs;
double *predictions = malloc(model_count * NUM_OUTPUTS *
NUM_TRAINING_STEPS * sizeof(double));
@@ -147,10 +171,19 @@ int main(void)
printf("\n\nRunning benchmark on %s\n",
models[model].model_path);
- if (run_benchmark(&config, reservoir, &state_matrix,
- models[model].model_path,
- models[model].subcircuit_name,
- model_predictions) < 0) {
+ int benchmark_status;
+ if (mode == BENCHMARK_ONLINE) {
+ benchmark_status = run_online_benchmark(
+ &config, reservoir, training_inputs,
+ NUM_TRAINING_STEPS, models[model].model_path,
+ models[model].subcircuit_name, model_predictions);
+ } else {
+ benchmark_status = run_benchmark(
+ &config, reservoir, &state_matrix,
+ models[model].model_path,
+ models[model].subcircuit_name, model_predictions);
+ }
+ if (benchmark_status < 0) {
fprintf(stderr, "Failed to run benchmark");
free(predictions);
free_reservoir_state_matrix(&state_matrix);
@@ -159,12 +192,12 @@ int main(void)
}
plot_reservoir_predictions(
- target_outputs, model_predictions, state_matrix.num_samples,
+ target_outputs, model_predictions, NUM_TRAINING_STEPS,
config.num_outputs, 0, models[model].model_path);
mean_squared_error[model] =
calculate_MSE(target_outputs, model_predictions,
- state_matrix.num_samples, config.num_outputs);
+ NUM_TRAINING_STEPS, config.num_outputs);
}
double *fixed_predictions = predictions;
@@ -174,8 +207,7 @@ int main(void)
model_predictions = predictions + model * predictions_per_model;
if (plot_model_delta(fixed_predictions, model_predictions,
- state_matrix.num_samples,
- config.num_outputs, 0,
+ NUM_TRAINING_STEPS, config.num_outputs, 0,
models[model].model_path) < 0) {
fprintf(stderr, "Failed to plot model delta\n");
}