Skip to content

Commit 5cf9d7e

Browse files
dance858claude
andauthored
values_version cache coherence for Jacobian/Hessian value caches (#113)
Replace the caller-must-refresh temporal contracts and the per-atom jacobian_csc_filled flags with a monotone values_version counter on matrix: writers bump, consumers refresh iff their recorded version differs. Staleness becomes impossible by construction and shared-operand double-refreshes dedupe for free. - eval_jacobian/eval_wsum_hess free-function wrappers run the atom impl (slots renamed *_impl) and bump the output's version; the jacobian wrapper skips the bump for an affine node already evaluated this parameter epoch, preserving today's refresh counts. - sparse_matrix csc_cache and stacked_pd csr_cache are version-guarded (the latter skips its per-call block memcpy in to_csr). - spd vtable fill adapters and the raw hess_term2 write sites bump their outputs so version-guarded reads stay fresh. - expr_set_needs_refresh gains a set_needs_refresh_children hook so the parameter-refresh walk reaches hstack's args[] children; without it a parameter-dependent affine child under hstack/vstack served stale spd Jacobian values after a parameter update (also fixes the pre-existing stale-CSC-mirror bug for args[] children on main). - new tests lock the semantics: bump-per-eval, affine bump-skip and re-arm, CSC mirror dedup, spd to_csr freshness, spd hess-term staleness regression, and the hstack parameter-refresh regression. Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
1 parent a887344 commit 5cf9d7e

110 files changed

Lines changed: 737 additions & 409 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎include/expr.h‎

Lines changed: 31 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
#include "utils/matrix.h"
2424
#include <stdbool.h>
2525
#include <stddef.h>
26+
#include <stdint.h>
2627
#include <string.h>
2728

2829
#define JAC_IDXS_NOT_SET -1
@@ -39,6 +40,7 @@ typedef void (*local_jacobian_fn)(struct expr *node, double *out);
3940
typedef void (*local_wsum_hess_fn)(struct expr *node, double *out, const double *w);
4041
typedef bool (*is_affine_fn)(const struct expr *node);
4142
typedef void (*free_type_data_fn)(struct expr *node);
43+
typedef void (*set_needs_refresh_children_fn)(struct expr *node);
4244

4345
/* Workspace for derivative computation */
4446
typedef struct
@@ -48,10 +50,19 @@ typedef struct
4850
CSC_matrix *jacobian_csc;
4951
int *csc_work; /* for CSR_matrix-CSC_matrix conversion */
5052

51-
/* jacobian_csc_filled is only used for affine functions to avoid redundant
52-
conversions. Could become relevant for non-affine functions if we start
53-
supporting common subexpressions on the Python side. */
54-
bool jacobian_csc_filled;
53+
/* jacobian->values_version that the jacobian_csc mirror reflects;
54+
expr_refresh_jacobian_csc refills iff it differs. */
55+
uint64_t jacobian_csc_seen;
56+
57+
/* node->is_affine(node), computed once by jacobian_init (affinity is
58+
structural, and the recursive is_affine is too costly per eval). */
59+
bool is_affine_cached;
60+
61+
/* True once eval_jacobian has run this parameter epoch; cleared by
62+
expr_set_needs_refresh. Only consulted for affine nodes, where it
63+
lets the eval_jacobian wrapper skip the values_version bump (same
64+
role the old jacobian_csc_filled latch played). */
65+
bool jacobian_evaluated;
5566
double *local_jac_diag; /* cached f'(g(x)) diagonal */
5667
matrix *hess_term1; /* Jg^T D Jg workspace */
5768
matrix *hess_term2; /* child wsum_hess workspace */
@@ -76,8 +87,8 @@ typedef struct expr
7687
forward_fn forward;
7788
jacobian_init_fn jacobian_init_impl;
7889
wsum_hess_init_fn wsum_hess_init_impl;
79-
eval_jacobian_fn eval_jacobian;
80-
wsum_hess_fn eval_wsum_hess;
90+
eval_jacobian_fn eval_jacobian_impl;
91+
wsum_hess_fn eval_wsum_hess_impl;
8192

8293
// ------------------------------------------------------------------------
8394
// other things
@@ -86,7 +97,11 @@ typedef struct expr
8697
local_jacobian_fn local_jacobian; /* used by elementwise univariate atoms*/
8798
local_wsum_hess_fn local_wsum_hess; /* used by elementwise univariate atoms*/
8899
free_type_data_fn free_type_data; /* Cleanup for type-specific fields */
89-
Expr_Work *work; /* derivative workspace */
100+
/* Recursion hook for expr_set_needs_refresh: atoms holding children
101+
outside left/right (hstack's args[]) set this so the parameter-refresh
102+
walk reaches them. NULL for binary/unary atoms. */
103+
set_needs_refresh_children_fn set_needs_refresh_children;
104+
Expr_Work *work; /* derivative workspace */
90105
/* Set to true on all nodes by problem_update_params() via
91106
expr_set_needs_refresh(). Atoms that cache parameter data
92107
(e.g. left_matmul_dense) check this flag before their forward
@@ -111,6 +126,15 @@ void free_expr(expr *node);
111126
void jacobian_init(expr *node);
112127
void wsum_hess_init(expr *node);
113128

129+
/* Eval wrappers: run the atom's eval_*_impl and bump the output matrix's
130+
* values_version so version-guarded caches (CSC mirrors, spd CSR views)
131+
* refresh. Always call these instead of the impl slots. */
132+
void eval_jacobian(expr *node);
133+
void eval_wsum_hess(expr *node, const double *w);
134+
135+
/* Refresh work->jacobian_csc from node->jacobian iff its values changed. */
136+
void expr_refresh_jacobian_csc(expr *node);
137+
114138
/* Initialize CSC_matrix form of the Jacobian from the CSR_matrix Jacobian.
115139
* Must be called after jacobian_init. */
116140
void jacobian_csc_init(expr *node);

‎include/utils/matrix.h‎

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
#include "CSC_matrix.h"
2222
#include "CSR_matrix.h"
2323
#include <stdbool.h>
24+
#include <stdint.h>
2425

2526
/* Broadcast shape used by the broadcast atom and its vtable methods. */
2627
typedef enum
@@ -77,7 +78,8 @@ typedef void (*matrix_transpose_fill_values_fn)(const matrix *A, matrix *AT);
7778
typedef CSR_matrix *(*matrix_to_csr_fn)(matrix *A);
7879

7980
/* Refresh any internal caches (e.g. a CSC_matrix mirror) so subsequent ATA /
80-
ATDA calls reflect the current values. */
81+
ATDA calls reflect the current values. Version-guarded: a no-op when the
82+
cache already matches values_version, so it is cheap to call when fresh. */
8183
typedef void (*matrix_refresh_csc_values_fn)(matrix *A);
8284

8385
/* Allocate C = A[indices, :] */
@@ -128,6 +130,14 @@ struct matrix
128130
bool is_permuted_dense;
129131
bool is_stacked_pd;
130132

133+
/* Monotone counter bumped whenever the matrix's values change. Consumers
134+
that mirror the values into a cache (CSC mirror, CSR view, ...) record
135+
the version they last saw and refresh iff it differs. Code that writes
136+
x directly must call matrix_values_changed on the OWNER of the buffer —
137+
aliased children (spd blocks, cache views) have no version of their
138+
own. */
139+
uint64_t values_version;
140+
131141
/* Operator ops */
132142
matrix_block_left_mult_vec_fn block_left_mult_vec;
133143
matrix_block_left_mult_sparsity_fn block_left_mult_sparsity;
@@ -160,6 +170,12 @@ struct matrix
160170
matrix_free_fn free_fn;
161171
};
162172

173+
/* Notify the library after writing A->x directly. */
174+
static inline void matrix_values_changed(matrix *A)
175+
{
176+
A->values_version++;
177+
}
178+
163179
/* Free helper */
164180
static inline void free_matrix(matrix *m)
165181
{

‎include/utils/sparse_matrix.h‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@ typedef struct sparse_matrix
2828
matrix base;
2929
CSR_matrix *csr;
3030
CSC_matrix *csc_cache;
31+
uint64_t csc_seen; /* base.values_version the csc_cache values reflect */
3132
int *csc_iwork;
3233
int *transpose_iwork; /* sized csr->n; allocated by sparse_transpose_alloc
3334
on the output sm and reused by

‎include/utils/stacked_pd.h‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,7 @@ typedef struct stacked_pd
4848

4949
/* lazily built CSR view */
5050
CSR_matrix *csr_cache;
51+
uint64_t csr_seen; /* base.values_version the csr_cache values reflect */
5152

5253
/* Private permuted_dense scratch owned by the kernel that produced
5354
this spd. Allocated by the producing _alloc, used (without

‎src/atoms/affine/add.c‎

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -50,11 +50,11 @@ static void jacobian_init_impl(expr *node)
5050
sum_matrices_alloc(node->left->jacobian, node->right->jacobian, node->jacobian);
5151
}
5252

53-
static void eval_jacobian(expr *node)
53+
static void eval_jacobian_impl(expr *node)
5454
{
5555
/* evaluate children's jacobians */
56-
node->left->eval_jacobian(node->left);
57-
node->right->eval_jacobian(node->right);
56+
eval_jacobian(node->left);
57+
eval_jacobian(node->right);
5858

5959
/* sum children's jacobians */
6060
sum_matrices_fill_values(node->left->jacobian, node->right->jacobian,
@@ -76,11 +76,11 @@ static void wsum_hess_init_impl(expr *node)
7676
node->wsum_hess);
7777
}
7878

79-
static void eval_wsum_hess(expr *node, const double *w)
79+
static void eval_wsum_hess_impl(expr *node, const double *w)
8080
{
8181
/* evaluate children's wsum_hess */
82-
node->left->eval_wsum_hess(node->left, w);
83-
node->right->eval_wsum_hess(node->right, w);
82+
eval_wsum_hess(node->left, w);
83+
eval_wsum_hess(node->right, w);
8484

8585
/* sum children's wsum_hess */
8686
sum_matrices_fill_values(node->left->wsum_hess, node->right->wsum_hess,
@@ -97,7 +97,8 @@ expr *new_add(expr *left, expr *right)
9797
assert(left->d1 == right->d1 && left->d2 == right->d2);
9898
expr *node = (expr *) sp_calloc(1, sizeof(expr));
9999
init_expr(node, left->d1, left->d2, left->n_vars, forward, jacobian_init_impl,
100-
eval_jacobian, is_affine, wsum_hess_init_impl, eval_wsum_hess, NULL);
100+
eval_jacobian_impl, is_affine, wsum_hess_init_impl,
101+
eval_wsum_hess_impl, NULL);
101102
node->left = left;
102103
node->right = right;
103104
expr_retain(left);

‎src/atoms/affine/broadcast.c‎

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -77,9 +77,9 @@ static void jacobian_init_impl(expr *node)
7777
x->jacobian->broadcast_alloc(x->jacobian, bcast->type, node->d1, node->d2);
7878
}
7979

80-
static void eval_jacobian(expr *node)
80+
static void eval_jacobian_impl(expr *node)
8181
{
82-
node->left->eval_jacobian(node->left);
82+
eval_jacobian(node->left);
8383

8484
/* fill values into the preallocated output. */
8585
broadcast_expr *bcast = (broadcast_expr *) node;
@@ -99,7 +99,7 @@ static void wsum_hess_init_impl(expr *node)
9999
node->work->dwork = sp_malloc(node->size * sizeof(double));
100100
}
101101

102-
static void eval_wsum_hess(expr *node, const double *w)
102+
static void eval_wsum_hess_impl(expr *node, const double *w)
103103
{
104104
broadcast_expr *bcast = (broadcast_expr *) node;
105105
expr *x = node->left;
@@ -139,7 +139,7 @@ static void eval_wsum_hess(expr *node, const double *w)
139139
}
140140
}
141141

142-
x->eval_wsum_hess(x, node->work->dwork);
142+
eval_wsum_hess(x, node->work->dwork);
143143
memcpy(node->wsum_hess->x, x->wsum_hess->x,
144144
node->wsum_hess->nnz * sizeof(double));
145145
}
@@ -183,7 +183,8 @@ expr *new_broadcast(expr *child, int d1, int d2)
183183
// initialize the rest of the expression
184184
// --------------------------------------------------------------------------
185185
init_expr(node, d1, d2, child->n_vars, forward, jacobian_init_impl,
186-
eval_jacobian, is_affine, wsum_hess_init_impl, eval_wsum_hess, NULL);
186+
eval_jacobian_impl, is_affine, wsum_hess_init_impl,
187+
eval_wsum_hess_impl, NULL);
187188
node->left = child;
188189
expr_retain(child);
189190
bcast->type = type;

‎src/atoms/affine/convolve.c‎

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -91,12 +91,12 @@ static void jacobian_init_impl(expr *node)
9191
new_sparse_matrix(csr_csc_matmul_alloc(cnode->T, cnode->Jchild_CSC));
9292
}
9393

94-
static void eval_jacobian(expr *node)
94+
static void eval_jacobian_impl(expr *node)
9595
{
9696
expr *child = node->left;
9797
convolve_expr *cnode = (convolve_expr *) node;
9898

99-
child->eval_jacobian(child);
99+
eval_jacobian(child);
100100

101101
/* J = T @ J_child */
102102
csr_to_csc_fill_values(child->jacobian->to_csr(child->jacobian),
@@ -115,7 +115,7 @@ static void wsum_hess_init_impl(expr *node)
115115
node->work->dwork = (double *) sp_malloc(cnode->n * sizeof(double));
116116
}
117117

118-
static void eval_wsum_hess(expr *node, const double *w)
118+
static void eval_wsum_hess_impl(expr *node, const double *w)
119119
{
120120
expr *child = node->left;
121121
convolve_expr *cnode = (convolve_expr *) node;
@@ -133,7 +133,7 @@ static void eval_wsum_hess(expr *node, const double *w)
133133
w_prime[j] = sum;
134134
}
135135

136-
child->eval_wsum_hess(child, w_prime);
136+
eval_wsum_hess(child, w_prime);
137137
memcpy(node->wsum_hess->x, child->wsum_hess->x,
138138
node->wsum_hess->nnz * sizeof(double));
139139
}
@@ -181,8 +181,8 @@ expr *new_convolve(expr *param_node, expr *child)
181181
convolve_expr *cnode = (convolve_expr *) sp_calloc(1, sizeof(convolve_expr));
182182
expr *node = &cnode->base;
183183
init_expr(node, d1, d2, child->n_vars, forward, jacobian_init_impl,
184-
eval_jacobian, is_affine, wsum_hess_init_impl, eval_wsum_hess,
185-
free_type_data);
184+
eval_jacobian_impl, is_affine, wsum_hess_init_impl,
185+
eval_wsum_hess_impl, free_type_data);
186186
node->left = child;
187187
expr_retain(child);
188188

‎src/atoms/affine/diag_vec.c‎

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -54,9 +54,9 @@ static void jacobian_init_impl(expr *node)
5454
node->jacobian = x->jacobian->diag_vec_alloc(x->jacobian);
5555
}
5656

57-
static void eval_jacobian(expr *node)
57+
static void eval_jacobian_impl(expr *node)
5858
{
59-
node->left->eval_jacobian(node->left);
59+
eval_jacobian(node->left);
6060

6161
/* fill the diagonal rows of the preallocated output. */
6262
node->left->jacobian->diag_vec_fill_values(node->left->jacobian, node->jacobian);
@@ -77,7 +77,7 @@ static void wsum_hess_init_impl(expr *node)
7777
node->wsum_hess = x->wsum_hess->copy_sparsity(x->wsum_hess);
7878
}
7979

80-
static void eval_wsum_hess(expr *node, const double *w)
80+
static void eval_wsum_hess_impl(expr *node, const double *w)
8181
{
8282
expr *x = node->left;
8383
int n = x->size;
@@ -89,7 +89,7 @@ static void eval_wsum_hess(expr *node, const double *w)
8989
}
9090

9191
/* Evaluate child's Hessian with extracted weights */
92-
x->eval_wsum_hess(x, node->work->dwork);
92+
eval_wsum_hess(x, node->work->dwork);
9393
memcpy(node->wsum_hess->x, x->wsum_hess->x,
9494
node->wsum_hess->nnz * sizeof(double));
9595
}
@@ -107,8 +107,9 @@ expr *new_diag_vec(expr *child)
107107
/* n is the number of elements (works for both row and column vectors) */
108108
int n = child->size;
109109
expr *node = (expr *) sp_calloc(1, sizeof(expr));
110-
init_expr(node, n, n, child->n_vars, forward, jacobian_init_impl, eval_jacobian,
111-
is_affine, wsum_hess_init_impl, eval_wsum_hess, NULL);
110+
init_expr(node, n, n, child->n_vars, forward, jacobian_init_impl,
111+
eval_jacobian_impl, is_affine, wsum_hess_init_impl,
112+
eval_wsum_hess_impl, NULL);
112113
node->left = child;
113114
expr_retain(child);
114115

‎src/atoms/affine/hstack.c‎

Lines changed: 17 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -87,15 +87,15 @@ static void jacobian_init_impl(expr *node)
8787
node->jacobian = new_sparse_matrix(A);
8888
}
8989

90-
static void eval_jacobian(expr *node)
90+
static void eval_jacobian_impl(expr *node)
9191
{
9292
hstack_expr *hnode = (hstack_expr *) node;
9393
int cursor = 0;
9494

9595
for (int i = 0; i < hnode->n_args; i++)
9696
{
9797
expr *child = hnode->args[i];
98-
child->eval_jacobian(child);
98+
eval_jacobian(child);
9999
/* to_csr needed for stacked_pd */
100100
CSR_matrix *child_csr = child->jacobian->to_csr(child->jacobian);
101101
memcpy(node->jacobian->x + cursor, child_csr->x,
@@ -148,7 +148,7 @@ static void wsum_hess_eval(expr *node, const double *w)
148148
for (int i = 0; i < hnode->n_args; i++)
149149
{
150150
expr *child = hnode->args[i];
151-
child->eval_wsum_hess(child, w + row_offset);
151+
eval_wsum_hess(child, w + row_offset);
152152
copy_CSR_matrix(H, hnode->CSR_work);
153153
sum_csr_fill_values(hnode->CSR_work,
154154
child->wsum_hess->to_csr(child->wsum_hess), H);
@@ -170,6 +170,17 @@ static bool is_affine(const expr *node)
170170
return true;
171171
}
172172

173+
/* Children live in args[], not left/right, so the parameter-refresh walk
174+
needs this hook to reach them. */
175+
static void set_needs_refresh_children(expr *node)
176+
{
177+
hstack_expr *hnode = (hstack_expr *) node;
178+
for (int i = 0; i < hnode->n_args; i++)
179+
{
180+
expr_set_needs_refresh(hnode->args[i]);
181+
}
182+
}
183+
173184
static void free_type_data(expr *node)
174185
{
175186
hstack_expr *hnode = (hstack_expr *) node;
@@ -199,9 +210,11 @@ expr *new_hstack(expr **args, int n_args, int n_vars)
199210
hstack_expr *hnode = (hstack_expr *) sp_calloc(1, sizeof(hstack_expr));
200211
expr *node = &hnode->base;
201212
init_expr(node, args[0]->d1, d2, n_vars, forward, jacobian_init_impl,
202-
eval_jacobian, is_affine, wsum_hess_init_impl, wsum_hess_eval,
213+
eval_jacobian_impl, is_affine, wsum_hess_init_impl, wsum_hess_eval,
203214
free_type_data);
204215

216+
node->set_needs_refresh_children = set_needs_refresh_children;
217+
205218
/* Set type-specific fields (deep copy args array) */
206219
hnode->args = (expr **) sp_calloc(n_args, sizeof(expr *));
207220
hnode->n_args = n_args;

0 commit comments

Comments
 (0)