From 8671410a0f9bdd4b405762ad4790c15d601b0fba Mon Sep 17 00:00:00 2001 From: Your Name Date: Fri, 28 Aug 2026 13:42:10 -0700 Subject: online crossbar inputs --- src/application.c | 82 ++++++++++++++++++++++++++++++++++++++----------------- 1 file changed, 57 insertions(+), 25 deletions(-) (limited to 'src/application.c') 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 #include #include +#include // 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"); } -- cgit v1.2.3