aboutsummaryrefslogtreecommitdiff
path: root/include
diff options
context:
space:
mode:
Diffstat (limited to 'include')
-rw-r--r--include/adjoint.h14
-rw-r--r--include/cnf.h24
-rw-r--r--include/dynmlp.h24
-rw-r--r--include/dynnet.h50
-rw-r--r--include/spiral.h5
-rw-r--r--include/tests.h9
-rw-r--r--include/train.h4
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,