aboutsummaryrefslogtreecommitdiff
path: root/neural_ode.c
diff options
context:
space:
mode:
authory-jan137 <yousefjan24000@gmail.com>2026-03-17 16:05:16 +0300
committery-jan137 <yousefjan24000@gmail.com>2026-03-17 16:05:16 +0300
commit4659574b206eda0178fb1d92a43c3e843c23ca36 (patch)
treeaa0766b83ec812311a75e120ece9546b84cb5e84 /neural_ode.c
parent90ea2deac346116af54e06fc2ff538619934b732 (diff)
Add Workspace for memory allocation
Diffstat (limited to 'neural_ode.c')
-rw-r--r--neural_ode.c157
1 files changed, 107 insertions, 50 deletions
diff --git a/neural_ode.c b/neural_ode.c
index 81fce2a..abf2aad 100644
--- a/neural_ode.c
+++ b/neural_ode.c
@@ -54,6 +54,44 @@ static double *vec_zeros(int n) { return (double *)xcalloc((size_t)n, sizeof(do
static void vec_zero(double *v, int n) { memset(v, 0, (size_t)n * sizeof(double)); }
static void vec_copy(const double *src, double *dst, int n) { memcpy(dst, src, (size_t)n * sizeof(double)); }
+typedef struct {
+ double *x; /* [D+1] for dynmlp_forward/vjp */
+ double *h_pre; /* [H] for dynmlp_forward/vjp */
+ double *h; /* [H] for dynmlp_forward/vjp */
+ double *dh; /* [H] for dynmlp_vjp */
+ double *dh_pre; /* [H] for dynmlp_vjp */
+ double *dx; /* [D+1] for dynmlp_vjp */
+ double *neg_a; /* [D] for adjoint_dynamics */
+ double *vjp_z; /* [D] for adjoint_dynamics */
+ double *vjp_theta; /* [nparams] for adjoint_dynamics */
+} Workspace;
+
+static Workspace workspace_alloc(int D, int H, int nparams) {
+ Workspace ws;
+ ws.x = vec_alloc(D + 1);
+ ws.h_pre = vec_alloc(H);
+ ws.h = vec_alloc(H);
+ ws.dh = vec_alloc(H);
+ ws.dh_pre = vec_alloc(H);
+ ws.dx = vec_alloc(D + 1);
+ ws.neg_a = vec_alloc(D);
+ ws.vjp_z = vec_alloc(D);
+ ws.vjp_theta = vec_alloc(nparams);
+ return ws;
+}
+
+static void workspace_free(Workspace *ws) {
+ free(ws->x);
+ free(ws->h_pre);
+ free(ws->h);
+ free(ws->dh);
+ free(ws->dh_pre);
+ free(ws->dx);
+ free(ws->neg_a);
+ free(ws->vjp_z);
+ free(ws->vjp_theta);
+}
+
static void vec_add_scaled(double *dst, double alpha, const double *v, int n) {
for (int i = 0; i < n; i++) dst[i] += alpha * v[i];
}
@@ -122,16 +160,17 @@ static void dynmlp_init(DynMLP *net, int D, int H, double *theta, RNG *r) {
}
static void dynmlp_forward(const DynMLP *net, const double *theta,
- const double *z, double t, double *out) {
+ const double *z, double t, double *out,
+ Workspace *ws) {
int D = net->D, H = net->H;
const double *W1 = theta + DYNMLP_W1(D, H);
const double *b1 = theta + DYNMLP_b1(D, H);
const double *W2 = theta + DYNMLP_W2(D, H);
const double *b2 = theta + DYNMLP_b2(D, H);
- double *x = vec_alloc(D + 1);
- double *h_pre = vec_alloc(H);
- double *h = vec_alloc(H);
+ double *x = ws->x;
+ double *h_pre = ws->h_pre;
+ double *h = ws->h;
vec_copy(z, x, D);
x[D] = t;
@@ -143,15 +182,14 @@ static void dynmlp_forward(const DynMLP *net, const double *theta,
mat_vec(W2, h, out, D, H);
vec_add_scaled(out, 1.0, b2, D);
-
- free(x); free(h_pre); free(h);
}
/* Vector-Jacobian product: vjp_z = v^T (∂f/∂z), vjp_theta += v^T (∂f/∂θ)
Note: vjp_theta is accumulated into, not overwritten. */
static void dynmlp_vjp(const DynMLP *net, const double *theta,
const double *z, double t, const double *v,
- double *vjp_z, double *vjp_theta) {
+ double *vjp_z, double *vjp_theta,
+ Workspace *ws) {
int D = net->D;
int H = net->H;
const double *W1 = theta + DYNMLP_W1(D, H);
@@ -162,35 +200,35 @@ static void dynmlp_vjp(const DynMLP *net, const double *theta,
double *dW2 = vjp_theta + DYNMLP_W2(D, H);
double *db2 = vjp_theta + DYNMLP_b2(D, H);
- double *x = vec_alloc(D + 1);
- double *h_pre = vec_alloc(H);
- double *h = vec_alloc(H);
+ double *x = ws->x;
+ double *h_pre = ws->h_pre;
+ double *h = ws->h;
vec_copy(z, x, D);
x[D] = t;
mat_vec(W1, x, h_pre, H, D + 1);
vec_add_scaled(h_pre, 1.0, b1, H);
-
+
for (int i = 0; i < H; i++)
- h[i] = tanh(h_pre[i]);
+ h[i] = tanh(h_pre[i]);
- double *dh = vec_zeros(H);
- double *dh_pre = vec_alloc(H);
- double *dx = vec_zeros(D + 1);
+ double *dh = ws->dh;
+ double *dh_pre = ws->dh_pre;
+ double *dx = ws->dx;
+ vec_zero(dh, H);
+ vec_zero(dx, D + 1);
mat_vec_T(W2, v, dh, D, H);
mat_outer_add(dW2, 1.0, v, h, D, H);
vec_add_scaled(db2, 1.0, v, D);
- for (int i = 0; i < H; i++)
- dh_pre[i] = (1.0 - h[i] * h[i]) * dh[i];
+ for (int i = 0; i < H; i++)
+ dh_pre[i] = (1.0 - h[i] * h[i]) * dh[i];
mat_vec_T(W1, dh_pre, dx, H, D + 1);
mat_outer_add(dW1, 1.0, dh_pre, x, H, D + 1);
vec_add_scaled(db1, 1.0, dh_pre, H);
vec_copy(dx, vjp_z, D);
-
- free(x); free(h_pre); free(h); free(dh); free(dh_pre); free(dx);
}
/* ============================================================
@@ -346,44 +384,41 @@ typedef struct {
const double *theta;
int state_dim;
int nparams;
+ Workspace *ws;
} AdjointCtx;
static void neural_ode_rhs(const double *state, double t, const double *params,
int dim, double *out, void *ctx) {
- (void)params;
+ (void)params;
(void)dim;
-
+
AdjointCtx *ac = (AdjointCtx *)ctx;
-
- dynmlp_forward(&ac->net, ac->theta, state, t, out);
+ dynmlp_forward(&ac->net, ac->theta, state, t, out, ac->ws);
}
static void adjoint_dynamics(const double *aug_state, double t, const double *params,
int aug_dim, double *aug_out, void *ctx) {
- (void)params;
+ (void)params;
(void)aug_dim;
-
+
AdjointCtx *ac = (AdjointCtx *)ctx;
int D = ac->state_dim;
int nparams = ac->nparams;
const double *z = aug_state;
const double *a = aug_state + D;
- dynmlp_forward(&ac->net, ac->theta, z, t, aug_out);
+ dynmlp_forward(&ac->net, ac->theta, z, t, aug_out, ac->ws);
- double *neg_a = vec_alloc(D);
- double *vjp_z = vec_alloc(D);
- double *vjp_theta = vec_zeros(nparams);
+ double *neg_a = ac->ws->neg_a;
+ double *vjp_z = ac->ws->vjp_z;
+ double *vjp_theta = ac->ws->vjp_theta;
+ vec_zero(vjp_theta, nparams);
for (int i = 0; i < D; i++) neg_a[i] = -a[i];
- dynmlp_vjp(&ac->net, ac->theta, z, t, neg_a, vjp_z, vjp_theta);
+ dynmlp_vjp(&ac->net, ac->theta, z, t, neg_a, vjp_z, vjp_theta, ac->ws);
vec_copy(vjp_z, aug_out + D, D);
- vec_copy(vjp_theta, aug_out + 2 * D, nparams);
-
- free(neg_a);
- free(vjp_z);
- free(vjp_theta);
+ vec_copy(vjp_theta, aug_out + 2 * D, nparams);
}
typedef struct {
@@ -412,7 +447,8 @@ static ForwardResult forward_solve(const DynMLP *net, const double *theta,
const double *z0, double t0, double t1,
double atol, double rtol, int num_checkpoints) {
int D = net->D;
- AdjointCtx ac = { *net, theta, D, net->nparams };
+ Workspace ws = workspace_alloc(D, net->H, net->nparams);
+ AdjointCtx ac = { *net, theta, D, net->nparams, &ws };
ForwardResult fr;
fr.num_checkpoints = num_checkpoints;
@@ -440,6 +476,7 @@ static ForwardResult forward_solve(const DynMLP *net, const double *theta,
fr.z1 = vec_alloc(D);
vec_copy(fr.states[num_checkpoints], fr.z1, D);
+ workspace_free(&ws);
return fr;
}
@@ -454,7 +491,8 @@ static AdjointResult adjoint_solve(const DynMLP *net, const double *theta,
vec_copy(fr->z1, aug, D);
vec_copy(dL_dz1, aug + D, D);
- AdjointCtx ac = { *net, theta, D, nparams };
+ Workspace ws = workspace_alloc(D, net->H, nparams);
+ AdjointCtx ac = { *net, theta, D, nparams, &ws };
int total_nfe = 0;
for (int k = fr->num_checkpoints; k >= 1; k--) {
@@ -476,6 +514,7 @@ static AdjointResult adjoint_solve(const DynMLP *net, const double *theta,
vec_copy(aug + D, ar.dL_dz0, D);
vec_copy(aug + 2 * D, ar.dL_dtheta, nparams);
+ workspace_free(&ws);
free(aug);
return ar;
}
@@ -541,7 +580,8 @@ static MultiObsAdjointResult adjoint_solve_multi(
int nparams = net->nparams;
int aug_dim = 2 * D + nparams;
- AdjointCtx ac = { *net, theta, D, nparams };
+ Workspace ws = workspace_alloc(D, net->H, nparams);
+ AdjointCtx ac = { *net, theta, D, nparams, &ws };
double *a = vec_alloc(D);
double *dtheta = vec_zeros(nparams);
@@ -574,6 +614,7 @@ static MultiObsAdjointResult adjoint_solve_multi(
result.dL_dtheta = dtheta;
result.nfe = total_nfe;
+ workspace_free(&ws);
free(aug);
return result;
}
@@ -597,10 +638,12 @@ MultiObsNeuralODEOutput neural_ode_forward_backward_multi(
double rtol)
{
int D = net->D;
- AdjointCtx ac = { *net, theta, D, net->nparams };
+ Workspace ws = workspace_alloc(D, net->H, net->nparams);
+ AdjointCtx ac = { *net, theta, D, net->nparams, &ws };
ODEResult fwd = ode_solve_times(neural_ode_rhs, z0, times, ntimes,
NULL, D, atol, rtol, &ac);
+ workspace_free(&ws);
double *dL_dz_each = vec_alloc(ntimes * D);
for (int i = 0; i < ntimes * D; i++)
@@ -611,7 +654,7 @@ MultiObsNeuralODEOutput neural_ode_forward_backward_multi(
free(dL_dz_each);
MultiObsNeuralODEOutput out;
- out.z_traj = fwd.y;
+ out.z_traj = fwd.y;
out.dL_dz0 = ar.dL_dz0;
out.dL_dtheta = ar.dL_dtheta;
out.nfe_forward = fwd.nfe;
@@ -774,7 +817,8 @@ static double evaluate(const DynMLP *net, const double *theta,
const Dataset *ds, double t0, double t1,
double atol, double rtol) {
int D = net->D;
- AdjointCtx ac = { *net, theta, D, net->nparams };
+ Workspace ws = workspace_alloc(D, net->H, net->nparams);
+ AdjointCtx ac = { *net, theta, D, net->nparams, &ws };
double total_loss = 0.0;
for (int i = 0; i < ds->num_samples; i++) {
ODEResult fwd = ode_solve(neural_ode_rhs, ds->z0[i], t0, t1,
@@ -785,6 +829,7 @@ static double evaluate(const DynMLP *net, const double *theta,
}
free(fwd.y);
}
+ workspace_free(&ws);
return total_loss / (double)ds->num_samples;
}
@@ -841,15 +886,17 @@ static void test_dynmlp_gradients(RNG *r) {
for (int i = 0; i < D; i++) v[i] = rng_normal(r);
double t = rng_normal(r);
+ Workspace ws = workspace_alloc(D, H, np);
+
double *vjp_z = vec_zeros(D);
double *vjp_theta = vec_zeros(np);
- dynmlp_vjp(&net, theta, z, t, v, vjp_z, vjp_theta);
+ dynmlp_vjp(&net, theta, z, t, v, vjp_z, vjp_theta, &ws);
double *num_vjp_z = vec_alloc(D);
for (int i = 0; i < D; i++) {
double zi = z[i];
- z[i] = zi + EPS; dynmlp_forward(&net, theta, z, t, out_p);
- z[i] = zi - EPS; dynmlp_forward(&net, theta, z, t, out_m);
+ z[i] = zi + EPS; dynmlp_forward(&net, theta, z, t, out_p, &ws);
+ z[i] = zi - EPS; dynmlp_forward(&net, theta, z, t, out_m, &ws);
z[i] = zi;
num_vjp_z[i] = (vec_dot(v, out_p, D) - vec_dot(v, out_m, D)) / (2.0 * EPS);
}
@@ -862,8 +909,8 @@ static void test_dynmlp_gradients(RNG *r) {
double *num_vjp_theta = vec_alloc(np);
for (int k = 0; k < np; k++) {
double tk = theta[k];
- theta[k] = tk + EPS; dynmlp_forward(&net, theta, z, t, out_p);
- theta[k] = tk - EPS; dynmlp_forward(&net, theta, z, t, out_m);
+ theta[k] = tk + EPS; dynmlp_forward(&net, theta, z, t, out_p, &ws);
+ theta[k] = tk - EPS; dynmlp_forward(&net, theta, z, t, out_m, &ws);
theta[k] = tk;
num_vjp_theta[k] = (vec_dot(v, out_p, D) - vec_dot(v, out_m, D)) / (2.0 * EPS);
}
@@ -876,6 +923,7 @@ static void test_dynmlp_gradients(RNG *r) {
printf("dL/dz: max_err=%.2e %s\n", max_err_z, max_err_z < TOL ? "PASS" : "FAIL");
printf("dL/dtheta: max_err=%.2e %s\n", max_err_theta, max_err_theta < TOL ? "PASS" : "FAIL");
+ workspace_free(&ws);
free(theta); free(z); free(v); free(out_p); free(out_m);
free(vjp_z); free(vjp_theta); free(num_vjp_z); free(num_vjp_theta);
}
@@ -898,8 +946,10 @@ static void test_adjoint_gradients(RNG *r) {
NeuralODEOutput out = neural_ode_forward_backward(&net, theta, z0, t0, t1,
target, atol, rtol, 10);
+ Workspace ws = workspace_alloc(D, H, np);
+
#define FWD_LOSS(z0_, theta_) ({ \
- AdjointCtx ac_ = { net, (theta_), D, np }; \
+ AdjointCtx ac_ = { net, (theta_), D, np, &ws }; \
ODEResult r_ = ode_solve(neural_ode_rhs, (z0_), t0, t1, NULL, D, atol, rtol, &ac_); \
double l_ = 0.0; \
for (int _i = 0; _i < D; _i++) { double _d = r_.y[_i] - target[_i]; l_ += 0.5*_d*_d; } \
@@ -947,6 +997,7 @@ static void test_adjoint_gradients(RNG *r) {
#undef FWD_LOSS
+ workspace_free(&ws);
free(out.z1); free(out.dL_dz0); free(out.dL_dtheta);
free(num_dL_dtheta); free(num_dL_dz0);
free(theta); free(z0); free(target);
@@ -972,7 +1023,8 @@ static void test_multi_obs_adjoint(RNG *r) {
MultiObsNeuralODEOutput out = neural_ode_forward_backward_multi(
&net, theta, z0, times, targets, ntimes, atol, rtol);
- AdjointCtx ac = { net, theta, D, np };
+ Workspace ws = workspace_alloc(D, H, np);
+ AdjointCtx ac = { net, theta, D, np, &ws };
/* Numerical dL/dtheta */
double *num_dL_dtheta = vec_alloc(np);
@@ -1052,6 +1104,7 @@ static void test_multi_obs_adjoint(RNG *r) {
printf("multi-obs adjoint dL/dz0: max_rel_err=%.2e %s\n",
rel_z0, rel_z0 < 1e-3 ? "PASS" : "FAIL");
+ workspace_free(&ws);
free(out.z_traj); free(out.dL_dz0); free(out.dL_dtheta);
free(num_dL_dtheta); free(num_dL_dz0);
free(theta); free(z0); free(targets);
@@ -1097,7 +1150,8 @@ static void test_training(RNG *r) {
printf("iter %3d loss=%.4f nfe_fwd=%d\n", iter + 1, res.loss, res.nfe_fwd);
}
- AdjointCtx ac = { net, theta, D, nparams };
+ Workspace ws = workspace_alloc(D, H, nparams);
+ AdjointCtx ac = { net, theta, D, nparams, &ws };
double final_loss = 0.0;
for (int i = 0; i < N; i++) {
ODEResult fwd = ode_solve(neural_ode_rhs, z0s[i], t0, t1, NULL, D, atol, rtol, &ac);
@@ -1110,6 +1164,7 @@ static void test_training(RNG *r) {
final_loss /= (double)N;
printf("Loss: %.4f\n", final_loss);
+ workspace_free(&ws);
adam_free(&adam);
free(batch_z0); free(batch_tgt);
for (int i = 0; i < N; i++) { free(z0s[i]); free(targets[i]); }
@@ -1189,7 +1244,8 @@ int main(void) {
printf("Total parameters: %d\n", nparams);
printf("\nSample predictions:\n");
- AdjointCtx ac = { net, theta, D, nparams };
+ Workspace ws = workspace_alloc(D, H, nparams);
+ AdjointCtx ac = { net, theta, D, nparams, &ws };
for (int s = 0; s < 5; s++) {
int idx = (int)(rng_next(&r) % (uint64_t)TEST_N);
ODEResult fwd = ode_solve(neural_ode_rhs,
@@ -1201,6 +1257,7 @@ int main(void) {
test_ds.target[idx][0], test_ds.target[idx][1]);
free(fwd.y);
}
+ workspace_free(&ws);
free(theta);
adam_free(&adam);