aboutsummaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
authory-jan137 <yousefjan24000@gmail.com>2026-03-28 15:03:31 +0300
committery-jan137 <yousefjan24000@gmail.com>2026-03-28 15:03:31 +0300
commitcbf9e2a4d8c7daafda0f8215cac9b0eea751d6d9 (patch)
treed1d3b34b253d46537f94b0306e3ddf6c8c2a70a8 /src
parent063b20be98747ba0b288545e653ef9489bf14a54 (diff)
Optimization passHEADmain
Diffstat (limited to 'src')
-rw-r--r--src/cnf.c154
-rw-r--r--src/dynnet.c34
-rw-r--r--src/ode_solver.c21
3 files changed, 123 insertions, 86 deletions
diff --git a/src/cnf.c b/src/cnf.c
index 08d8dd8..2864598 100644
--- a/src/cnf.c
+++ b/src/cnf.c
@@ -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;
}