Skip to main content

dace_rs/
eval.rs

1//! Polynomial evaluation: compiled evaluation trees, partial evaluation,
2//! and variable substitution.
3//!
4//! Ports `core/daceeval.c`: the Horner-style evaluation tree of
5//! `daceEvalTree` (as the C++ `compiledDA`), partial evaluation of one
6//! variable (`daceEvalVariable`), variable replacement/scaling/translation,
7//! and the monomial-wise dot product `daceEvalMonomials`.
8
9use std::sync::Arc;
10
11use crate::context::{Context, npown_i64};
12use crate::da::{Da, RawTerm};
13use crate::kernels::{keep, multiply, weighted_sum};
14
15/// Pack a dense coefficient array (length `nmmax`) into a sparse `Da`,
16/// flushing `|c| <= eps` and re-zeroing the array (`dacePack`, fast path:
17/// the truncation order is enforced by the callers, as in C).
18pub(crate) fn pack(ctx: &Arc<Context>, cc: &mut [f64]) -> Da {
19    let (eps, _nocut) = crate::context::eps_nocut();
20    let mut terms = Vec::new();
21    for (i, c) in cc.iter_mut().enumerate() {
22        if keep(*c, eps) {
23            terms.push(RawTerm {
24                idx: i as u32,
25                c: *c,
26            });
27        }
28        *c = 0.0;
29    }
30    Da {
31        ctx: ctx.clone(),
32        terms,
33    }
34}
35
36impl Da {
37    /// Partial evaluation: replace independent variable `var` (1-based) by
38    /// the value `val` (`daceEvalVariable`). Out-of-range variables warn
39    /// and return zero.
40    pub fn plug(&self, var: u32, val: f64) -> Da {
41        let ctx = &self.ctx;
42        if !(1..=ctx.nvmax).contains(&var) {
43            log::warn!("DACE error 624: invalid independent variable {var} in plug");
44            return Da::new();
45        }
46        let (_eps, nocut) = crate::context::eps_nocut();
47        let ibase = ctx.nomax + 1;
48        let j = if var > ctx.nv1 {
49            var - 1 - ctx.nv1
50        } else {
51            var - 1
52        };
53        let idiv = npown_i64(ibase, j);
54        let in_second_half = var > ctx.nv1;
55
56        let mut p = vec![1.0; ctx.nomax as usize + 1];
57        for i in 1..p.len() {
58            p[i] = p[i - 1] * val;
59        }
60        let mut cc = vec![0.0; ctx.nmmax as usize];
61
62        for t in &self.terms {
63            let ic1 = ctx.ie1[t.idx as usize];
64            let ic2 = ctx.ie2[t.idx as usize];
65            let ipow = if in_second_half {
66                (ic2 / idiv) % ibase
67            } else {
68                (ic1 / idiv) % ibase
69            };
70            let j = if in_second_half {
71                ctx.ia1[ic1 as usize] + ctx.ia2[(ic2 - ipow * idiv) as usize]
72            } else {
73                ctx.ia1[(ic1 - ipow * idiv) as usize] + ctx.ia2[ic2 as usize]
74            };
75            if ctx.order_of(j) <= nocut {
76                cc[j as usize] += t.c * p[ipow as usize];
77            }
78        }
79
80        pack(ctx, &mut cc)
81    }
82
83    /// Replace independent variable `from` by `val` times independent
84    /// variable `to` (`daceReplaceVariable`); `from == to` scales the
85    /// variable by `val`. Out-of-range variables warn and return zero.
86    ///
87    /// Divergence from C: the C implementation indexes its 0-based exponent
88    /// array with the 1-based variable numbers, so it actually replaces
89    /// variable `from + 1` by `val ยท (variable to + 1)` and silently does
90    /// nothing when `from == nvmax`. This implementation follows the
91    /// documented (1-based) semantics instead.
92    pub fn replace_variable(&self, from: u32, to: u32, val: f64) -> Da {
93        let ctx = &self.ctx;
94        if !(1..=ctx.nvmax).contains(&from) || !(1..=ctx.nvmax).contains(&to) {
95            log::warn!("DACE error 624: invalid independent variable in replace_variable");
96            return Da::new();
97        }
98        if from == to {
99            return self.scale_variable(from, val);
100        }
101
102        let mut pows = vec![1.0; ctx.nomax as usize + 1];
103        for i in 0..ctx.nomax as usize {
104            pows[i + 1] = pows[i] * val;
105        }
106        let mut p = vec![0u32; ctx.nvmax as usize];
107        let mut cc = vec![0.0; ctx.nmmax as usize];
108        for t in &self.terms {
109            ctx.decode_into(t.idx, &mut p);
110            p[to as usize - 1] += p[from as usize - 1];
111            let c = pows[p[from as usize - 1] as usize] * t.c;
112            p[from as usize - 1] = 0;
113            let idx = ctx.encode(&p).expect("order preserved by replacement");
114            cc[idx as usize] += c;
115        }
116        pack(ctx, &mut cc)
117    }
118
119    /// Scale independent variable `var` by `val`: `x_var -> val * x_var`
120    /// (`daceScaleVariable`). Out-of-range variables warn and return zero.
121    pub fn scale_variable(&self, var: u32, val: f64) -> Da {
122        let ctx = &self.ctx;
123        if !(1..=ctx.nvmax).contains(&var) {
124            log::warn!("DACE error 624: invalid independent variable {var} in scale_variable");
125            return Da::new();
126        }
127        let mut pows = vec![1.0; ctx.nomax as usize + 1];
128        for i in 0..ctx.nomax as usize {
129            pows[i + 1] = pows[i] * val;
130        }
131        let ibase = ctx.nomax + 1;
132        let j = if var > ctx.nv1 {
133            var - 1 - ctx.nv1
134        } else {
135            var - 1
136        };
137        let idiv = npown_i64(ibase, j);
138        let in_second_half = var > ctx.nv1;
139
140        let mut terms = self.terms.clone();
141        for t in terms.iter_mut() {
142            let ipow = if in_second_half {
143                (ctx.ie2[t.idx as usize] / idiv) % ibase
144            } else {
145                (ctx.ie1[t.idx as usize] / idiv) % ibase
146            };
147            t.c *= pows[ipow as usize];
148        }
149        Da {
150            ctx: ctx.clone(),
151            terms,
152        }
153    }
154
155    /// Translate independent variable `var` to `a*x + c`
156    /// (`daceTranslateVariable`). Out-of-range variables warn and return
157    /// zero.
158    pub fn translate_variable(&self, var: u32, a: f64, c: f64) -> Da {
159        let ctx = &self.ctx;
160        if !(1..=ctx.nvmax).contains(&var) {
161            log::warn!("DACE error 624: invalid independent variable {var} in translate_variable");
162            return Da::new();
163        }
164        let n1 = ctx.nomax as usize;
165
166        let mut powa = vec![1.0; n1 + 1];
167        let mut powc = vec![1.0; n1 + 1];
168        for i in 0..n1 {
169            powa[i + 1] = powa[i] * a;
170            powc[i + 1] = powc[i] * c;
171        }
172
173        // binomial coefficients n choose k
174        let mut binomial = vec![0.0; (n1 + 1) * (n1 + 1)];
175        for n in 0..=n1 {
176            binomial[n * (n1 + 1)] = 1.0;
177            binomial[n * (n1 + 1) + n] = 1.0;
178            for k in 1..n {
179                binomial[n * (n1 + 1) + k] =
180                    binomial[(n - 1) * (n1 + 1) + k - 1] + binomial[(n - 1) * (n1 + 1) + k];
181            }
182        }
183
184        let mut p = vec![0u32; ctx.nvmax as usize];
185        let mut cc = vec![0.0; ctx.nmmax as usize];
186        for t in &self.terms {
187            ctx.decode_into(t.idx, &mut p);
188            let n = p[(var - 1) as usize];
189
190            // shortcut the case when the monomial doesn't depend on var
191            if n == 0 {
192                cc[t.idx as usize] += t.c;
193                continue;
194            }
195
196            for k in 0..=n {
197                let idx = ctx.encode(&p).expect("order preserved by translation");
198                cc[idx as usize] += t.c
199                    * binomial[n as usize * (n1 + 1) + k as usize]
200                    * powa[(n - k) as usize]
201                    * powc[k as usize];
202                // C decrements unconditionally (unsigned wrap), but the
203                // value is only read after the next decode; guard instead.
204                if p[(var - 1) as usize] > 0 {
205                    p[(var - 1) as usize] -= 1;
206                }
207            }
208        }
209
210        pack(ctx, &mut cc)
211    }
212
213    /// Evaluate by providing the value of each monomial in `values`: the
214    /// monomial-wise dot product of the two DAs (`daceEvalMonomials`).
215    pub fn eval_monomials(&self, values: &Da) -> f64 {
216        Da::assert_same_context(self, values);
217        let mut res = 0.0;
218        let mut ib = values.terms.iter().peekable();
219        for ta in &self.terms {
220            while let Some(tb) = ib.peek() {
221                if tb.idx < ta.idx {
222                    ib.next();
223                } else {
224                    break;
225                }
226            }
227            match ib.peek() {
228                Some(tb) if tb.idx == ta.idx => res += tb.c * ta.c,
229                Some(_) => {}
230                None => break,
231            }
232        }
233        res
234    }
235
236    /// Evaluate at a point: compile and evaluate the tree, returning the
237    /// first component (C++ `DA::eval`).
238    pub fn eval(&self, args: &[f64]) -> f64 {
239        self.compile().eval(args)[0]
240    }
241
242    /// Evaluate with DA arguments (contraction): substitute each argument
243    /// DA for the corresponding variable.
244    pub fn eval_da(&self, args: &[Da]) -> Da {
245        self.compile().eval_da(args)[0].clone()
246    }
247
248    /// Compile into a reusable evaluation tree (C++ `compiledDA`).
249    pub fn compile(&self) -> CompiledDa {
250        CompiledDa::from_das(std::slice::from_ref(self))
251    }
252}
253
254/// A compiled evaluation tree over one or more DAs (C++ `compiledDA`):
255/// precomputed Horner-tree coefficients that can be evaluated repeatedly
256/// at low cost.
257#[derive(Debug, Clone)]
258pub struct CompiledDa {
259    /// Number of component DAs.
260    pub dim: u32,
261    /// Maximum order of the tree.
262    pub ord: u32,
263    /// Number of variables used.
264    pub vars: u32,
265    /// Number of terms (including the root).
266    pub terms: u32,
267    /// Packed coefficients: two unused slots, `dim` constants, then per
268    /// term `(level, variable, dim coefficients)` with 1-based indices.
269    pub(crate) ac: Vec<f64>,
270}
271
272const _: () = {
273    const fn assert_send_sync<T: Send + Sync>() {}
274    assert_send_sync::<CompiledDa>();
275};
276
277impl CompiledDa {
278    /// Compile several DAs (same context) into one shared tree
279    /// (`daceEvalTree`).
280    ///
281    /// # Panics
282    ///
283    /// Panics with [`crate::DaceError`] when the DAs belong to different contexts.
284    pub fn from_das(das: &[Da]) -> CompiledDa {
285        for da in das {
286            Da::assert_same_context(&das[0], da);
287        }
288        let ctx = das[0].ctx.clone();
289        let count = das.len();
290        let mut nc = vec![0u32; ctx.nmmax as usize];
291
292        // mark all used monomials as new
293        for da in das {
294            for t in &da.terms {
295                nc[t.idx as usize] = 2;
296            }
297        }
298
299        // make sure each term has a parent
300        nc[0] = 1; // constant part is the root, doesn't need a parent
301        let mut p = vec![0u32; ctx.nvmax as usize];
302        for i in 1..ctx.nmmax as usize {
303            if nc[i] != 2 {
304                continue;
305            }
306            nc[i] = 1;
307            ctx.decode_into(i as u32, &mut p);
308            // generate an ancestor tree for this entry
309            let mut parent: i64;
310            loop {
311                parent = -1;
312                // find a parent
313                for j in 0..ctx.nvmax as usize {
314                    if p[j] == 0 {
315                        continue;
316                    }
317                    p[j] -= 1;
318                    if nc[ctx.encode(&p).expect("valid parent") as usize] != 0 {
319                        // parent already exists => done
320                        parent = -1;
321                        break;
322                    }
323                    p[j] += 1;
324                    parent = j as i64;
325                }
326                // no parent found => create foster parent
327                if parent >= 0 {
328                    p[parent as usize] -= 1;
329                    nc[ctx.encode(&p).expect("valid foster parent") as usize] = 1;
330                } else {
331                    break;
332                }
333            }
334        }
335
336        // constant terms are always stored
337        nc[0] = 3;
338        let mut nord = 0u32;
339        let mut nvar = 0u32;
340        let mut nterm = 1u32;
341        let mut ac: Vec<f64> = Vec::new();
342        ac.push(0.0);
343        ac.push(0.0);
344        for da in das {
345            ac.push(da.cons());
346        }
347
348        // higher order terms
349        p[0] = 1;
350        for slot in p.iter_mut().skip(1) {
351            *slot = 0;
352        }
353        let mut stack = vec![0u32; ctx.nomax as usize];
354        let mut sp: i64 = 0;
355        stack[0] = 0;
356        while sp >= 0 {
357            let ic = ctx.encode(&p).expect("valid monomial");
358            if nc[ic as usize] == 1 {
359                // store entry
360                nc[ic as usize] = 3;
361                nord = nord.max(sp as u32 + 1);
362                nvar = nvar.max(stack[sp as usize] + 1);
363                nterm += 1;
364                ac.push((sp + 1) as f64); // +1 for 1-based indices (as in C)
365                ac.push(f64::from(stack[sp as usize] + 1));
366                for da in das {
367                    ac.push(da.get_coefficient0(ic));
368                }
369
370                // step forward if we can
371                if sp < ctx.nomax as i64 - 1 {
372                    sp += 1;
373                    stack[sp as usize] = 0;
374                    p[0] += 1;
375                    continue;
376                }
377            }
378            if stack[sp as usize] < ctx.nvmax - 1 {
379                // step sideways
380                let s = stack[sp as usize];
381                p[s as usize] -= 1;
382                stack[sp as usize] = s + 1;
383                p[(s + 1) as usize] += 1;
384            } else {
385                // step back
386                let s = stack[sp as usize];
387                p[s as usize] -= 1;
388                sp -= 1;
389            }
390        }
391
392        CompiledDa {
393            dim: count as u32,
394            ord: nord,
395            vars: nvar,
396            terms: nterm,
397            ac,
398        }
399    }
400
401    /// Evaluate the tree at a point (`compiledDA::eval<double>`): one
402    /// result per component DA.
403    pub fn eval(&self, args: &[f64]) -> Vec<f64> {
404        let narg = args.len();
405        let mut p = self.ac.iter().skip(2);
406        let mut xm = vec![0.0; self.ord as usize + 1];
407
408        // prepare temporary powers
409        xm[0] = 1.0;
410        // constant part
411        let mut res: Vec<f64> = p.by_ref().take(self.dim as usize).copied().collect();
412        // higher order terms
413        for _ in 1..self.terms {
414            let jl = *p.next().expect("tree level") as usize;
415            let jv = *p.next().expect("tree variable") as usize - 1;
416            xm[jl] = if jv < narg {
417                xm[jl - 1] * args[jv]
418            } else {
419                0.0
420            };
421            for r in res.iter_mut().take(self.dim as usize) {
422                *r += xm[jl] * p.next().expect("tree coefficient");
423            }
424        }
425        res
426    }
427
428    /// Evaluate the tree with DA arguments (contraction,
429    /// `compiledDA::eval<DA>`), including the skip logic for unused
430    /// variables.
431    pub fn eval_da(&self, args: &[Da]) -> Vec<Da> {
432        let narg = args.len();
433        let mut jlskip = self.ord + 1;
434        let mut p = self.ac.iter().skip(2);
435        let mut xm: Vec<Da> = (0..=self.ord).map(|_| Da::new()).collect();
436
437        // prepare temporary powers
438        xm[0] = Da::constant(1.0);
439        // constant part
440        let mut res: Vec<Da> = p
441            .by_ref()
442            .take(self.dim as usize)
443            .map(|c| Da::constant(*c))
444            .collect();
445        // higher order terms
446        for _ in 1..self.terms {
447            let jl = *p.next().expect("tree level") as u32;
448            let jv = *p.next().expect("tree variable") as u32 - 1;
449            if jl > jlskip {
450                p.by_ref().take(self.dim as usize).for_each(drop);
451                continue;
452            }
453            if jv as usize >= narg {
454                jlskip = jl;
455                p.by_ref().take(self.dim as usize).for_each(drop);
456                continue;
457            }
458            jlskip = self.ord + 1;
459            xm[jl as usize] = multiply(&xm[(jl - 1) as usize], &args[jv as usize]);
460            for r in res.iter_mut().take(self.dim as usize) {
461                let coef = *p.next().expect("tree coefficient");
462                if coef != 0.0 {
463                    *r = weighted_sum(r, 1.0, &xm[jl as usize], coef);
464                }
465            }
466        }
467        res
468    }
469}
470
471#[cfg(test)]
472mod tests {
473    use super::*;
474    use crate::test_support::CONTEXT_LOCK;
475
476    #[test]
477    fn eval_matches_direct_evaluation() {
478        let _g = CONTEXT_LOCK.lock();
479        crate::context::init(5, 3).unwrap();
480        let x = Da::variable(1);
481        let y = Da::variable(2);
482        let z = Da::variable(3);
483
484        let f = 1.0 + 2.0 * x.clone() * y.clone() - 0.5 * z.clone() * z.clone() + 0.25 * x.clone();
485
486        // direct polynomial evaluation at a point
487        let (px, py, pz) = (1.3, -0.7, 0.9);
488        let direct = 1.0 + 2.0 * px * py - 0.5 * pz * pz + 0.25 * px;
489        assert!((f.eval(&[px, py, pz]) - direct).abs() < 1e-12);
490
491        // CompiledDa eval agrees with Da::eval on random points
492        let compiled = f.compile();
493        for trial in 0..10 {
494            let t = 0.1 * trial as f64;
495            let args = [t, -0.3 + 0.05 * t, 0.7 - 0.1 * t];
496            assert!((compiled.eval(&args)[0] - f.eval(&args)).abs() < 1e-13);
497        }
498
499        // multi-component compilation
500        let g = y.clone() - z.clone();
501        let c2 = CompiledDa::from_das(&[f.clone(), g.clone()]);
502        let r = c2.eval(&[px, py, pz]);
503        assert_eq!(r.len(), 2);
504        assert!((r[1] - (py - pz)).abs() < 1e-13);
505
506        // eval_da contraction: f(x, y, x) with z := x
507        let sub = f.eval_da(&[x.clone(), y.clone(), x.clone()]);
508        // f(x,y,x) coefficient checks: 2xy term unchanged, -0.5x^2 replaces -0.5z^2
509        assert!((sub.get_coefficient(&[1, 1, 0]) - 2.0).abs() < 1e-13);
510        assert!((sub.get_coefficient(&[2, 0, 0]) + 0.5).abs() < 1e-13);
511
512        // plug then eval == eval with that coordinate fixed
513        let plugged = f.plug(2, py);
514        assert!((plugged.eval(&[px, 0.0, pz]) - direct).abs() < 1e-12);
515
516        // translate_variable on x by c gives cons() == c
517        let t = x.clone().translate_variable(1, 1.0, 0.75);
518        assert!((t.cons() - 0.75).abs() < 1e-15);
519        assert!((t.get_coefficient(&[1, 0, 0]) - 1.0).abs() < 1e-15);
520
521        // scale_variable multiplies the right exponents
522        let s = (x.clone() * y.clone()).scale_variable(1, 3.0);
523        assert!((s.get_coefficient(&[1, 1, 0]) - 3.0).abs() < 1e-15);
524
525        // replace_variable y -> 2x
526        let r = (x.clone() * y.clone()).replace_variable(2, 1, 2.0);
527        assert!((r.get_coefficient(&[2, 0, 0]) - 2.0).abs() < 1e-15);
528        assert_eq!(r.size(), 1);
529
530        // eval_monomials == monomial-wise dot product
531        let a = 1.0 + x.clone();
532        let b = 2.0 + 3.0 * x.clone();
533        assert!((a.eval_monomials(&b) - (1.0 * 2.0 + 1.0 * 3.0)).abs() < 1e-15);
534    }
535}