aboutsummaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
authory-jan137 <yousefjan24000@gmail.com>2026-03-18 11:44:05 +0300
committery-jan137 <yousefjan24000@gmail.com>2026-03-18 11:44:05 +0300
commit161352cf4504c2113ad148ce5b5917a5bed36c06 (patch)
treef5c30d3893e5d68f8830a01d947577d0759062eb /src
parent00cd5b38f26e40fbb0dd21e9682e9a9580fe3e78 (diff)
Add continuous normalizing flow
Diffstat (limited to 'src')
-rw-r--r--src/cnf.c233
-rw-r--r--src/cnf_train.c128
-rw-r--r--src/main.c10
-rw-r--r--src/test_cnf.c240
4 files changed, 611 insertions, 0 deletions
diff --git a/src/cnf.c b/src/cnf.c
new file mode 100644
index 0000000..9b121e6
--- /dev/null
+++ b/src/cnf.c
@@ -0,0 +1,233 @@
+#include "cnf.h"
+
+#include <stdlib.h>
+
+/* Compute tr(df/dz) via central finite differences.
+ * Uses z_p (D) as scratch; writes intermediate outputs to f_p and f_m (D). */
+static double compute_trace_fd(const DynMLP *net, const double *theta,
+ const double *z, double t, Workspace *ws,
+ double eps, double *z_p, double *f_p, double *f_m) {
+ int D = net->D;
+ double tr = 0.0;
+ for (int i = 0; i < D; i++) {
+ vec_copy(z, z_p, D);
+ z_p[i] += eps;
+ dynmlp_forward(net, theta, z_p, t, f_p, ws);
+ z_p[i] = z[i] - eps;
+ dynmlp_forward(net, theta, z_p, t, f_m, ws);
+ tr += (f_p[i] - f_m[i]) / (2.0 * eps);
+ }
+ return tr;
+}
+
+/* ---- Forward augmented ODE ---- */
+typedef struct {
+ DynMLP net;
+ const double *theta;
+ int D;
+ Workspace *ws;
+ double eps;
+ double *z_p, *f_p, *f_m; /* D */
+} CNFFwdCtx;
+
+static void cnf_fwd_rhs(const double *aug, double t, const double *params,
+ int aug_dim, double *out, void *ctx) {
+ (void)params; (void)aug_dim;
+ CNFFwdCtx *c = (CNFFwdCtx *)ctx;
+ int D = c->D;
+
+ dynmlp_forward(&c->net, c->theta, aug, t, out, c->ws);
+ out[D] = -compute_trace_fd(&c->net, c->theta, aug, t, c->ws,
+ c->eps, c->z_p, c->f_p, c->f_m);
+}
+
+/* aug_state = [z(D), a_z(D), g(nparams)]
+ * a_logp (= dL/d(logdet)) is a constant stored in ctx since d(logdet)/dt
+ * does not depend on logdet itself. */
+typedef struct {
+ DynMLP net;
+ const double *theta;
+ int D;
+ int nparams;
+ Workspace *ws;
+ double eps;
+ double a_logp;
+ /* scratch buffers */
+ double *z_oj; /* D: outer perturbation (trace grad z & theta) */
+ double *z_p2, *f_p2, *f_m2; /* D: inner scratch for nested compute_trace_fd */
+ double *neg_a_z; /* D */
+ double *dg_tmp; /* nparams */
+ double *da_z_tmp; /* D (vjp_z scratch, not used as output) */
+ double *v_ei; /* D: one-hot for trace grad theta */
+ double *vjp_z_tmp; /* D */
+ double *vjp_th_p; /* nparams */
+ double *vjp_th_m; /* nparams */
+} CNFAdjCtx;
+
+static void cnf_adj_rhs(const double *aug, double t, const double *params,
+ int aug_dim, double *out, void *ctx) {
+ (void)params; (void)aug_dim;
+ CNFAdjCtx *cc = (CNFAdjCtx *)ctx;
+ int D = cc->D;
+ int nparams = cc->nparams;
+ const double *z = aug;
+ const double *a_z = aug + D;
+
+ dynmlp_forward(&cc->net, cc->theta, z, t, out, cc->ws);
+
+ for (int i = 0; i < D; i++) cc->neg_a_z[i] = -a_z[i];
+ vec_zero(cc->dg_tmp, nparams);
+ dynmlp_vjp(&cc->net, cc->theta, z, t, cc->neg_a_z,
+ out + D, cc->dg_tmp, cc->ws);
+ vec_copy(cc->dg_tmp, out + 2 * D, nparams);
+
+ if (cc->a_logp == 0.0) return; /* no trace gradient terms needed */
+
+ for (int j = 0; j < D; j++) {
+ vec_copy(z, cc->z_oj, D);
+ cc->z_oj[j] += cc->eps;
+ double tr_p = compute_trace_fd(&cc->net, cc->theta, cc->z_oj, t, cc->ws,
+ cc->eps, cc->z_p2, cc->f_p2, cc->f_m2);
+ cc->z_oj[j] = z[j] - cc->eps;
+ double tr_m = compute_trace_fd(&cc->net, cc->theta, cc->z_oj, t, cc->ws,
+ cc->eps, cc->z_p2, cc->f_p2, cc->f_m2);
+ out[D + j] += cc->a_logp * (tr_p - tr_m) / (2.0 * cc->eps);
+ }
+
+ vec_zero(cc->v_ei, D);
+ for (int i = 0; i < D; i++) {
+ cc->v_ei[i] = 1.0;
+
+ vec_copy(z, cc->z_oj, D);
+ cc->z_oj[i] += cc->eps;
+ vec_zero(cc->vjp_th_p, nparams);
+ dynmlp_vjp(&cc->net, cc->theta, cc->z_oj, t, cc->v_ei,
+ cc->vjp_z_tmp, cc->vjp_th_p, cc->ws);
+
+ cc->z_oj[i] = z[i] - cc->eps;
+ vec_zero(cc->vjp_th_m, nparams);
+ dynmlp_vjp(&cc->net, cc->theta, cc->z_oj, t, cc->v_ei,
+ cc->vjp_z_tmp, cc->vjp_th_m, cc->ws);
+
+ cc->v_ei[i] = 0.0;
+
+ double scale = cc->a_logp / (2.0 * cc->eps);
+ for (int k = 0; k < nparams; k++)
+ out[2 * D + k] += scale * (cc->vjp_th_p[k] - cc->vjp_th_m[k]);
+ }
+}
+
+/* ---- Public API ---- */
+
+void cnf_init(CNF *cnf, int D, int H, double *theta, RNG *r) {
+ dynmlp_init(&cnf->net, D, H, theta, r);
+ cnf->nparams = dynmlp_nparams(D, H);
+ cnf->trace_eps = 1e-5;
+}
+
+CNFSampleResult cnf_sample(const CNF *cnf, const double *theta,
+ const double *z0, double t0, double t1,
+ double atol, double rtol) {
+ int D = cnf->net.D;
+ Workspace ws = workspace_alloc(D, cnf->net.H, cnf->nparams);
+ double *z_p = vec_alloc(D), *f_p = vec_alloc(D), *f_m = vec_alloc(D);
+ CNFFwdCtx ctx = { cnf->net, theta, D, &ws, cnf->trace_eps, z_p, f_p, f_m };
+
+ double *aug0 = vec_zeros(D + 1);
+ vec_copy(z0, aug0, D);
+
+ ODEResult res = ode_solve(cnf_fwd_rhs, aug0, t0, t1,
+ NULL, D + 1, atol, rtol, &ctx);
+ free(aug0);
+ free(z_p); free(f_p); free(f_m);
+ workspace_free(&ws);
+
+ CNFSampleResult out;
+ out.z1 = vec_alloc(D);
+ vec_copy(res.y, out.z1, D);
+ out.delta_logp = res.y[D];
+ out.nfe = res.nfe;
+ free(res.y);
+ return out;
+}
+
+CNFLogProbResult cnf_log_prob(const CNF *cnf, const double *theta,
+ const double *z1, double t0, double t1,
+ double atol, double rtol) {
+ int D = cnf->net.D;
+ Workspace ws = workspace_alloc(D, cnf->net.H, cnf->nparams);
+ double *z_p = vec_alloc(D), *f_p = vec_alloc(D), *f_m = vec_alloc(D);
+ CNFFwdCtx ctx = { cnf->net, theta, D, &ws, cnf->trace_eps, z_p, f_p, f_m };
+
+ double *aug1 = vec_zeros(D + 1);
+ vec_copy(z1, aug1, D);
+ ODEResult res = ode_solve(cnf_fwd_rhs, aug1, t1, t0,
+ NULL, D + 1, atol, rtol, &ctx);
+ free(aug1);
+ free(z_p); free(f_p); free(f_m);
+ workspace_free(&ws);
+
+ /* Backward integral gives ∫_{t1}^{t0} -tr dt = -delta_logp_fwd,
+ * so delta_logp = log p(z1) - log p(z0) = -res.y[D]. */
+ CNFLogProbResult out;
+ out.z0 = vec_alloc(D);
+ vec_copy(res.y, out.z0, D);
+ out.delta_logp = -res.y[D];
+ out.nfe = res.nfe;
+ free(res.y);
+ return out;
+}
+
+CNFBackwardResult cnf_backward(const CNF *cnf, const double *theta,
+ const double *z1,
+ double dL_dlogp, const double *dL_dz1_in,
+ double t0, double t1,
+ double atol, double rtol) {
+ int D = cnf->net.D;
+ int nparams = cnf->nparams;
+ int aug_dim = 2 * D + nparams;
+
+ Workspace ws = workspace_alloc(D, cnf->net.H, nparams);
+
+ CNFAdjCtx ctx;
+ ctx.net = cnf->net;
+ ctx.theta = theta;
+ ctx.D = D;
+ ctx.nparams = nparams;
+ ctx.ws = &ws;
+ ctx.eps = cnf->trace_eps;
+ ctx.a_logp = dL_dlogp;
+ ctx.z_oj = vec_alloc(D);
+ ctx.z_p2 = vec_alloc(D);
+ ctx.f_p2 = vec_alloc(D);
+ ctx.f_m2 = vec_alloc(D);
+ ctx.neg_a_z = vec_alloc(D);
+ ctx.dg_tmp = vec_alloc(nparams);
+ ctx.da_z_tmp = vec_alloc(D);
+ ctx.v_ei = vec_zeros(D);
+ ctx.vjp_z_tmp = vec_alloc(D);
+ ctx.vjp_th_p = vec_alloc(nparams);
+ ctx.vjp_th_m = vec_alloc(nparams);
+
+ double *aug = vec_zeros(aug_dim);
+ vec_copy(z1, aug, D);
+ vec_copy(dL_dz1_in, aug + D, D);
+
+ ODEResult res = ode_solve(cnf_adj_rhs, aug, t1, t0,
+ NULL, aug_dim, atol, rtol, &ctx);
+ free(aug);
+
+ CNFBackwardResult out;
+ out.dL_dz0 = vec_alloc(D);
+ out.dL_dtheta = vec_alloc(nparams);
+ vec_copy(res.y + D, out.dL_dz0, D);
+ vec_copy(res.y + 2 * D, out.dL_dtheta, nparams);
+ out.nfe = res.nfe;
+ free(res.y);
+
+ free(ctx.z_oj); free(ctx.z_p2); free(ctx.f_p2); free(ctx.f_m2);
+ free(ctx.neg_a_z); free(ctx.dg_tmp); free(ctx.da_z_tmp);
+ free(ctx.v_ei); free(ctx.vjp_z_tmp); free(ctx.vjp_th_p); free(ctx.vjp_th_m);
+ workspace_free(&ws);
+ return out;
+}
diff --git a/src/cnf_train.c b/src/cnf_train.c
new file mode 100644
index 0000000..21b357f
--- /dev/null
+++ b/src/cnf_train.c
@@ -0,0 +1,128 @@
+#include "cnf_train.h"
+#include "cnf.h"
+#include "adam.h"
+
+#include <stdio.h>
+#include <stdlib.h>
+#include <math.h>
+
+static const double TARGET_MU[4][2] = {
+ { 1.5, 0.0}, {-1.5, 0.0},
+ { 0.0, 1.5}, { 0.0, -1.5}
+};
+static const double TARGET_SIGMA2 = 0.2;
+static const int TARGET_K = 4;
+
+static double log_p_target(const double *z) {
+ double lc[4], mx = -1e300;
+ for (int k = 0; k < TARGET_K; k++) {
+ double dx = z[0] - TARGET_MU[k][0], dy = z[1] - TARGET_MU[k][1];
+ lc[k] = -0.5 * (dx * dx + dy * dy) / TARGET_SIGMA2;
+ if (lc[k] > mx) mx = lc[k];
+ }
+ double s = 0.0;
+ for (int k = 0; k < TARGET_K; k++) s += exp(lc[k] - mx);
+ return mx + log(s / (double)TARGET_K)
+ - log(2.0 * M_PI * TARGET_SIGMA2);
+}
+
+static void grad_log_p_target(const double *z, double *g) {
+ double lc[4], mx = -1e300;
+ for (int k = 0; k < TARGET_K; k++) {
+ double dx = z[0] - TARGET_MU[k][0], dy = z[1] - TARGET_MU[k][1];
+ lc[k] = -0.5 * (dx * dx + dy * dy) / TARGET_SIGMA2;
+ if (lc[k] > mx) mx = lc[k];
+ }
+ double w[4], wsum = 0.0;
+ for (int k = 0; k < TARGET_K; k++) { w[k] = exp(lc[k] - mx); wsum += w[k]; }
+ g[0] = g[1] = 0.0;
+ for (int k = 0; k < TARGET_K; k++) {
+ double wk = w[k] / wsum;
+ g[0] += wk * (-(z[0] - TARGET_MU[k][0]) / TARGET_SIGMA2);
+ g[1] += wk * (-(z[1] - TARGET_MU[k][1]) / TARGET_SIGMA2);
+ }
+}
+
+/* ---- Training loop ---- */
+void cnf_train_demo(RNG *r) {
+ const int D = 2, H = 32;
+ const int ITERS = 200, BATCH = 16, LOG_EVERY = 20;
+ const double t0 = 0.0, t1 = 1.0;
+ const double atol = 1e-4, rtol = 1e-4;
+ const double LR = 1e-3;
+
+ int nparams = dynmlp_nparams(D, H);
+ double *theta = vec_alloc(nparams);
+ CNF cnf;
+ cnf_init(&cnf, D, H, theta, r);
+ Adam adam = adam_init(nparams, LR, 0.9, 0.999, 1e-8);
+
+ printf("\n--- CNF density matching (4-Gaussian target, D=2, H=%d) ---\n\n", H);
+ printf("%-6s %-14s %-10s %-10s\n", "Iter", "Loss", "FwdNFE", "BwdNFE");
+ printf("--------------------------------------------------\n");
+
+ for (int iter = 1; iter <= ITERS; iter++) {
+ double *dL_dtheta = vec_zeros(nparams);
+ double loss = 0.0;
+ int total_nfe_fwd = 0, total_nfe_bwd = 0;
+
+ for (int b = 0; b < BATCH; b++) {
+ double z0[2];
+ for (int i = 0; i < D; i++) z0[i] = rng_normal(r);
+
+ CNFSampleResult sr = cnf_sample(&cnf, theta, z0, t0, t1, atol, rtol);
+ double *z1 = sr.z1;
+ total_nfe_fwd += sr.nfe;
+
+ double log_pb = -0.5 * (z0[0]*z0[0] + z0[1]*z0[1])
+ - (double)D * 0.5 * log(2.0 * M_PI);
+
+ double log_pm = log_pb + sr.delta_logp;
+ double log_pt = log_p_target(z1);
+
+ * = log_pb + delta_logp - log_pt
+ * (KL(p_model || p_target) estimator) */
+ loss += log_pm - log_pt;
+
+ double dL_dz1[2];
+ grad_log_p_target(z1, dL_dz1);
+ dL_dz1[0] = -dL_dz1[0];
+ dL_dz1[1] = -dL_dz1[1];
+
+ CNFBackwardResult br = cnf_backward(&cnf, theta, z1,
+ 1.0, dL_dz1,
+ t0, t1, atol, rtol);
+ total_nfe_bwd += br.nfe;
+ vec_add_scaled(dL_dtheta, 1.0 / BATCH, br.dL_dtheta, nparams);
+ free(br.dL_dz0); free(br.dL_dtheta);
+ free(z1);
+ }
+
+ adam_update(&adam, theta, dL_dtheta);
+ free(dL_dtheta);
+
+ if (iter % LOG_EVERY == 0) {
+ printf("%-6d %-14.4f %-10d %-10d\n",
+ iter, loss / BATCH,
+ total_nfe_fwd / BATCH, total_nfe_bwd / BATCH);
+ fflush(stdout);
+ }
+ }
+
+ printf("\n--------------------------------------------------\n");
+
+ const int EVAL_N = 200;
+ double eval_loss = 0.0;
+ for (int i = 0; i < EVAL_N; i++) {
+ double z0[2];
+ for (int j = 0; j < D; j++) z0[j] = rng_normal(r);
+ CNFSampleResult sr = cnf_sample(&cnf, theta, z0, t0, t1, atol, rtol);
+ eval_loss += log_p_target(sr.z1);
+ free(sr.z1);
+ }
+ printf("Eval mean log p_target of samples: %.4f\n", eval_loss / EVAL_N);
+ printf("Total parameters: %d\n", nparams);
+
+ adam_free(&adam);
+ free(theta);
+}
diff --git a/src/main.c b/src/main.c
index 2cae095..291a7bb 100644
--- a/src/main.c
+++ b/src/main.c
@@ -6,6 +6,8 @@
#include "train.h"
#include "spiral.h"
#include "tests.h"
+#include "test_cnf.h"
+#include "cnf_train.h"
#include <stdio.h>
#include <stdlib.h>
@@ -21,6 +23,14 @@ int main(void) {
test_multi_obs_adjoint(&r);
test_training(&r);
+ printf("\n--- CNF tests ---\n");
+ test_cnf_trace(&r);
+ test_cnf_invertibility(&r);
+ test_cnf_gradients(&r);
+ test_cnf_training(&r);
+
+ cnf_train_demo(&r);
+
printf("\n--- Training demo (spiral) ---\n\n");
r = rng_init((uint64_t)time(NULL));
diff --git a/src/test_cnf.c b/src/test_cnf.c
new file mode 100644
index 0000000..da84869
--- /dev/null
+++ b/src/test_cnf.c
@@ -0,0 +1,240 @@
+#include "test_cnf.h"
+#include "cnf.h"
+#include "adjoint.h"
+#include "adam.h"
+
+#include <stdio.h>
+#include <stdlib.h>
+#include <math.h>
+
+/* ---- Test 1: Trace via FD agrees with exact trace for linear f ---- */
+void test_cnf_trace(RNG *r) {
+ const int D = 3, H = 8;
+ const double atol = 1e-8, rtol = 1e-8;
+ const double EPS_FD = 1e-5, TOL = 1e-4;
+
+ int nparams = dynmlp_nparams(D, H);
+ double *theta = vec_alloc(nparams);
+ DynMLP net;
+ dynmlp_init(&net, D, H, theta, r);
+
+ double z[3];
+ for (int i = 0; i < D; i++) z[i] = rng_normal(r);
+ double t = 0.5;
+
+ Workspace ws = workspace_alloc(D, H, nparams);
+ double *z_p = vec_alloc(D), *f_p = vec_alloc(D), *f_m = vec_alloc(D);
+
+ double tr_fd = 0.0;
+ for (int i = 0; i < D; i++) {
+ vec_copy(z, z_p, D);
+ z_p[i] += EPS_FD;
+ dynmlp_forward(&net, theta, z_p, t, f_p, &ws);
+ z_p[i] = z[i] - EPS_FD;
+ dynmlp_forward(&net, theta, z_p, t, f_m, &ws);
+ tr_fd += (f_p[i] - f_m[i]) / (2.0 * EPS_FD);
+ }
+
+ double *out0 = vec_alloc(D);
+ dynmlp_forward(&net, theta, z, t, out0, &ws);
+ double tr_jac = 0.0;
+ for (int i = 0; i < D; i++) {
+ double *col_p = vec_alloc(D), *col_m = vec_alloc(D);
+ vec_copy(z, z_p, D);
+ z_p[i] += EPS_FD;
+ dynmlp_forward(&net, theta, z_p, t, col_p, &ws);
+ z_p[i] = z[i] - EPS_FD;
+ dynmlp_forward(&net, theta, z_p, t, col_m, &ws);
+ /* J[:,i] = (col_p - col_m) / (2*eps); diagonal entry = row i */
+ tr_jac += (col_p[i] - col_m[i]) / (2.0 * EPS_FD);
+ free(col_p); free(col_m);
+ }
+
+ double err = fabs(tr_fd - tr_jac);
+ printf("CNF trace vs full Jacobian trace: err=%.2e tr_fd=%.6f tr_jac=%.6f %s\n",
+ err, tr_fd, tr_jac, err < TOL ? "PASS" : "FAIL");
+
+ CNF cnf;
+ cnf.net = net;
+ cnf.nparams = nparams;
+ cnf.trace_eps = EPS_FD;
+
+ double dt = 1e-4;
+ CNFSampleResult sr = cnf_sample(&cnf, theta, z, 0.0, dt, atol, rtol);
+ double delta_approx = -tr_fd * dt;
+ double err2 = fabs(sr.delta_logp - delta_approx);
+ printf("CNF delta_logp Euler approx: err=%.2e %s\n",
+ err2, err2 < 1e-2 ? "PASS" : "FAIL");
+ free(sr.z1);
+
+ workspace_free(&ws);
+ free(theta); free(z_p); free(f_p); free(f_m); free(out0);
+}
+
+/* ---- Test 2: cnf_sample followed by cnf_log_prob recovers z0 ---- */
+void test_cnf_invertibility(RNG *r) {
+ const int D = 2, H = 16;
+ const double atol = 1e-8, rtol = 1e-8, TOL = 1e-4;
+ const double t0 = 0.0, t1 = 1.0;
+
+ int nparams = dynmlp_nparams(D, H);
+ double *theta = vec_alloc(nparams);
+ CNF cnf;
+ cnf_init(&cnf, D, H, theta, r);
+
+ double z0[2];
+ for (int i = 0; i < D; i++) z0[i] = rng_normal(r);
+
+ CNFSampleResult sr = cnf_sample(&cnf, theta, z0, t0, t1, atol, rtol);
+
+ CNFLogProbResult lr = cnf_log_prob(&cnf, theta, sr.z1, t0, t1, atol, rtol);
+
+ double err = 0.0;
+ for (int i = 0; i < D; i++) {
+ double d = lr.z0[i] - z0[i];
+ err += d * d;
+ }
+ err = sqrt(err);
+
+ double logp_err = fabs(sr.delta_logp - lr.delta_logp);
+
+ printf("CNF invertibility: z0_err=%.2e logp_err=%.2e nfe_fwd=%d nfe_bwd=%d %s\n",
+ err, logp_err, sr.nfe, lr.nfe,
+ (err < TOL && logp_err < TOL) ? "PASS" : "FAIL");
+
+ free(theta); free(sr.z1); free(lr.z0);
+}
+
+/* ---- Test 3: adjoint gradients via finite differences ---- */
+void test_cnf_gradients(RNG *r) {
+ const int D = 2, H = 8;
+ const double EPS = 1e-5, atol = 1e-7, rtol = 1e-7;
+ const double t0 = 0.0, t1 = 0.5;
+
+ int nparams = dynmlp_nparams(D, H);
+ double *theta = vec_alloc(nparams);
+ CNF cnf;
+ cnf_init(&cnf, D, H, theta, r);
+
+ double z0[2];
+ for (int i = 0; i < D; i++) z0[i] = rng_normal(r);
+
+ CNFSampleResult sr = cnf_sample(&cnf, theta, z0, t0, t1, atol, rtol);
+ double *z1 = sr.z1;
+
+ double dL_dlogp = -1.0;
+ double *dL_dz1 = vec_alloc(D);
+ for (int i = 0; i < D; i++) dL_dz1[i] = z1[i];
+
+ CNFBackwardResult br = cnf_backward(&cnf, theta, z1, dL_dlogp, dL_dz1,
+ t0, t1, atol, rtol);
+
+ double *num_grad = vec_alloc(nparams);
+ for (int k = 0; k < nparams; k++) {
+ double tk = theta[k];
+
+ theta[k] = tk + EPS;
+ CNFSampleResult sp = cnf_sample(&cnf, theta, z0, t0, t1, atol, rtol);
+ double lp = 0.0;
+ for (int i = 0; i < D; i++) lp += 0.5 * sp.z1[i] * sp.z1[i];
+ lp -= sp.delta_logp;
+ free(sp.z1);
+
+ theta[k] = tk - EPS;
+ CNFSampleResult sm = cnf_sample(&cnf, theta, z0, t0, t1, atol, rtol);
+ double lm = 0.0;
+ for (int i = 0; i < D; i++) lm += 0.5 * sm.z1[i] * sm.z1[i];
+ lm -= sm.delta_logp;
+ free(sm.z1);
+
+ theta[k] = tk;
+ num_grad[k] = (lp - lm) / (2.0 * EPS);
+ }
+
+ double max_num = 0.0, max_err = 0.0;
+ for (int k = 0; k < nparams; k++) {
+ if (fabs(num_grad[k]) > max_num) max_num = fabs(num_grad[k]);
+ double e = fabs(br.dL_dtheta[k] - num_grad[k]);
+ if (e > max_err) max_err = e;
+ }
+ double rel = max_err / (max_num + 1e-8);
+ printf("CNF adjoint dL/dtheta: max_rel_err=%.2e nfe_bwd=%d %s\n",
+ rel, br.nfe, rel < 1e-2 ? "PASS" : "FAIL");
+
+ free(theta); free(z1); free(dL_dz1);
+ free(br.dL_dz0); free(br.dL_dtheta);
+ free(num_grad);
+}
+
+/* ---- Test 4: training loss decreases on a simple target ---- */
+void test_cnf_training(RNG *r) {
+ const int D = 2, H = 16;
+ const int ITERS = 30;
+ const double t0 = 0.0, t1 = 1.0;
+ const double atol = 1e-4, rtol = 1e-4;
+
+ int nparams = dynmlp_nparams(D, H);
+ double *theta = vec_alloc(nparams);
+ CNF cnf;
+ cnf_init(&cnf, D, H, theta, r);
+ Adam adam = adam_init(nparams, 1e-3, 0.9, 0.999, 1e-8);
+
+ const double mu[2][2] = {{-1.5, 0.0}, {1.5, 0.0}};
+ const double sigma2 = 0.25;
+
+ double first_loss = 0.0, last_loss = 0.0;
+
+ for (int iter = 0; iter < ITERS; iter++) {
+ const int BATCH = 8;
+ double *dL_dtheta_acc = vec_zeros(nparams);
+ double loss = 0.0;
+
+ for (int b = 0; b < BATCH; b++) {
+ double z0[2];
+ for (int i = 0; i < D; i++) z0[i] = rng_normal(r);
+
+ CNFSampleResult sr = cnf_sample(&cnf, theta, z0, t0, t1, atol, rtol);
+ double *z1 = sr.z1;
+
+ double lc[2];
+ for (int k = 0; k < 2; k++) {
+ double dx = z1[0] - mu[k][0], dy = z1[1] - mu[k][1];
+ lc[k] = -0.5 * (dx * dx + dy * dy) / sigma2;
+ }
+ double mx = lc[0] > lc[1] ? lc[0] : lc[1];
+ double log_p_target = mx + log(0.5 * (exp(lc[0] - mx) + exp(lc[1] - mx)));
+
+ double log_p_base = -0.5 * (z0[0]*z0[0] + z0[1]*z0[1])
+ - (double)D * 0.5 * log(2.0 * M_PI);
+
+ loss += -log_p_target + log_p_base + sr.delta_logp;
+
+ double dL_dz1[2];
+ double w0 = exp(lc[0] - mx), w1 = exp(lc[1] - mx);
+ double wsum = w0 + w1;
+ dL_dz1[0] = -(-w0 * (z1[0] - mu[0][0]) / sigma2
+ - w1 * (z1[0] - mu[1][0]) / sigma2) / wsum;
+ dL_dz1[1] = -(-w0 * (z1[1] - mu[0][1]) / sigma2
+ - w1 * (z1[1] - mu[1][1]) / sigma2) / wsum;
+
+ CNFBackwardResult br = cnf_backward(&cnf, theta, z1,
+ 1.0, dL_dz1,
+ t0, t1, atol, rtol);
+ vec_add_scaled(dL_dtheta_acc, 1.0 / BATCH, br.dL_dtheta, nparams);
+ free(br.dL_dz0); free(br.dL_dtheta);
+ free(z1);
+ }
+
+ if (iter == 0) first_loss = loss / BATCH;
+ if (iter == ITERS - 1) last_loss = loss / BATCH;
+
+ adam_update(&adam, theta, dL_dtheta_acc);
+ free(dL_dtheta_acc);
+ }
+
+ printf("CNF training loss: first=%.4f last=%.4f %s\n",
+ first_loss, last_loss, last_loss < first_loss ? "PASS" : "FAIL");
+
+ adam_free(&adam);
+ free(theta);
+}