aboutsummaryrefslogtreecommitdiff
path: root/include
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 /include
parent00cd5b38f26e40fbb0dd21e9682e9a9580fe3e78 (diff)
Add continuous normalizing flow
Diffstat (limited to 'include')
-rw-r--r--include/cnf.h45
-rw-r--r--include/cnf_train.h5
-rw-r--r--include/test_cnf.h8
3 files changed, 58 insertions, 0 deletions
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);