From 063b20be98747ba0b288545e653ef9489bf14a54 Mon Sep 17 00:00:00 2001 From: y-jan137 Date: Wed, 18 Mar 2026 14:04:47 +0300 Subject: Replace DynMLP with DynNet --- include/cnf.h | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) (limited to 'include/cnf.h') diff --git a/include/cnf.h b/include/cnf.h index 5dc5688..a9ab5bd 100644 --- a/include/cnf.h +++ b/include/cnf.h @@ -1,44 +1,44 @@ #pragma once -#include "dynmlp.h" +#include "dynnet.h" #include "ode_solver.h" -#include "utils.h" typedef struct { - DynMLP net; + DynNet *net; int nparams; - double trace_eps; /* finite-difference epsilon for trace computation */ + double trace_eps; } CNF; typedef struct { - double *z1; + double *z1; double delta_logp; int nfe; } CNFSampleResult; typedef struct { - double *z0; - double delta_logp; + double *z0; + double delta_logp; int nfe; } CNFLogProbResult; typedef struct { - double *dL_dz0; - double *dL_dtheta; + double *dL_dz0; + double *dL_dtheta; int nfe; } CNFBackwardResult; void cnf_init(CNF *cnf, int D, int H, double *theta, RNG *r); +void cnf_free(CNF *cnf); -CNFSampleResult cnf_sample(const CNF *cnf, const double *theta, +CNFSampleResult cnf_sample(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, +CNFLogProbResult cnf_log_prob(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, +CNFBackwardResult cnf_backward(CNF *cnf, const double *theta, const double *z1, double dL_dlogp, const double *dL_dz1_in, double t0, double t1, -- cgit v1.2.3