diff options
| author | y-jan137 <yousefjan24000@gmail.com> | 2026-03-28 15:03:31 +0300 |
|---|---|---|
| committer | y-jan137 <yousefjan24000@gmail.com> | 2026-03-28 15:03:31 +0300 |
| commit | cbf9e2a4d8c7daafda0f8215cac9b0eea751d6d9 (patch) | |
| tree | d1d3b34b253d46537f94b0306e3ddf6c8c2a70a8 | |
| parent | 063b20be98747ba0b288545e653ef9489bf14a54 (diff) | |
| -rw-r--r-- | include/cnf.h | 2 | ||||
| -rw-r--r-- | src/cnf.c | 154 | ||||
| -rw-r--r-- | src/dynnet.c | 34 | ||||
| -rw-r--r-- | src/ode_solver.c | 21 |
4 files changed, 125 insertions, 86 deletions
diff --git a/include/cnf.h b/include/cnf.h index a9ab5bd..ae455f0 100644 --- a/include/cnf.h +++ b/include/cnf.h @@ -7,6 +7,8 @@ typedef struct { DynNet *net; int nparams; double trace_eps; + RNG rng; + int n_hutchinson; /* 0=exact trace (default), 1=Hutchinson estimator */ } CNF; typedef struct { @@ -48,15 +48,13 @@ typedef struct { double *ws; double eps; double a_logp; - double *z_oj; - double *z_p2, *f_p2, *f_m2; - double *neg_a_z; - double *dg_tmp; - double *da_z_tmp; - double *v_ei; - double *vjp_z_tmp; - double *vjp_th_p; - double *vjp_th_m; + int n_hutchinson; + /* Shared scratch */ + double *z_oj, *neg_a_z, *dg_tmp, *vjp_th_p, *vjp_th_m; + /* Exact-trace scratch (n_hutchinson == 0) */ + double *z_p2, *f_p2, *f_m2, *da_z_tmp, *v_ei, *vjp_z_tmp; + /* Hutchinson scratch (n_hutchinson > 0) */ + double *v, *vjp_z_p, *vjp_z_m; } CNFAdjCtx; static void cnf_adj_rhs(const double *aug, double t, const double *params, @@ -78,37 +76,60 @@ static void cnf_adj_rhs(const double *aug, double t, const double *params, if (cc->a_logp == 0.0) return; - 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; + double scale = cc->a_logp / (2.0 * cc->eps); - vec_copy(z, cc->z_oj, D); - cc->z_oj[i] += cc->eps; + if (cc->n_hutchinson > 0) { + /* Hutchinson estimator: tr(J) ≈ v^T J v (Rademacher v in cc->v) + * ∂(v^T J v)/∂z = (J(z+ε*v)^T v − J(z−ε*v)^T v) / (2ε) + * ∂(v^T J v)/∂θ = (VJP_θ(z+ε*v, v) − VJP_θ(z−ε*v, v)) / (2ε) + * Both fall out of 2 dynnet_vjp calls at z±ε*v with tangent v. */ + double *v = cc->v; + for (int i = 0; i < D; i++) cc->z_oj[i] = z[i] + cc->eps * v[i]; vec_zero(cc->vjp_th_p, nparams); - dynnet_vjp(cc->net, cc->theta, cc->z_oj, t, cc->v_ei, - cc->vjp_z_tmp, cc->vjp_th_p, cc->ws); + dynnet_vjp(cc->net, cc->theta, cc->z_oj, t, v, + cc->vjp_z_p, cc->vjp_th_p, cc->ws); - cc->z_oj[i] = z[i] - cc->eps; + for (int i = 0; i < D; i++) cc->z_oj[i] = z[i] - cc->eps * v[i]; vec_zero(cc->vjp_th_m, nparams); - dynnet_vjp(cc->net, cc->theta, cc->z_oj, t, cc->v_ei, - cc->vjp_z_tmp, cc->vjp_th_m, cc->ws); + dynnet_vjp(cc->net, cc->theta, cc->z_oj, t, v, + cc->vjp_z_m, cc->vjp_th_m, cc->ws); - cc->v_ei[i] = 0.0; - - double scale = cc->a_logp / (2.0 * cc->eps); + for (int j = 0; j < D; j++) + out[D + j] += scale * (cc->vjp_z_p[j] - cc->vjp_z_m[j]); for (int k = 0; k < nparams; k++) out[2 * D + k] += scale * (cc->vjp_th_p[k] - cc->vjp_th_m[k]); + } else { + 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); + dynnet_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); + dynnet_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; + + for (int k = 0; k < nparams; k++) + out[2 * D + k] += scale * (cc->vjp_th_p[k] - cc->vjp_th_m[k]); + } } } @@ -122,9 +143,11 @@ void cnf_init(CNF *cnf, int D, int H, double *theta, RNG *r) { dynnet_add_linear(net, D); dynnet_finalize(net); dynnet_init_params(net, theta, r); - cnf->net = net; - cnf->nparams = net->total_params; - cnf->trace_eps = 1e-5; + cnf->net = net; + cnf->nparams = net->total_params; + cnf->trace_eps = 1e-5; + cnf->rng = *r; + cnf->n_hutchinson = 0; } void cnf_free(CNF *cnf) { @@ -195,24 +218,36 @@ CNFBackwardResult cnf_backward(CNF *cnf, const double *theta, double *ws = vec_alloc(cnf->net->total_workspace); 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); + 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.n_hutchinson = cnf->n_hutchinson; + ctx.z_oj = vec_alloc(D); + ctx.neg_a_z = vec_alloc(D); + ctx.dg_tmp = vec_alloc(nparams); + ctx.vjp_th_p = vec_alloc(nparams); + ctx.vjp_th_m = vec_alloc(nparams); + if (cnf->n_hutchinson > 0) { + ctx.v = vec_alloc(D); + ctx.vjp_z_p = vec_alloc(D); + ctx.vjp_z_m = vec_alloc(D); + for (int i = 0; i < D; i++) + ctx.v[i] = (rng_next(&cnf->rng) & 1) ? 1.0 : -1.0; + ctx.z_p2 = ctx.f_p2 = ctx.f_m2 = NULL; + ctx.da_z_tmp = ctx.v_ei = ctx.vjp_z_tmp = NULL; + } else { + ctx.z_p2 = vec_alloc(D); + ctx.f_p2 = vec_alloc(D); + ctx.f_m2 = vec_alloc(D); + ctx.da_z_tmp = vec_alloc(D); + ctx.v_ei = vec_zeros(D); + ctx.vjp_z_tmp = vec_alloc(D); + ctx.v = ctx.vjp_z_p = ctx.vjp_z_m = NULL; + } double *aug = vec_zeros(aug_dim); vec_copy(z1, aug, D); @@ -230,9 +265,14 @@ CNFBackwardResult cnf_backward(CNF *cnf, const double *theta, 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); + free(ctx.z_oj); free(ctx.neg_a_z); free(ctx.dg_tmp); + free(ctx.vjp_th_p); free(ctx.vjp_th_m); + if (cnf->n_hutchinson > 0) { + free(ctx.v); free(ctx.vjp_z_p); free(ctx.vjp_z_m); + } else { + free(ctx.z_p2); free(ctx.f_p2); free(ctx.f_m2); + free(ctx.da_z_tmp); free(ctx.v_ei); free(ctx.vjp_z_tmp); + } free(ws); return out; } diff --git a/src/dynnet.c b/src/dynnet.c index 0f95ef9..c158364 100644 --- a/src/dynnet.c +++ b/src/dynnet.c @@ -71,24 +71,26 @@ static int linear_ws(const void *cfg) { (void)cfg; return 0; } static void tanh_forward(const void *cfg, const double *theta, const double *in, double *out, double *workspace) { - (void)theta; (void)workspace; + (void)theta; const ActCfg *c = (const ActCfg *)cfg; - for (int i = 0; i < c->dim; i++) out[i] = tanh(in[i]); + for (int i = 0; i < c->dim; i++) { workspace[i] = tanh(in[i]); out[i] = workspace[i]; } } static void tanh_vjp(const void *cfg, const double *theta, const double *in, const double *v_out, double *v_in, double *v_theta, double *workspace) { - (void)theta; (void)v_theta; (void)workspace; + (void)theta; (void)v_theta; (void)in; const ActCfg *c = (const ActCfg *)cfg; for (int i = 0; i < c->dim; i++) { - double h = tanh(in[i]); + double h = workspace[i]; v_in[i] = (1.0 - h * h) * v_out[i]; } } static int act_nparams(const void *cfg) { (void)cfg; return 0; } -static int act_ws(const void *cfg) { (void)cfg; return 0; } +static int tanh_ws(const void *cfg) { return ((const ActCfg *)cfg)->dim; } +static int softplus_ws(const void *cfg) { return ((const ActCfg *)cfg)->dim; } +static int swish_ws(const void *cfg) { return ((const ActCfg *)cfg)->dim; } static inline double softplus(double x) { @@ -101,35 +103,35 @@ static inline double sigmoid(double x) { static void softplus_forward(const void *cfg, const double *theta, const double *in, double *out, double *workspace) { - (void)theta; (void)workspace; + (void)theta; const ActCfg *c = (const ActCfg *)cfg; - for (int i = 0; i < c->dim; i++) out[i] = softplus(in[i]); + for (int i = 0; i < c->dim; i++) { workspace[i] = sigmoid(in[i]); out[i] = softplus(in[i]); } } static void softplus_vjp(const void *cfg, const double *theta, const double *in, const double *v_out, double *v_in, double *v_theta, double *workspace) { - (void)theta; (void)v_theta; (void)workspace; + (void)theta; (void)v_theta; (void)in; const ActCfg *c = (const ActCfg *)cfg; for (int i = 0; i < c->dim; i++) - v_in[i] = sigmoid(in[i]) * v_out[i]; + v_in[i] = workspace[i] * v_out[i]; } static void swish_forward(const void *cfg, const double *theta, const double *in, double *out, double *workspace) { - (void)theta; (void)workspace; + (void)theta; const ActCfg *c = (const ActCfg *)cfg; - for (int i = 0; i < c->dim; i++) out[i] = in[i] * sigmoid(in[i]); + for (int i = 0; i < c->dim; i++) { workspace[i] = sigmoid(in[i]); out[i] = in[i] * workspace[i]; } } static void swish_vjp(const void *cfg, const double *theta, const double *in, const double *v_out, double *v_in, double *v_theta, double *workspace) { - (void)theta; (void)v_theta; (void)workspace; + (void)theta; (void)v_theta; const ActCfg *c = (const ActCfg *)cfg; for (int i = 0; i < c->dim; i++) { - double s = sigmoid(in[i]); + double s = workspace[i]; v_in[i] = (s + in[i] * s * (1.0 - s)) * v_out[i]; } } @@ -258,9 +260,9 @@ static int residual_ws(const void *cfg) { static const LayerOps linear_ops = { linear_forward, linear_vjp, linear_nparams, linear_ws }; -static const LayerOps tanh_ops = { tanh_forward, tanh_vjp, act_nparams, act_ws }; -static const LayerOps softplus_ops = { softplus_forward, softplus_vjp, act_nparams, act_ws }; -static const LayerOps swish_ops = { swish_forward, swish_vjp, act_nparams, act_ws }; +static const LayerOps tanh_ops = { tanh_forward, tanh_vjp, act_nparams, tanh_ws }; +static const LayerOps softplus_ops = { softplus_forward, softplus_vjp, act_nparams, softplus_ws }; +static const LayerOps swish_ops = { swish_forward, swish_vjp, act_nparams, swish_ws }; static const LayerOps layernorm_ops = { layernorm_forward, layernorm_vjp, layernorm_nparams, layernorm_ws }; static const LayerOps time_concat_ops = { time_concat_forward, time_concat_vjp, time_concat_nparams, time_concat_ws }; static const LayerOps residual_ops = { residual_forward, residual_vjp, residual_nparams, residual_ws }; diff --git a/src/ode_solver.c b/src/ode_solver.c index 566cb2a..6c6963d 100644 --- a/src/ode_solver.c +++ b/src/ode_solver.c @@ -24,12 +24,13 @@ static const double dp_e[7] = { ODEResult ode_solve(ode_rhs_fn f, const double *y0, double t0, double t1, const double *params, int dim, double atol, double rtol, void *ctx) { - double **k = (double **)xmalloc(7 * sizeof(double *)); - for (int i = 0; i < 7; i++) k[i] = vec_alloc(dim); - double *y = vec_alloc(dim); - double *y5 = vec_alloc(dim); - double *err = vec_alloc(dim); - double *stg = vec_alloc(dim); + double *buf = vec_alloc(11 * dim); + double *k[7]; + for (int i = 0; i < 7; i++) k[i] = buf + i * dim; + double *y = buf + 7 * dim; + double *y5 = buf + 8 * dim; + double *err = buf + 9 * dim; + double *stg = buf + 10 * dim; ODEResult res = { vec_alloc(dim), 0 }; vec_copy(y0, y, dim); @@ -109,13 +110,7 @@ ODEResult ode_solve(ode_rhs_fn f, const double *y0, double t0, double t1, } vec_copy(y, res.y, dim); - for (int i = 0; i < 7; i++) free(k[i]); - - free(k); - free(y); - free(y5); - free(err); - free(stg); + free(buf); return res; } |