|
23 | 23 |
|
24 | 24 | /* Constant vector elementwise multiplication: y = a \circ child */ |
25 | 25 |
|
| 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 | + |
26 | 31 | static void forward(expr *node, const double *u) |
27 | 32 | { |
28 | 33 | 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); |
30 | 35 |
|
31 | 36 | /* child's forward pass */ |
32 | 37 | child->forward(child, u); |
@@ -54,7 +59,7 @@ static void jacobian_init(expr *node) |
54 | 59 | static void eval_jacobian(expr *node) |
55 | 60 | { |
56 | 61 | 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); |
58 | 63 |
|
59 | 64 | /* evaluate x */ |
60 | 65 | x->eval_jacobian(x); |
@@ -87,7 +92,7 @@ static void wsum_hess_init(expr *node) |
87 | 92 | static void eval_wsum_hess(expr *node, const double *w) |
88 | 93 | { |
89 | 94 | 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); |
91 | 96 |
|
92 | 97 | /* scale weights w by a */ |
93 | 98 | for (int i = 0; i < node->size; i++) |
@@ -128,6 +133,39 @@ expr *new_const_vector_mult(const double *a, expr *child) |
128 | 133 | /* copy a vector */ |
129 | 134 | vnode->a = (double *) malloc(child->size * sizeof(double)); |
130 | 135 | 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); |
131 | 169 |
|
132 | 170 | return node; |
133 | 171 | } |
0 commit comments