Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 25 additions & 8 deletions include/expr.h
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,7 @@ typedef void (*local_jacobian_fn)(struct expr *node, double *out);
typedef void (*local_wsum_hess_fn)(struct expr *node, double *out, const double *w);
typedef bool (*is_affine_fn)(const struct expr *node);
typedef void (*free_type_data_fn)(struct expr *node);
typedef void (*set_needs_refresh_children_fn)(struct expr *node);
typedef bool (*set_needs_refresh_children_fn)(struct expr *node);

/* Workspace for derivative computation */
typedef struct
Expand Down Expand Up @@ -97,13 +97,27 @@ typedef struct expr
local_jacobian_fn local_jacobian; /* used by elementwise univariate atoms*/
local_wsum_hess_fn local_wsum_hess; /* used by elementwise univariate atoms*/
free_type_data_fn free_type_data; /* Cleanup for type-specific fields */
/* Recursion hook for expr_set_needs_refresh: atoms holding children
outside left/right (hstack's args[]) set this so the parameter-refresh
walk reaches them. NULL for binary/unary atoms. */
/* Recursion hook for expr_set_needs_refresh, which on its own only
descends left/right. Nodes that hold parameter-bearing subtrees
elsewhere install this so the walk reaches them: the parameter leaf
(reports itself), hstack (args[]) and the coefficient atoms
scalar_mult, vector_mult, kron, convolve, left_matmul and quad_form
(param_source). NULL everywhere else.

Contract: the hook MUST return whether the nodes it reaches contain
an updatable parameter (param_id >= 0). The walk ORs the result into
has_params; a hook that returns false lets the subtree be memoized
parameter-free and pruned from the second update on. */
set_needs_refresh_children_fn set_needs_refresh_children;
Expr_Work *work; /* derivative workspace */
/* Set to true on all nodes by problem_update_params() via
expr_set_needs_refresh(). Atoms that cache parameter data
/* Does this subtree contain an updatable parameter? Starts true so
nothing is pruned before the first walk; expr_set_needs_refresh
refines and memoizes it from the left/right walk and the
set_needs_refresh_children hook. */
bool has_params;
/* Set to true by problem_update_params() via expr_set_needs_refresh()
on every node whose subtree contains an updatable parameter;
parameter-free subtrees are skipped. Atoms that cache parameter data
(e.g. left_matmul_dense) check this flag before their forward
pass: if true, they refresh their cached matrices from
param_source->value and clear the flag to false. */
Expand Down Expand Up @@ -139,8 +153,11 @@ void expr_refresh_jacobian_csc(expr *node);
* Must be called after jacobian_init. */
void jacobian_csc_init(expr *node);

/* Recursively set needs_parameter_refresh on node and all children */
void expr_set_needs_refresh(expr *node);
/* Mark the subtree dirty and report whether it contains an updatable
* parameter. A subtree known to be parameter-free is skipped entirely: its
* values and derivatives cannot have changed, so re-arming it would only
* force a recompute of a result we already hold. */
bool expr_set_needs_refresh(expr *node);

/* Reference counting helpers */
void expr_retain(expr *node);
Expand Down
11 changes: 8 additions & 3 deletions src/atoms/affine/convolve.c
Original file line number Diff line number Diff line change
Expand Up @@ -39,9 +39,6 @@ static void forward(expr *node, const double *u)

if (cnode->base.needs_parameter_refresh)
{
/* Composite sources hold gated nodes of their own (promote, nested
mults): mark the whole side subtree before re-evaluating it. */
expr_set_needs_refresh(cnode->param_source);
cnode->param_source->forward(cnode->param_source, NULL);
/* refresh the convolution matrix values if it exists (necessary to check
for null in case someone calls forward before initializing the jacobian,
Expand Down Expand Up @@ -143,6 +140,13 @@ static bool is_affine(const expr *node)
return node->left->is_affine(node->left);
}

/* param_source lives outside left/right, so the refresh walk reaches it
here -- and reports whether it actually holds an updatable parameter. */
static bool set_needs_refresh_param_source(expr *node)
{
return expr_set_needs_refresh(((convolve_expr *) node)->param_source);
}

static void free_type_data(expr *node)
{
convolve_expr *cnode = (convolve_expr *) node;
Expand Down Expand Up @@ -197,6 +201,7 @@ expr *new_convolve(expr *param_node, expr *child)
/* Ensure first forward() pulls current param values through any
broadcast/promote wrappers and reflects them in T (once T is built). */
cnode->base.needs_parameter_refresh = true;
cnode->base.set_needs_refresh_children = set_needs_refresh_param_source;

return node;
}
6 changes: 4 additions & 2 deletions src/atoms/affine/hstack.c
Original file line number Diff line number Diff line change
Expand Up @@ -172,13 +172,15 @@ static bool is_affine(const expr *node)

/* Children live in args[], not left/right, so the parameter-refresh walk
needs this hook to reach them. */
static void set_needs_refresh_children(expr *node)
static bool set_needs_refresh_children(expr *node)
{
hstack_expr *hnode = (hstack_expr *) node;
bool child_has_params = false;
for (int i = 0; i < hnode->n_args; i++)
{
expr_set_needs_refresh(hnode->args[i]);
child_has_params |= expr_set_needs_refresh(hnode->args[i]);
}
return child_has_params;
}

static void free_type_data(expr *node)
Expand Down
11 changes: 8 additions & 3 deletions src/atoms/affine/kron.c
Original file line number Diff line number Diff line change
Expand Up @@ -44,9 +44,6 @@ static void refresh_param_values(kron_expr *knode)
return;
}

/* Composite sources hold gated nodes of their own (promote, nested
mults): mark the whole side subtree before re-evaluating it. */
expr_set_needs_refresh(knode->param_source);
knode->param_source->forward(knode->param_source, NULL);
knode->base.needs_parameter_refresh = false;
}
Expand Down Expand Up @@ -176,6 +173,13 @@ static bool is_affine(const expr *node)
return node->left->is_affine(node->left);
}

/* param_source lives outside left/right, so the refresh walk reaches it
here -- and reports whether it actually holds an updatable parameter. */
static bool set_needs_refresh_param_source(expr *node)
{
return expr_set_needs_refresh(((kron_expr *) node)->param_source);
}

static void free_type_data(expr *node)
{
kron_expr *knode = (kron_expr *) node;
Expand Down Expand Up @@ -214,6 +218,7 @@ static kron_expr *new_kron_common(expr *param_node, expr *child, int p, int q, i
}

knode->base.needs_parameter_refresh = true;
knode->base.set_needs_refresh_children = set_needs_refresh_param_source;
return knode;
}

Expand Down
11 changes: 8 additions & 3 deletions src/atoms/affine/left_matmul.c
Original file line number Diff line number Diff line change
Expand Up @@ -70,9 +70,6 @@ static void forward(expr *node, const double *u)
/* call forward on param_source if it exists and needs refresh */
if (lnode->param_source != NULL && lnode->base.needs_parameter_refresh)
{
/* Composite sources hold gated nodes of their own (promote, nested
mults): mark the whole side subtree before re-evaluating it. */
expr_set_needs_refresh(lnode->param_source);
lnode->param_source->forward(lnode->param_source, NULL);
}

Expand All @@ -94,6 +91,13 @@ static bool is_affine(const expr *node)
return node->left->is_affine(node->left);
}

/* param_source lives outside left/right, so the refresh walk reaches it
here -- and reports whether it actually holds an updatable parameter. */
static bool set_needs_refresh_param_source(expr *node)
{
return expr_set_needs_refresh(((left_matmul_expr *) node)->param_source);
}

static void free_type_data(expr *node)
{
left_matmul_expr *lnode = (left_matmul_expr *) node;
Expand Down Expand Up @@ -336,6 +340,7 @@ expr *new_left_matmul_dense(expr *param_node, expr *u, int m, int n,
lnode->A = new_permuted_dense_full(m, n, NULL);
lnode->AT = new_permuted_dense_full(n, m, NULL);
node->needs_parameter_refresh = true;
node->set_needs_refresh_children = set_needs_refresh_param_source;
}
/* constant matrix case */
else
Expand Down
7 changes: 7 additions & 0 deletions src/atoms/affine/parameter.c
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,12 @@ static void eval_wsum_hess_impl(expr *node, const double *w)
(void) w;
}

/* A leaf has no children, so it reports its own parameter-ness here. */
static bool set_needs_refresh_children(expr *node)
{
return ((parameter_expr *) node)->param_id >= 0;
}

static bool is_affine(const expr *node)
{
(void) node;
Expand All @@ -68,6 +74,7 @@ expr *new_parameter(int d1, int d2, int param_id, int n_vars, const double *valu

// TODO we should assert that the values array has the correct size.
pnode->param_id = param_id;
node->set_needs_refresh_children = set_needs_refresh_children;

if (values == NULL)
{
Expand Down
11 changes: 8 additions & 3 deletions src/atoms/affine/scalar_mult.c
Original file line number Diff line number Diff line change
Expand Up @@ -35,9 +35,6 @@ static void forward(expr *node, const double *u)
its values) */
if (snode->base.needs_parameter_refresh)
{
/* Composite sources hold gated nodes of their own (promote, nested
mults): mark the whole side subtree before re-evaluating it. */
expr_set_needs_refresh(snode->param_source);
snode->param_source->forward(snode->param_source, NULL);
snode->base.needs_parameter_refresh = false;
}
Expand Down Expand Up @@ -108,6 +105,13 @@ static bool is_affine(const expr *node)
return node->left->is_affine(node->left);
}

/* param_source lives outside left/right, so the refresh walk reaches it
here -- and reports whether it actually holds an updatable parameter. */
static bool set_needs_refresh_param_source(expr *node)
{
return expr_set_needs_refresh(((scalar_mult_expr *) node)->param_source);
}

static void free_type_data(expr *node)
{
scalar_mult_expr *snode = (scalar_mult_expr *) node;
Expand Down Expand Up @@ -135,6 +139,7 @@ expr *new_scalar_mult(expr *param_node, expr *child)

/* special case for handling broadcasting of constants correctly */
mult_node->base.needs_parameter_refresh = true;
mult_node->base.set_needs_refresh_children = set_needs_refresh_param_source;

return node;
}
11 changes: 8 additions & 3 deletions src/atoms/affine/vector_mult.c
Original file line number Diff line number Diff line change
Expand Up @@ -35,9 +35,6 @@ static void forward(expr *node, const double *u)
its values) */
if (vnode->base.needs_parameter_refresh)
{
/* Composite sources hold gated nodes of their own (promote, nested
mults): mark the whole side subtree before re-evaluating it. */
expr_set_needs_refresh(vnode->param_source);
vnode->param_source->forward(vnode->param_source, NULL);
vnode->base.needs_parameter_refresh = false;
}
Expand Down Expand Up @@ -109,6 +106,13 @@ static void eval_wsum_hess_impl(expr *node, const double *w)
node->wsum_hess->nnz * sizeof(double));
}

/* param_source lives outside left/right, so the refresh walk reaches it
here -- and reports whether it actually holds an updatable parameter. */
static bool set_needs_refresh_param_source(expr *node)
{
return expr_set_needs_refresh(((vector_mult_expr *) node)->param_source);
}

static void free_type_data(expr *node)
{
vector_mult_expr *vnode = (vector_mult_expr *) node;
Expand Down Expand Up @@ -141,6 +145,7 @@ expr *new_vector_mult(expr *param_node, expr *child)

/* special case for handling broadcasting of constants correctly */
vnode->base.needs_parameter_refresh = true;
vnode->base.set_needs_refresh_children = set_needs_refresh_param_source;

return node;
}
11 changes: 8 additions & 3 deletions src/atoms/other/quad_form.c
Original file line number Diff line number Diff line change
Expand Up @@ -56,9 +56,6 @@ static void forward(expr *node, const double *u)
/* refresh Q from the parameter if needed (no-op on the constant/sparse path) */
if (qnode->param_source != NULL && node->needs_parameter_refresh)
{
/* Composite sources hold gated nodes of their own (promote, nested
mults): mark the whole side subtree before re-evaluating it. */
expr_set_needs_refresh(qnode->param_source);
qnode->param_source->forward(qnode->param_source, NULL);
}
refresh_param_values_qf(qnode);
Expand Down Expand Up @@ -340,6 +337,13 @@ static void eval_wsum_hess_dense(expr *node, const double *w)
}
}

/* param_source lives outside left/right, so the refresh walk reaches it
here -- and reports whether it actually holds an updatable parameter. */
static bool set_needs_refresh_param_source(expr *node)
{
return expr_set_needs_refresh(((quad_form_expr *) node)->param_source);
}

static void free_type_data(expr *node)
{
quad_form_expr *qnode = (quad_form_expr *) node;
Expand Down Expand Up @@ -421,6 +425,7 @@ expr *new_quad_form_dense(expr *child, int n, const double *P_data,
/* Q is filled from the parameter on the first forward pass. */
qnode->Q = new_permuted_dense_full(n, n, NULL);
node->needs_parameter_refresh = true;
node->set_needs_refresh_children = set_needs_refresh_param_source;
}
else
{
Expand Down
23 changes: 18 additions & 5 deletions src/expr.c
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ void init_expr(expr *node, int d1, int d2, int n_vars, forward_fn forward,
node->eval_wsum_hess_impl = eval_wsum_hess;
node->free_type_data = free_type_data;
node->work = (Expr_Work *) sp_calloc(1, sizeof(Expr_Work));
node->has_params = true; /* assume dirty until the first walk refines it */
}

void jacobian_csc_init(expr *node)
Expand Down Expand Up @@ -151,20 +152,32 @@ void eval_wsum_hess(expr *node, const double *w)
matrix_values_changed(node->wsum_hess);
}

void expr_set_needs_refresh(expr *node)
bool expr_set_needs_refresh(expr *node)
{
if (node == NULL) return;
if (node == NULL) return false;

/* Known parameter-free: nothing below can have changed, so leave the
node's jacobian_evaluated latch set and skip the whole subtree. */
if (!node->has_params) return false;

node->needs_parameter_refresh = true;

/* Re-arm the eval_jacobian wrapper's values_version bump: the next eval
after a parameter update may change even an affine node's values. */
node->work->jacobian_evaluated = false;
expr_set_needs_refresh(node->left);
expr_set_needs_refresh(node->right);

bool child_has_params = expr_set_needs_refresh(node->left);
child_has_params |= expr_set_needs_refresh(node->right);
if (node->set_needs_refresh_children != NULL)
{
node->set_needs_refresh_children(node);
child_has_params |= node->set_needs_refresh_children(node);
}

/* The hook reports nodes the left/right walk cannot see: a parameter
leaf reports itself, hstack its args[], the coefficient atoms their
param_source. So this assignment is the whole answer. */
node->has_params = child_has_params;
return node->has_params;
}

void expr_retain(expr *node)
Expand Down
6 changes: 6 additions & 0 deletions tests/all_tests.c
Original file line number Diff line number Diff line change
Expand Up @@ -116,6 +116,7 @@
#ifdef PROFILE_ONLY
#include "profiling/profile_BTA_pd_csr_vs_csc.h"
#include "profiling/profile_hessian_exp_AX.h"
#include "profiling/profile_lasso.h"
#include "profiling/profile_left_matmul.h"
#include "profiling/profile_log_reg.h"
#include "profiling/profile_memory.h"
Expand Down Expand Up @@ -217,6 +218,10 @@ int main(void)
mu_run_test(test_values_version_csc_mirror_dedup, tests_run);
mu_run_test(test_values_version_stacked_pd_to_csr, tests_run);
mu_run_test(test_values_version_param_under_hstack, tests_run);
mu_run_test(test_refresh_prunes_param_free, tests_run);
mu_run_test(test_refresh_rearms_param_dependent, tests_run);
mu_run_test(test_refresh_prunes_fixed_constant, tests_run);
mu_run_test(test_refresh_rearms_updatable_constant, tests_run);
mu_run_test(test_values_version_spd_hess_terms, tests_run);
/* commented out - see test_quad_form.h */
// mu_run_test(test_quad_form2, tests_run);
Expand Down Expand Up @@ -612,6 +617,7 @@ int main(void)

#ifdef PROFILE_ONLY
printf("\n--- Profiling Tests ---\n");
mu_run_test(profile_lasso, tests_run);
mu_run_test(profile_left_matmul, tests_run);
mu_run_test(profile_log_reg, tests_run);
mu_run_test(profile_trimmed_log_reg, tests_run);
Expand Down
Loading
Loading