diff options
| -rw-r--r-- | README.md | 2 | ||||
| -rw-r--r-- | include/cnf.h | 45 | ||||
| -rw-r--r-- | include/cnf_train.h | 5 | ||||
| -rw-r--r-- | include/test_cnf.h | 8 | ||||
| -rw-r--r-- | src/cnf.c | 233 | ||||
| -rw-r--r-- | src/cnf_train.c | 128 | ||||
| -rw-r--r-- | src/main.c | 10 | ||||
| -rw-r--r-- | src/test_cnf.c | 240 |
8 files changed, 670 insertions, 1 deletions
@@ -4,7 +4,7 @@ A implementation of [Neural Ordinary Differential Equations by Chen et al. (2018 ## Build ```bash -cc -O2 -Wall -Wextra -Iinclude src/utils.c src/dynmlp.c src/ode_solver.c src/adjoint.c src/adam.c src/train.c src/spiral.c src/tests.c src/main.c -lm -o neural_ode +cc -O2 -Wall -Wextra -Iinclude src/utils.c src/dynmlp.c src/ode_solver.c src/adjoint.c src/adam.c src/train.c src/spiral.c src/tests.c src/cnf.c src/cnf_train.c src/test_cnf.c src/main.c -lm -o neural_ode ``` ## Run diff --git a/include/cnf.h b/include/cnf.h new file mode 100644 index 0000000..5dc5688 --- /dev/null +++ b/include/cnf.h @@ -0,0 +1,45 @@ +#pragma once + +#include "dynmlp.h" +#include "ode_solver.h" +#include "utils.h" + +typedef struct { + DynMLP net; + int nparams; + double trace_eps; /* finite-difference epsilon for trace computation */ +} CNF; + +typedef struct { + double *z1; + double delta_logp; + int nfe; +} CNFSampleResult; + +typedef struct { + double *z0; + double delta_logp; + int nfe; +} CNFLogProbResult; + +typedef struct { + double *dL_dz0; + double *dL_dtheta; + int nfe; +} CNFBackwardResult; + +void cnf_init(CNF *cnf, int D, int H, double *theta, RNG *r); + +CNFSampleResult cnf_sample(const 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, + const double *z1, double t0, double t1, + double atol, double rtol); + +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); diff --git a/include/cnf_train.h b/include/cnf_train.h new file mode 100644 index 0000000..6b0ad09 --- /dev/null +++ b/include/cnf_train.h @@ -0,0 +1,5 @@ +#pragma once + +#include "utils.h" + +void cnf_train_demo(RNG *r); diff --git a/include/test_cnf.h b/include/test_cnf.h new file mode 100644 index 0000000..4960f0d --- /dev/null +++ b/include/test_cnf.h @@ -0,0 +1,8 @@ +#pragma once + +#include "utils.h" + +void test_cnf_trace(RNG *r); +void test_cnf_invertibility(RNG *r); +void test_cnf_gradients(RNG *r); +void test_cnf_training(RNG *r); 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); +} @@ -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); +} |