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 --- src/cnf_train.c | 16 +++++++++++++--- 1 file changed, 13 insertions(+), 3 deletions(-) (limited to 'src/cnf_train.c') diff --git a/src/cnf_train.c b/src/cnf_train.c index 21b357f..d88aa3c 100644 --- a/src/cnf_train.c +++ b/src/cnf_train.c @@ -51,7 +51,17 @@ void cnf_train_demo(RNG *r) { const double atol = 1e-4, rtol = 1e-4; const double LR = 1e-3; - int nparams = dynmlp_nparams(D, H); + int nparams = 0; + { + DynNet *tmp = dynnet_create(D); + dynnet_add_time_concat(tmp); + dynnet_add_linear(tmp, H); + dynnet_add_tanh(tmp); + dynnet_add_linear(tmp, D); + dynnet_finalize(tmp); + nparams = tmp->total_params; + dynnet_free(tmp); + } double *theta = vec_alloc(nparams); CNF cnf; cnf_init(&cnf, D, H, theta, r); @@ -80,8 +90,7 @@ void cnf_train_demo(RNG *r) { 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) */ + /* KL(p_model || p_target) estimator */ loss += log_pm - log_pt; double dL_dz1[2]; @@ -124,5 +133,6 @@ void cnf_train_demo(RNG *r) { printf("Total parameters: %d\n", nparams); adam_free(&adam); + cnf_free(&cnf); free(theta); } -- cgit v1.2.3