Skip to content

Commit 2913dee

Browse files
Transurgeonclaude
andcommitted
Add parameter support to C diff engine
Add parameter node type and parameter-aware variants of scalar mult, vector mult, and left matmul. Parameters store an offset into a global theta vector and can be updated via problem_update_params without rebuilding the expression tree. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
1 parent fa65481 commit 2913dee

9 files changed

Lines changed: 346 additions & 9 deletions

File tree

‎include/affine.h‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@ expr *new_trace(expr *child);
3434

3535
expr *new_constant(int d1, int d2, int n_vars, const double *values);
3636
expr *new_variable(int d1, int d2, int var_id, int n_vars);
37+
expr *new_parameter(int d1, int d2, int param_id, int n_vars);
3738

3839
expr *new_index(expr *child, int d1, int d2, const int *indices, int n_idxs);
3940
expr *new_reshape(expr *child, int d1, int d2);

‎include/bivariate.h‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,4 +42,13 @@ expr *new_const_scalar_mult(double a, expr *child);
4242
/* Constant vector elementwise multiplication: a ∘ f(x) where a is constant */
4343
expr *new_const_vector_mult(const double *a, expr *child);
4444

45+
/* Left matrix multiplication with parameter source: P @ f(x) where P is a parameter */
46+
expr *new_left_param_matmul(expr *param_node, expr *u, int A_m, int A_n);
47+
48+
/* Parameter scalar multiplication: p * f(x) where p is a parameter */
49+
expr *new_param_scalar_mult(expr *param_node, expr *child);
50+
51+
/* Parameter vector elementwise multiplication: p ∘ f(x) where p is a parameter */
52+
expr *new_param_vector_mult(expr *param_node, expr *child);
53+
4554
#endif /* BIVARIATE_H */

‎include/problem.h‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -59,6 +59,11 @@ typedef struct problem
5959
* hessian are called */
6060
bool jacobian_called;
6161

62+
/* Parameter tracking for fast parameter updates */
63+
expr **param_nodes; /* weak references to parameter nodes in tree */
64+
int n_param_nodes;
65+
int n_params; /* total scalar parameters */
66+
6267
/* Statistics for performance measurement */
6368
Diff_engine_stats stats;
6469
bool verbose;
@@ -78,4 +83,9 @@ void problem_gradient(problem *prob);
7883
void problem_jacobian(problem *prob);
7984
void problem_hessian(problem *prob, double obj_w, const double *w);
8085

86+
/* Parameter support */
87+
void problem_register_params(problem *prob, expr **param_nodes,
88+
int n_param_nodes, int n_params);
89+
void problem_update_params(problem *prob, const double *theta);
90+
8191
#endif

‎include/subexpr.h‎

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,13 @@
2525
/* Forward declaration */
2626
struct int_double_pair;
2727

28+
/* Parameter node: like constant but with updatable values via problem_update_params */
29+
typedef struct parameter_expr
30+
{
31+
expr base;
32+
int param_id; /* offset into global theta vector */
33+
} parameter_expr;
34+
2835
/* Type-specific expression structures that "inherit" from expr */
2936

3037
/* Linear operator: y = A * x + b */
@@ -110,6 +117,8 @@ typedef struct left_matmul_expr
110117
CSR_Matrix *A;
111118
CSR_Matrix *AT;
112119
CSC_Matrix *CSC_work;
120+
expr *param_source; /* if non-NULL, refresh A/AT values from param_source->value */
121+
int src_m, src_n; /* original (non-block-diag) matrix dimensions */
113122
} left_matmul_expr;
114123

115124
/* Right matrix multiplication: y = f(x) * A where f(x) is an expression.
@@ -128,13 +137,15 @@ typedef struct const_scalar_mult_expr
128137
{
129138
expr base;
130139
double a;
140+
expr *param_source; /* if non-NULL, read a from param_source->value[0] */
131141
} const_scalar_mult_expr;
132142

133143
/* Constant vector elementwise multiplication: y = a \circ child for constant a */
134144
typedef struct const_vector_mult_expr
135145
{
136146
expr base;
137-
double *a; /* length equals node->size */
147+
double *a; /* length equals node->size */
148+
expr *param_source; /* if non-NULL, use param_source->value instead of a */
138149
} const_vector_mult_expr;
139150

140151
/* Index/slicing: y = child[indices] where indices is a list of flat positions */

‎src/affine/parameter.c‎

Lines changed: 74 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,74 @@
1+
/*
2+
* Copyright 2026 Daniel Cederberg and William Zhang
3+
*
4+
* This file is part of the DNLP-differentiation-engine project.
5+
*
6+
* Licensed under the Apache License, Version 2.0 (the "License");
7+
* you may not use this file except in compliance with the License.
8+
* You may obtain a copy of the License at
9+
*
10+
* http://www.apache.org/licenses/LICENSE-2.0
11+
*
12+
* Unless required by applicable law or agreed to in writing, software
13+
* distributed under the License is distributed on an "AS IS" BASIS,
14+
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
15+
* See the License for the specific language governing permissions and
16+
* limitations under the License.
17+
*/
18+
19+
/* Parameter leaf node: behaviorally identical to constant (zero derivatives
20+
w.r.t. variables), but its values are updatable via problem_update_params.
21+
This allows re-solving with different parameter values without rebuilding
22+
the expression tree. */
23+
24+
#include "affine.h"
25+
#include "subexpr.h"
26+
#include <stdlib.h>
27+
#include <string.h>
28+
29+
static void forward(expr *node, const double *u)
30+
{
31+
/* Values are set by problem_update_params, not by forward pass */
32+
(void)node;
33+
(void)u;
34+
}
35+
36+
static void jacobian_init(expr *node)
37+
{
38+
/* Parameter jacobian is all zeros: size x n_vars with 0 nonzeros */
39+
node->jacobian = new_csr_matrix(node->size, node->n_vars, 0);
40+
}
41+
42+
static void eval_jacobian(expr *node)
43+
{
44+
/* Parameter jacobian never changes */
45+
(void)node;
46+
}
47+
48+
static void wsum_hess_init(expr *node)
49+
{
50+
/* Parameter Hessian is all zeros */
51+
node->wsum_hess = new_csr_matrix(node->n_vars, node->n_vars, 0);
52+
}
53+
54+
static void eval_wsum_hess(expr *node, const double *w)
55+
{
56+
(void)node;
57+
(void)w;
58+
}
59+
60+
static bool is_affine(const expr *node)
61+
{
62+
(void)node;
63+
return true;
64+
}
65+
66+
expr *new_parameter(int d1, int d2, int param_id, int n_vars)
67+
{
68+
parameter_expr *pnode = (parameter_expr *)calloc(1, sizeof(parameter_expr));
69+
init_expr(&pnode->base, d1, d2, n_vars, forward, jacobian_init, eval_jacobian,
70+
is_affine, wsum_hess_init, eval_wsum_hess, NULL);
71+
pnode->param_id = param_id;
72+
/* values will be populated by problem_update_params */
73+
return &pnode->base;
74+
}

‎src/bivariate/const_scalar_mult.c‎

Lines changed: 26 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,11 @@
2424

2525
/* Constant scalar multiplication: y = a * child where a is a constant double */
2626

27+
static inline double get_scalar(const const_scalar_mult_expr *sn)
28+
{
29+
return sn->param_source ? sn->param_source->value[0] : sn->a;
30+
}
31+
2732
static void forward(expr *node, const double *u)
2833
{
2934
expr *child = node->left;
@@ -32,7 +37,7 @@ static void forward(expr *node, const double *u)
3237
child->forward(child, u);
3338

3439
/* local forward pass: multiply each element by scalar a */
35-
double a = ((const_scalar_mult_expr *) node)->a;
40+
double a = get_scalar((const_scalar_mult_expr *) node);
3641
for (int i = 0; i < node->size; i++)
3742
{
3843
node->value[i] = a * child->value[i];
@@ -55,7 +60,7 @@ static void jacobian_init(expr *node)
5560
static void eval_jacobian(expr *node)
5661
{
5762
expr *child = node->left;
58-
double a = ((const_scalar_mult_expr *) node)->a;
63+
double a = get_scalar((const_scalar_mult_expr *) node);
5964

6065
/* evaluate child */
6166
child->eval_jacobian(child);
@@ -85,7 +90,7 @@ static void eval_wsum_hess(expr *node, const double *w)
8590
expr *x = node->left;
8691
x->eval_wsum_hess(x, w);
8792

88-
double a = ((const_scalar_mult_expr *) node)->a;
93+
double a = get_scalar((const_scalar_mult_expr *) node);
8994
for (int j = 0; j < x->wsum_hess->nnz; j++)
9095
{
9196
node->wsum_hess->x[j] = a * x->wsum_hess->x[j];
@@ -108,7 +113,25 @@ expr *new_const_scalar_mult(double a, expr *child)
108113
eval_jacobian, is_affine, wsum_hess_init, eval_wsum_hess, NULL);
109114
node->left = child;
110115
mult_node->a = a;
116+
mult_node->param_source = NULL;
117+
expr_retain(child);
118+
119+
return node;
120+
}
121+
122+
expr *new_param_scalar_mult(expr *param_node, expr *child)
123+
{
124+
const_scalar_mult_expr *mult_node =
125+
(const_scalar_mult_expr *) calloc(1, sizeof(const_scalar_mult_expr));
126+
expr *node = &mult_node->base;
127+
128+
init_expr(node, child->d1, child->d2, child->n_vars, forward, jacobian_init,
129+
eval_jacobian, is_affine, wsum_hess_init, eval_wsum_hess, NULL);
130+
node->left = child;
131+
mult_node->a = param_node->value[0]; /* initial value */
132+
mult_node->param_source = param_node;
111133
expr_retain(child);
134+
expr_retain(param_node);
112135

113136
return node;
114137
}

‎src/bivariate/const_vector_mult.c‎

Lines changed: 41 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -23,10 +23,15 @@
2323

2424
/* Constant vector elementwise multiplication: y = a \circ child */
2525

26+
static inline const double *get_vector(const const_vector_mult_expr *vn)
27+
{
28+
return vn->param_source ? vn->param_source->value : vn->a;
29+
}
30+
2631
static void forward(expr *node, const double *u)
2732
{
2833
expr *child = node->left;
29-
const double *a = ((const_vector_mult_expr *) node)->a;
34+
const double *a = get_vector((const_vector_mult_expr *) node);
3035

3136
/* child's forward pass */
3237
child->forward(child, u);
@@ -54,7 +59,7 @@ static void jacobian_init(expr *node)
5459
static void eval_jacobian(expr *node)
5560
{
5661
expr *x = node->left;
57-
const double *a = ((const_vector_mult_expr *) node)->a;
62+
const double *a = get_vector((const_vector_mult_expr *) node);
5863

5964
/* evaluate x */
6065
x->eval_jacobian(x);
@@ -87,7 +92,7 @@ static void wsum_hess_init(expr *node)
8792
static void eval_wsum_hess(expr *node, const double *w)
8893
{
8994
expr *x = node->left;
90-
const double *a = ((const_vector_mult_expr *) node)->a;
95+
const double *a = get_vector((const_vector_mult_expr *) node);
9196

9297
/* scale weights w by a */
9398
for (int i = 0; i < node->size; i++)
@@ -128,6 +133,39 @@ expr *new_const_vector_mult(const double *a, expr *child)
128133
/* copy a vector */
129134
vnode->a = (double *) malloc(child->size * sizeof(double));
130135
memcpy(vnode->a, a, child->size * sizeof(double));
136+
vnode->param_source = NULL;
137+
138+
return node;
139+
}
140+
141+
static void free_param_type_data(expr *node)
142+
{
143+
const_vector_mult_expr *vnode = (const_vector_mult_expr *) node;
144+
/* a is not owned when param_source is set */
145+
free(vnode->a);
146+
if (vnode->param_source)
147+
{
148+
free_expr(vnode->param_source);
149+
}
150+
}
151+
152+
expr *new_param_vector_mult(expr *param_node, expr *child)
153+
{
154+
const_vector_mult_expr *vnode =
155+
(const_vector_mult_expr *) calloc(1, sizeof(const_vector_mult_expr));
156+
expr *node = &vnode->base;
157+
158+
init_expr(node, child->d1, child->d2, child->n_vars, forward, jacobian_init,
159+
eval_jacobian, is_affine, wsum_hess_init, eval_wsum_hess,
160+
free_param_type_data);
161+
node->left = child;
162+
expr_retain(child);
163+
164+
/* Still allocate a copy for initial values (used before first update_params) */
165+
vnode->a = (double *) malloc(child->size * sizeof(double));
166+
memcpy(vnode->a, param_node->value, child->size * sizeof(double));
167+
vnode->param_source = param_node;
168+
expr_retain(param_node);
131169

132170
return node;
133171
}

0 commit comments

Comments
 (0)