aboutsummaryrefslogtreecommitdiff
path: root/include
diff options
context:
space:
mode:
Diffstat (limited to 'include')
-rw-r--r--include/adam.h16
-rw-r--r--include/adjoint.h44
-rw-r--r--include/dynmlp.h24
-rw-r--r--include/ode_solver.h17
-rw-r--r--include/spiral.h17
-rw-r--r--include/tests.h9
-rw-r--r--include/train.h16
-rw-r--r--include/utils.h39
8 files changed, 182 insertions, 0 deletions
diff --git a/include/adam.h b/include/adam.h
new file mode 100644
index 0000000..37f5767
--- /dev/null
+++ b/include/adam.h
@@ -0,0 +1,16 @@
+#pragma once
+
+typedef struct {
+ double *m;
+ double *v;
+ int nparams;
+ double lr;
+ double beta1;
+ double beta2;
+ double eps;
+ int t;
+} Adam;
+
+Adam adam_init(int nparams, double lr, double beta1, double beta2, double eps);
+void adam_update(Adam *a, double *theta, const double *grad);
+void adam_free(Adam *a);
diff --git a/include/adjoint.h b/include/adjoint.h
new file mode 100644
index 0000000..d9e71a9
--- /dev/null
+++ b/include/adjoint.h
@@ -0,0 +1,44 @@
+#pragma once
+
+#include "utils.h"
+#include "dynmlp.h"
+#include "ode_solver.h"
+
+typedef struct {
+ DynMLP net;
+ const double *theta;
+ int state_dim;
+ int nparams;
+ Workspace *ws;
+} AdjointCtx;
+
+typedef struct {
+ double *z1;
+ double *dL_dz0;
+ double *dL_dtheta;
+ int nfe_forward;
+ int nfe_backward;
+} NeuralODEOutput;
+
+typedef struct {
+ double *z_traj;
+ double *dL_dz0;
+ double *dL_dtheta;
+ int nfe_forward;
+ int nfe_backward;
+} MultiObsNeuralODEOutput;
+
+void neural_ode_rhs(const double *state, double t, const double *params,
+ int dim, double *out, void *ctx);
+
+NeuralODEOutput neural_ode_forward_backward(
+ const DynMLP *net, const double *theta,
+ const double *z0, double t0, double t1,
+ const double *target, double atol, double rtol,
+ int num_checkpoints);
+
+MultiObsNeuralODEOutput neural_ode_forward_backward_multi(
+ const DynMLP *net, const double *theta,
+ const double *z0, const double *times,
+ const double *targets, int ntimes,
+ double atol, double rtol);
diff --git a/include/dynmlp.h b/include/dynmlp.h
new file mode 100644
index 0000000..058271e
--- /dev/null
+++ b/include/dynmlp.h
@@ -0,0 +1,24 @@
+#pragma once
+
+#include "utils.h"
+
+typedef struct {
+ int D;
+ int H;
+ int nparams;
+} DynMLP;
+
+#define DYNMLP_W1(D, H) (0)
+#define DYNMLP_b1(D, H) ((D + 1) * (H))
+#define DYNMLP_W2(D, H) ((D + 1) * (H) + (H))
+#define DYNMLP_b2(D, H) ((D + 1) * (H) + (H) + (H) * (D))
+
+int dynmlp_nparams(int D, int H);
+void dynmlp_init(DynMLP *net, int D, int H, double *theta, RNG *r);
+void dynmlp_forward(const DynMLP *net, const double *theta,
+ const double *z, double t, double *out,
+ Workspace *ws);
+void dynmlp_vjp(const DynMLP *net, const double *theta,
+ const double *z, double t, const double *v,
+ double *vjp_z, double *vjp_theta,
+ Workspace *ws);
diff --git a/include/ode_solver.h b/include/ode_solver.h
new file mode 100644
index 0000000..91956b0
--- /dev/null
+++ b/include/ode_solver.h
@@ -0,0 +1,17 @@
+#pragma once
+
+typedef void (*ode_rhs_fn)(const double *state, double t, const double *params,
+ int dim, double *out, void *ctx);
+
+typedef struct {
+ double *y;
+ int nfe;
+} ODEResult;
+
+ODEResult ode_solve(ode_rhs_fn f, const double *y0, double t0, double t1,
+ const double *params, int dim, double atol, double rtol,
+ void *ctx);
+
+ODEResult ode_solve_times(ode_rhs_fn f, const double *y0, const double *times,
+ int ntimes, const double *params, int dim,
+ double atol, double rtol, void *ctx);
diff --git a/include/spiral.h b/include/spiral.h
new file mode 100644
index 0000000..5ed315e
--- /dev/null
+++ b/include/spiral.h
@@ -0,0 +1,17 @@
+#pragma once
+
+#include "utils.h"
+#include "dynmlp.h"
+
+typedef struct {
+ double **z0;
+ double **target;
+ int num_samples;
+} Dataset;
+
+Dataset generate_spiral_dataset(int num_samples, double t0, double t1,
+ double noise_std, RNG *r);
+void dataset_free(Dataset *ds);
+double evaluate(const DynMLP *net, const double *theta,
+ const Dataset *ds, double t0, double t1,
+ double atol, double rtol);
diff --git a/include/tests.h b/include/tests.h
new file mode 100644
index 0000000..794c7d6
--- /dev/null
+++ b/include/tests.h
@@ -0,0 +1,9 @@
+#pragma once
+
+#include "utils.h"
+
+void test_ode_solver(void);
+void test_dynmlp_gradients(RNG *r);
+void test_adjoint_gradients(RNG *r);
+void test_multi_obs_adjoint(RNG *r);
+void test_training(RNG *r);
diff --git a/include/train.h b/include/train.h
new file mode 100644
index 0000000..9cf613e
--- /dev/null
+++ b/include/train.h
@@ -0,0 +1,16 @@
+#pragma once
+
+#include "dynmlp.h"
+#include "adam.h"
+
+typedef struct {
+ double loss;
+ int nfe_fwd;
+ int nfe_bwd;
+} TrainStepResult;
+
+TrainStepResult train_step(const DynMLP *net, double *theta,
+ const double **z0s, const double **targets,
+ double t0, double t1, int batch_size,
+ Adam *adam, double atol, double rtol,
+ int num_checkpoints);
diff --git a/include/utils.h b/include/utils.h
new file mode 100644
index 0000000..e0901ff
--- /dev/null
+++ b/include/utils.h
@@ -0,0 +1,39 @@
+#pragma once
+
+#include <stdint.h>
+#include <stddef.h>
+
+typedef struct { uint64_t state; } RNG;
+
+typedef struct {
+ double *x;
+ double *h_pre;
+ double *h;
+ double *dh;
+ double *dh_pre;
+ double *dx;
+ double *neg_a;
+ double *vjp_z;
+ double *vjp_theta;
+} Workspace;
+
+RNG rng_init(uint64_t seed);
+uint64_t rng_next(RNG *r);
+double rng_uniform(RNG *r);
+double rng_normal(RNG *r);
+
+void *xmalloc(size_t n);
+void *xcalloc(size_t count, size_t size);
+
+double *vec_alloc(int n);
+double *vec_zeros(int n);
+void vec_zero(double *v, int n);
+void vec_copy(const double *src, double *dst, int n);
+void vec_add_scaled(double *dst, double alpha, const double *v, int n);
+double vec_dot(const double *a, const double *b, int n);
+void mat_vec(const double *M, const double *x, double *dst, int rows, int cols);
+void mat_vec_T(const double *M, const double *v, double *dst, int rows, int cols);
+void mat_outer_add(double *M, double alpha, const double *a, const double *b, int rows, int cols);
+
+Workspace workspace_alloc(int D, int H, int nparams);
+void workspace_free(Workspace *ws);