diff options
Diffstat (limited to 'include')
| -rw-r--r-- | include/adjoint.h | 14 | ||||
| -rw-r--r-- | include/cnf.h | 24 | ||||
| -rw-r--r-- | include/dynmlp.h | 24 | ||||
| -rw-r--r-- | include/dynnet.h | 50 | ||||
| -rw-r--r-- | include/spiral.h | 5 | ||||
| -rw-r--r-- | include/tests.h | 9 | ||||
| -rw-r--r-- | include/train.h | 4 |
7 files changed, 80 insertions, 50 deletions
diff --git a/include/adjoint.h b/include/adjoint.h index d9e71a9..f9fa37a 100644 --- a/include/adjoint.h +++ b/include/adjoint.h @@ -1,15 +1,13 @@ #pragma once -#include "utils.h" -#include "dynmlp.h" +#include "dynnet.h" #include "ode_solver.h" typedef struct { - DynMLP net; + DynNet *net; const double *theta; - int state_dim; - int nparams; - Workspace *ws; + double *ws; /* size net->total_workspace */ + double *neg_a; /* size net->D */ } AdjointCtx; typedef struct { @@ -32,13 +30,13 @@ 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, + DynNet *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, + DynNet *net, const double *theta, const double *z0, const double *times, const double *targets, int ntimes, double atol, double rtol); diff --git a/include/cnf.h b/include/cnf.h index 5dc5688..a9ab5bd 100644 --- a/include/cnf.h +++ b/include/cnf.h @@ -1,44 +1,44 @@ #pragma once -#include "dynmlp.h" +#include "dynnet.h" #include "ode_solver.h" -#include "utils.h" typedef struct { - DynMLP net; + DynNet *net; int nparams; - double trace_eps; /* finite-difference epsilon for trace computation */ + double trace_eps; } CNF; typedef struct { - double *z1; + double *z1; double delta_logp; int nfe; } CNFSampleResult; typedef struct { - double *z0; - double delta_logp; + double *z0; + double delta_logp; int nfe; } CNFLogProbResult; typedef struct { - double *dL_dz0; - double *dL_dtheta; + double *dL_dz0; + double *dL_dtheta; int nfe; } CNFBackwardResult; void cnf_init(CNF *cnf, int D, int H, double *theta, RNG *r); +void cnf_free(CNF *cnf); -CNFSampleResult cnf_sample(const CNF *cnf, const double *theta, +CNFSampleResult cnf_sample(CNF *cnf, const double *theta, const double *z0, double t0, double t1, double atol, double rtol); -CNFLogProbResult cnf_log_prob(const CNF *cnf, const double *theta, +CNFLogProbResult cnf_log_prob(CNF *cnf, const double *theta, const double *z1, double t0, double t1, double atol, double rtol); -CNFBackwardResult cnf_backward(const CNF *cnf, const double *theta, +CNFBackwardResult cnf_backward(CNF *cnf, const double *theta, const double *z1, double dL_dlogp, const double *dL_dz1_in, double t0, double t1, diff --git a/include/dynmlp.h b/include/dynmlp.h deleted file mode 100644 index 058271e..0000000 --- a/include/dynmlp.h +++ /dev/null @@ -1,24 +0,0 @@ -#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/dynnet.h b/include/dynnet.h new file mode 100644 index 0000000..966c239 --- /dev/null +++ b/include/dynnet.h @@ -0,0 +1,50 @@ +#pragma once + +#include "utils.h" + +typedef struct { + void (*forward)(const void *cfg, const double *theta, + const double *in, double *out, double *workspace); + void (*vjp)(const void *cfg, const double *theta, + const double *in, const double *v_out, + double *v_in, double *v_theta, double *workspace); + int (*nparams)(const void *cfg); + int (*workspace_size)(const void *cfg); +} LayerOps; + +typedef struct DynNet { + int num_layers; + const LayerOps **ops; + void **layer_configs; + int *param_offsets; + int *ws_offsets; + int *in_dims; + int total_params; + int total_workspace; + int v_buf_offset; + int max_dim; + int D; + double current_t; + /* build state (non-null before finalize) */ + void *_pre; + int _pre_n, _pre_cap, _cur_dim; +} DynNet; + +DynNet *dynnet_create(int D); +void dynnet_add_linear(DynNet *net, int out_dim); +void dynnet_add_tanh(DynNet *net); +void dynnet_add_softplus(DynNet *net); +void dynnet_add_swish(DynNet *net); +void dynnet_add_layernorm(DynNet *net); +void dynnet_add_residual_begin(DynNet *net); +void dynnet_add_residual_end(DynNet *net); +void dynnet_add_time_concat(DynNet *net); +void dynnet_finalize(DynNet *net); +void dynnet_init_params(DynNet *net, double *theta, RNG *r); +void dynnet_free(DynNet *net); + +void dynnet_forward(DynNet *net, const double *theta, + const double *z, double t, double *out, double *workspace); +void dynnet_vjp(DynNet *net, const double *theta, + const double *z, double t, const double *v, + double *vjp_z, double *vjp_theta, double *workspace); diff --git a/include/spiral.h b/include/spiral.h index 5ed315e..2b8e3ea 100644 --- a/include/spiral.h +++ b/include/spiral.h @@ -1,7 +1,6 @@ #pragma once -#include "utils.h" -#include "dynmlp.h" +#include "dynnet.h" typedef struct { double **z0; @@ -12,6 +11,6 @@ typedef struct { 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, +double evaluate(DynNet *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 index 794c7d6..d3e2acd 100644 --- a/include/tests.h +++ b/include/tests.h @@ -3,7 +3,14 @@ #include "utils.h" void test_ode_solver(void); -void test_dynmlp_gradients(RNG *r); +void test_dynnet_gradients(RNG *r); void test_adjoint_gradients(RNG *r); void test_multi_obs_adjoint(RNG *r); void test_training(RNG *r); +void test_dynnet_deep_gradients(RNG *r); +void test_dynnet_residual_gradients(RNG *r); +void test_dynnet_layernorm_gradients(RNG *r); +void test_dynnet_time_inject_gradients(RNG *r); +void test_dynnet_swish_gradients(RNG *r); +void test_dynnet_softplus_gradients(RNG *r); +void test_adjoint_deep(RNG *r); diff --git a/include/train.h b/include/train.h index 9cf613e..94c9666 100644 --- a/include/train.h +++ b/include/train.h @@ -1,6 +1,6 @@ #pragma once -#include "dynmlp.h" +#include "dynnet.h" #include "adam.h" typedef struct { @@ -9,7 +9,7 @@ typedef struct { int nfe_bwd; } TrainStepResult; -TrainStepResult train_step(const DynMLP *net, double *theta, +TrainStepResult train_step(DynNet *net, double *theta, const double **z0s, const double **targets, double t0, double t1, int batch_size, Adam *adam, double atol, double rtol, |