From 161352cf4504c2113ad148ce5b5917a5bed36c06 Mon Sep 17 00:00:00 2001 From: y-jan137 Date: Wed, 18 Mar 2026 11:44:05 +0300 Subject: Add continuous normalizing flow --- include/cnf.h | 45 +++++++++++++++++++++++++++++++++++++++++++++ include/cnf_train.h | 5 +++++ include/test_cnf.h | 8 ++++++++ 3 files changed, 58 insertions(+) create mode 100644 include/cnf.h create mode 100644 include/cnf_train.h create mode 100644 include/test_cnf.h (limited to 'include') 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); -- cgit v1.2.3