From 00cd5b38f26e40fbb0dd21e9682e9a9580fe3e78 Mon Sep 17 00:00:00 2001 From: y-jan137 Date: Wed, 18 Mar 2026 09:56:29 +0300 Subject: Reformat --- include/adam.h | 16 ++++++++++++++++ include/adjoint.h | 44 ++++++++++++++++++++++++++++++++++++++++++++ include/dynmlp.h | 24 ++++++++++++++++++++++++ include/ode_solver.h | 17 +++++++++++++++++ include/spiral.h | 17 +++++++++++++++++ include/tests.h | 9 +++++++++ include/train.h | 16 ++++++++++++++++ include/utils.h | 39 +++++++++++++++++++++++++++++++++++++++ 8 files changed, 182 insertions(+) create mode 100644 include/adam.h create mode 100644 include/adjoint.h create mode 100644 include/dynmlp.h create mode 100644 include/ode_solver.h create mode 100644 include/spiral.h create mode 100644 include/tests.h create mode 100644 include/train.h create mode 100644 include/utils.h (limited to 'include') 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 +#include + +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); -- cgit v1.2.3