Skip to main content

dace_rs/
da.rs

1//! The DA type: a truncated multivariate Taylor polynomial.
2
3use std::cell::Cell;
4use std::ops::{Add, AddAssign, Div, DivAssign, Mul, MulAssign, Neg, Sub, SubAssign};
5use std::sync::Arc;
6
7use crate::context::Context;
8use crate::elementary::minv;
9use crate::error::{codes, dace_panic};
10use crate::kernels::{multiply, weighted_sum};
11use crate::monomial::Monomial;
12
13/// One stored term of a [`Da`]: packed monomial index plus coefficient.
14#[derive(Clone, Copy, Debug)]
15pub(crate) struct RawTerm {
16    pub idx: u32,
17    pub c: f64,
18}
19
20/// A truncated multivariate Taylor polynomial ("differential algebra" value).
21///
22/// Stored as a sparse list of `(monomial index, coefficient)` terms, sorted by
23/// ascending monomial index; the index packs the exponent vector into the
24/// canonical position determined by the active [context][crate::context].
25///
26/// Values created under one context keep working after
27/// [`init`][crate::init] is called again (they hold on to their original
28/// context); mixing values from different contexts in one operation panics.
29/// See the crate-level [Multithreading][crate#multithreading] section for
30/// cross-thread use.
31#[derive(Clone, Debug)]
32pub struct Da {
33    pub(crate) ctx: Arc<Context>,
34    /// Sorted ascending by `idx`.
35    pub(crate) terms: Vec<RawTerm>,
36}
37
38const _: () = {
39    const fn assert_send_sync<T: Send + Sync>() {}
40    assert_send_sync::<Da>();
41};
42
43impl Da {
44    /// The zero polynomial (also [`Da::default`]).
45    ///
46    /// # Panics
47    ///
48    /// Panics with [`crate::DaceError`] if DACE has not been initialized.
49    pub fn new() -> Da {
50        Da {
51            ctx: Context::current(),
52            terms: Vec::new(),
53        }
54    }
55
56    /// The constant polynomial `c` (`daceCreateConstant`); `|c| <= eps` gives
57    /// the zero polynomial.
58    pub fn constant(c: f64) -> Da {
59        Da::variable_scaled(0, c)
60    }
61
62    /// The independent DA variable number `var` (1-based), i.e. the identity
63    /// in that variable (`daceCreateVariable`).
64    ///
65    /// Divergence from C: an out-of-range `var` logs a warning and returns the
66    /// zero polynomial instead of raising C error 624.
67    pub fn variable(var: u32) -> Da {
68        Da::variable_scaled(var, 1.0)
69    }
70
71    /// Alias of [`Da::variable`] (`DA::identity` in the C++ interface).
72    pub fn identity(var: u32) -> Da {
73        Da::variable(var)
74    }
75
76    fn variable_scaled(var: u32, ckon: f64) -> Da {
77        let ctx = Context::current();
78        if var > ctx.nvmax {
79            log::warn!("DACE error 624: invalid independent variable {var}; returning zero DA");
80            return Da {
81                ctx,
82                terms: Vec::new(),
83            };
84        }
85        let (eps, _nocut) = crate::context::eps_nocut();
86        if ckon.abs() <= eps {
87            return Da {
88                ctx,
89                terms: Vec::new(),
90            };
91        }
92        // Set up the encoded exponents (dacebasic.c:80-99).
93        let base = ctx.nomax + 1;
94        let (ic1, ic2) = if var == 0 {
95            (0, 0)
96        } else if var > ctx.nv1 {
97            (0, crate::context::npown_i64(base, var - 1 - ctx.nv1))
98        } else {
99            (crate::context::npown_i64(base, var - 1), 0)
100        };
101        let idx = ctx.ia1[ic1 as usize] + ctx.ia2[ic2 as usize];
102        Da {
103            ctx,
104            terms: vec![RawTerm { idx, c: ckon }],
105        }
106    }
107
108    /// The polynomial consisting of the single monomial `c * jj[0]^jj[1] * ...`
109    /// (`daceCreateMonomial`).
110    ///
111    /// `jj` is padded with zeros or truncated to the number of DA variables
112    /// (with a warning). Terms whose total order exceeds the maximum
113    /// computation order, or with an exponent above it, are dropped with a
114    /// warning (divergence from C, which encodes them as the constant term
115    /// after raising error 622). `|c| <= eps` gives the zero polynomial.
116    pub fn monomial(jj: &[u32], c: f64) -> Da {
117        let ctx = Context::current();
118        let (eps, _nocut) = crate::context::eps_nocut();
119        if c.abs() <= eps {
120            return Da {
121                ctx,
122                terms: Vec::new(),
123            };
124        }
125        let jj = fix_exponent_length(&ctx, jj);
126        match ctx.encode(&jj) {
127            Some(idx) => Da {
128                ctx,
129                terms: vec![RawTerm { idx, c }],
130            },
131            None => {
132                log::warn!(
133                    "DACE error 622: monomial order too large in Da::monomial; term dropped"
134                );
135                Da {
136                    ctx,
137                    terms: Vec::new(),
138                }
139            }
140        }
141    }
142
143    /// A DA with randomly filled coefficients (`daceCreateRandom`).
144    ///
145    /// `cmu` is the filling factor: `|cmu|` is the fraction of non-zero
146    /// coefficients; `cmu < 0` draws coefficients in `[-1, 1]`, `cmu > 0`
147    /// weights them to decay exponentially with order from 1.0 towards the
148    /// machine epsilon.
149    ///
150    /// Divergence from C: the C library uses libc `rand()` (platform- and
151    /// seed-dependent); this implementation uses a deterministic 64-bit LCG
152    /// per thread, so results are reproducible everywhere.
153    pub fn random(cmu: f64) -> Da {
154        let ctx = Context::current();
155        let (_eps, nocut) = crate::context::eps_nocut();
156        let mut terms = Vec::new();
157        for i in 0..ctx.nmmax {
158            if ctx.ieo[i as usize] <= nocut && dace_random() < cmu.abs() {
159                let c = if cmu < 0.0 {
160                    2.0 * dace_random() - 1.0
161                } else {
162                    let w = ctx
163                        .epsmac
164                        .powf(f64::from(ctx.ieo[i as usize]) / f64::from(nocut));
165                    w * (2.0 * dace_random() - 1.0)
166                };
167                terms.push(RawTerm { idx: i, c });
168            }
169        }
170        Da { ctx, terms }
171    }
172
173    // -----------------------------------------------------------------------
174    // Inspection
175    // -----------------------------------------------------------------------
176
177    /// The constant part of the polynomial (`daceGetConstant`).
178    pub fn cons(&self) -> f64 {
179        match self.terms.first() {
180            Some(t) if t.idx == 0 => t.c,
181            _ => 0.0,
182        }
183    }
184
185    /// The linear coefficients, one per DA variable (`daceGetLinear`).
186    pub fn linear(&self) -> Vec<f64> {
187        let mut jj = vec![0u32; self.ctx.nvmax as usize];
188        let mut c = vec![0.0; self.ctx.nvmax as usize];
189        for (i, ci) in c.iter_mut().enumerate() {
190            jj[i] = 1;
191            *ci = self.get_coefficient(&jj);
192            jj[i] = 0;
193        }
194        c
195    }
196
197    /// The gradient: derivatives with respect to all DA variables
198    /// (`DA::gradient` in the C++ interface).
199    pub fn gradient(&self) -> Vec<Da> {
200        (1..=self.ctx.nvmax).map(|i| self.deriv(i)).collect()
201    }
202
203    /// The number of stored (non-zero) monomials (`daceGetLength`).
204    pub fn size(&self) -> usize {
205        self.terms.len()
206    }
207
208    /// The coefficient of the monomial with exponents `jj`
209    /// (`daceGetCoefficient`). `jj` is padded/truncated to the number of DA
210    /// variables; invalid exponents return 0.0 with a warning.
211    pub fn get_coefficient(&self, jj: &[u32]) -> f64 {
212        let jj = fix_exponent_length(&self.ctx, jj);
213        match self.ctx.encode(&jj) {
214            Some(ic) => self.get_coefficient0(ic),
215            None => {
216                log::warn!(
217                    "DACE error 622: monomial order too large in get_coefficient; returning 0.0"
218                );
219                0.0
220            }
221        }
222    }
223
224    /// The coefficient of the monomial with packed index `ic`.
225    pub(crate) fn get_coefficient0(&self, ic: u32) -> f64 {
226        match self.terms.binary_search_by_key(&ic, |t| t.idx) {
227            Ok(pos) => self.terms[pos].c,
228            Err(_) => 0.0,
229        }
230    }
231
232    /// Set the coefficient of the monomial with exponents `jj`
233    /// (`daceSetCoefficient`): sets it, replaces it, or removes the monomial
234    /// when `|c| <= eps`, keeping the term list sorted.
235    pub fn set_coefficient(&mut self, jj: &[u32], c: f64) {
236        let jj = fix_exponent_length(&self.ctx, jj);
237        match self.ctx.encode(&jj) {
238            Some(ic) => self.set_coefficient0(ic, c),
239            None => {
240                log::warn!("DACE error 622: monomial order too large in set_coefficient; ignored");
241            }
242        }
243    }
244
245    /// Set the coefficient of the monomial with packed index `ic`.
246    pub(crate) fn set_coefficient0(&mut self, ic: u32, c: f64) {
247        let (eps, _nocut) = crate::context::eps_nocut();
248        match self.terms.binary_search_by_key(&ic, |t| t.idx) {
249            Ok(pos) => {
250                if crate::kernels::keep(c, eps) {
251                    self.terms[pos].c = c;
252                } else {
253                    self.terms.remove(pos);
254                }
255            }
256            Err(pos) => {
257                if crate::kernels::keep(c, eps) {
258                    self.terms.insert(pos, RawTerm { idx: ic, c });
259                }
260            }
261        }
262    }
263
264    /// The monomial at 1-based position `pos` in the stored term list
265    /// (`DA::getMonomial` in the C++ interface); `None` when out of range.
266    /// The ordering is implementation-dependent.
267    pub fn get_monomial(&self, pos: usize) -> Option<Monomial> {
268        self.terms.get(pos.wrapping_sub(1)).map(|t| Monomial {
269            jj: self.ctx.decode(t.idx),
270            c: t.c,
271        })
272    }
273
274    /// Iterate over all stored monomials, in stored order.
275    pub fn iter_monomials(&self) -> impl Iterator<Item = Monomial> + '_ {
276        let ctx = self.ctx.clone();
277        self.terms.iter().map(move |t| Monomial {
278            jj: ctx.decode(t.idx),
279            c: t.c,
280        })
281    }
282
283    /// Whether any coefficient is NaN (`daceIsNan`).
284    pub fn is_nan(&self) -> bool {
285        self.terms.iter().any(|t| t.c.is_nan())
286    }
287
288    /// Whether any coefficient is infinite (`daceIsInf`).
289    pub fn is_inf(&self) -> bool {
290        self.terms.iter().any(|t| t.c.is_infinite())
291    }
292
293    // -----------------------------------------------------------------------
294    // Calculus (dacemath.c:504-649)
295    // -----------------------------------------------------------------------
296
297    /// Derivative with respect to independent variable `var` (1-based)
298    /// (`daceDifferentiate`). Out-of-range variables warn and return zero.
299    pub fn deriv(&self, var: u32) -> Da {
300        let ctx = &self.ctx;
301        if !(1..=ctx.nvmax).contains(&var) {
302            log::warn!(
303                "DACE error 624: invalid independent variable {var} in deriv; returning zero DA"
304            );
305            return Da::new();
306        }
307        let (_eps, nocut) = crate::context::eps_nocut();
308        let ibase = ctx.nomax + 1;
309        let j = if var > ctx.nv1 {
310            var - 1 - ctx.nv1
311        } else {
312            var - 1
313        };
314        let idiv = crate::context::npown_i64(ibase, j);
315        let in_second_half = var > ctx.nv1;
316        let mut terms = Vec::with_capacity(self.terms.len());
317        for t in &self.terms {
318            let ic1 = ctx.ie1[t.idx as usize];
319            let ic2 = ctx.ie2[t.idx as usize];
320            let ipow = if in_second_half {
321                (ic2 / idiv) % ibase
322            } else {
323                (ic1 / idiv) % ibase
324            };
325            if ipow == 0 || ctx.order_of(t.idx) > nocut + 1 {
326                continue;
327            }
328            let idx = if in_second_half {
329                ctx.ia1[ic1 as usize] + ctx.ia2[(ic2 - idiv) as usize]
330            } else {
331                ctx.ia1[(ic1 - idiv) as usize] + ctx.ia2[ic2 as usize]
332            };
333            terms.push(RawTerm {
334                idx,
335                c: t.c * f64::from(ipow),
336            });
337        }
338        Da {
339            ctx: ctx.clone(),
340            terms,
341        }
342    }
343
344    /// Repeated derivative with respect to `vars[0]`, then `vars[1]`, ...
345    pub fn deriv_vars(&self, vars: &[u32]) -> Da {
346        let mut d = self.clone();
347        for &v in vars {
348            d = d.deriv(v);
349        }
350        d
351    }
352
353    /// Integral with respect to independent variable `var` (1-based)
354    /// (`daceIntegrate`); the integration constant is zero. Out-of-range
355    /// variables warn and return zero.
356    pub fn integ(&self, var: u32) -> Da {
357        let ctx = &self.ctx;
358        if !(1..=ctx.nvmax).contains(&var) {
359            log::warn!(
360                "DACE error 624: invalid independent variable {var} in integ; returning zero DA"
361            );
362            return Da::new();
363        }
364        let (eps, nocut) = crate::context::eps_nocut();
365        let ibase = ctx.nomax + 1;
366        let j = if var > ctx.nv1 {
367            var - 1 - ctx.nv1
368        } else {
369            var - 1
370        };
371        let idiv = crate::context::npown_i64(ibase, j);
372        let in_second_half = var > ctx.nv1;
373        let mut terms = Vec::with_capacity(self.terms.len());
374        for t in &self.terms {
375            if ctx.order_of(t.idx) >= nocut {
376                continue;
377            }
378            let ic1 = ctx.ie1[t.idx as usize];
379            let ic2 = ctx.ie2[t.idx as usize];
380            let ipow = if in_second_half {
381                (ic2 / idiv) % ibase
382            } else {
383                (ic1 / idiv) % ibase
384            };
385            let ccc = t.c / f64::from(ipow + 1);
386            if crate::kernels::keep(ccc, eps) {
387                let idx = if in_second_half {
388                    ctx.ia1[ic1 as usize] + ctx.ia2[(ic2 + idiv) as usize]
389                } else {
390                    ctx.ia1[(ic1 + idiv) as usize] + ctx.ia2[ic2 as usize]
391                };
392                terms.push(RawTerm { idx, c: ccc });
393            }
394        }
395        Da {
396            ctx: ctx.clone(),
397            terms,
398        }
399    }
400
401    /// Repeated integral with respect to `vars[0]`, then `vars[1]`, ...
402    pub fn integ_vars(&self, vars: &[u32]) -> Da {
403        let mut d = self.clone();
404        for &v in vars {
405            d = d.integ(v);
406        }
407        d
408    }
409
410    /// Keep only terms of order between `min_order` and `max_order`
411    /// inclusive (`daceTrim`).
412    pub fn trim(&self, min_order: u32, max_order: u32) -> Da {
413        let terms = self
414            .terms
415            .iter()
416            .filter(|t| {
417                let io = self.ctx.order_of(t.idx);
418                io >= min_order && io <= max_order
419            })
420            .copied()
421            .collect();
422        Da {
423            ctx: self.ctx.clone(),
424            terms,
425        }
426    }
427
428    /// The multiplicative inverse `1/self` (`daceMultiplicativeInverse`):
429    /// direct alternating series below truncation order 5, Newton iteration
430    /// above.
431    ///
432    /// # Panics
433    ///
434    /// Panics with [`crate::DaceError`] code 641 ("Dividing by zero") when the
435    /// constant part of `self` is zero.
436    pub fn minv(&self) -> Da {
437        minv(self)
438    }
439
440    /// The square `self * self` (`daceSquare`).
441    pub fn sqr(&self) -> Da {
442        multiply(self, self)
443    }
444
445    /// Divide by `var^p`, when every monomial's exponent in `var` is at
446    /// least `p` (`daceDivideByVariable`): exact polynomial division on the
447    /// exponents, coefficients unchanged. `p == 0` returns a copy.
448    ///
449    /// Out-of-range variables warn and return zero. Division is impossible
450    /// when some exponent is too small (`p > nomax`, or the DA is non-zero
451    /// with insufficient exponents).
452    ///
453    /// # Panics
454    ///
455    /// Panics with [`crate::DaceError`] code 642 ("Inverse does not exists") when
456    /// the division is impossible on a non-zero DA.
457    pub fn divide_variable(&self, var: u32, p: u32) -> Da {
458        let ctx = &self.ctx;
459        if !(1..=ctx.nvmax).contains(&var) {
460            log::warn!(
461                "DACE error 624: invalid independent variable {var} in divide_variable; returning zero DA"
462            );
463            return Da::new();
464        }
465        if p == 0 {
466            return self.clone();
467        }
468        if self.terms.is_empty() {
469            return Da::new();
470        }
471        if p > ctx.nomax {
472            crate::error::dace_panic(642, "Inverse does not exists");
473        }
474        let ibase = ctx.nomax + 1;
475        let j = if var > ctx.nv1 {
476            var - 1 - ctx.nv1
477        } else {
478            var - 1
479        };
480        let idiv = crate::context::npown_i64(ibase, j);
481        let in_second_half = var > ctx.nv1;
482        let mut terms = Vec::with_capacity(self.terms.len());
483        for t in &self.terms {
484            let ic1 = ctx.ie1[t.idx as usize];
485            let ic2 = ctx.ie2[t.idx as usize];
486            let ipow = if in_second_half {
487                (ic2 / idiv) % ibase
488            } else {
489                (ic1 / idiv) % ibase
490            };
491            if ipow < p {
492                crate::error::dace_panic(642, "Inverse does not exists");
493            }
494            let idx = if in_second_half {
495                ctx.ia1[ic1 as usize] + ctx.ia2[(ic2 - p * idiv) as usize]
496            } else {
497                ctx.ia1[(ic1 - p * idiv) as usize] + ctx.ia2[ic2 as usize]
498            };
499            terms.push(RawTerm { idx, c: t.c });
500        }
501        Da {
502            ctx: ctx.clone(),
503            terms,
504        }
505    }
506
507    /// Multiply with `other` monomial-by-monomial: the coefficient-wise
508    /// product over matching monomial indices (`daceMultiplyMonomials`).
509    pub fn multiply_monomials(&self, other: &Da) -> Da {
510        Da::assert_same_context(self, other);
511        let mut terms = Vec::new();
512        let mut ib = other.terms.iter().peekable();
513        'outer: for ta in &self.terms {
514            // Advance b to the first term with idx >= ta.idx (as in C).
515            while let Some(tb) = ib.peek() {
516                if tb.idx < ta.idx {
517                    ib.next();
518                } else {
519                    break;
520                }
521            }
522            match ib.peek() {
523                Some(tb) if tb.idx == ta.idx => {
524                    terms.push(RawTerm {
525                        idx: ta.idx,
526                        c: ta.c * tb.c,
527                    });
528                }
529                Some(_) => continue 'outer,
530                None => break 'outer,
531            }
532        }
533        Da {
534            ctx: self.ctx.clone(),
535            terms,
536        }
537    }
538
539    pub(crate) fn assert_same_context(a: &Da, b: &Da) {
540        if !Arc::ptr_eq(&a.ctx, &b.ctx) {
541            std::panic::panic_any(crate::error::DaceError::new(
542                codes::NOT_INITIALIZED,
543                "mixed DACE contexts (was init() called again?)",
544            ));
545        }
546    }
547}
548
549impl Default for Da {
550    fn default() -> Da {
551        Da::new()
552    }
553}
554
555/// Pad with zeros or truncate an exponent vector to the context's number of
556/// variables, warning on length mismatch.
557fn fix_exponent_length(ctx: &Context, jj: &[u32]) -> Vec<u32> {
558    let nvar = ctx.nvmax as usize;
559    if jj.len() == nvar {
560        return jj.to_vec();
561    }
562    if jj.len() > nvar {
563        log::warn!("DACE info: exponent vector longer than the number of variables; truncating");
564        jj[..nvar].to_vec()
565    } else {
566        log::warn!("DACE info: exponent vector shorter than the number of variables; zero-padding");
567        let mut v = vec![0u32; nvar];
568        v[..jj.len()].copy_from_slice(jj);
569        v
570    }
571}
572
573// ---------------------------------------------------------------------------
574// Operator traits
575// ---------------------------------------------------------------------------
576
577impl Add for Da {
578    type Output = Da;
579    fn add(self, rhs: Da) -> Da {
580        Da::assert_same_context(&self, &rhs);
581        weighted_sum(&self, 1.0, &rhs, 1.0)
582    }
583}
584
585impl Sub for Da {
586    type Output = Da;
587    fn sub(self, rhs: Da) -> Da {
588        Da::assert_same_context(&self, &rhs);
589        weighted_sum(&self, 1.0, &rhs, -1.0)
590    }
591}
592
593impl Mul for Da {
594    type Output = Da;
595    fn mul(self, rhs: Da) -> Da {
596        Da::assert_same_context(&self, &rhs);
597        multiply(&self, &rhs)
598    }
599}
600
601impl Div for Da {
602    type Output = Da;
603    /// # Panics
604    ///
605    /// Panics with [`crate::DaceError`] code 641 when `rhs` has a zero constant
606    /// part (division by zero), via the multiplicative inverse.
607    fn div(self, rhs: Da) -> Da {
608        Da::assert_same_context(&self, &rhs);
609        multiply(&self, &rhs.minv())
610    }
611}
612
613impl Neg for Da {
614    type Output = Da;
615    fn neg(self) -> Da {
616        weighted_sum(&self, -1.0, &self, 0.0)
617    }
618}
619
620impl Add<f64> for Da {
621    type Output = Da;
622    fn add(self, rhs: f64) -> Da {
623        weighted_sum(&self, 1.0, &Da::constant(rhs), 1.0)
624    }
625}
626
627impl Sub<f64> for Da {
628    type Output = Da;
629    fn sub(self, rhs: f64) -> Da {
630        weighted_sum(&self, 1.0, &Da::constant(rhs), -1.0)
631    }
632}
633
634impl Mul<f64> for Da {
635    type Output = Da;
636    fn mul(self, rhs: f64) -> Da {
637        weighted_sum(&self, rhs, &self, 0.0)
638    }
639}
640
641impl Div<f64> for Da {
642    type Output = Da;
643    /// # Panics
644    ///
645    /// Panics with [`crate::DaceError`] code 641 when `rhs == 0.0`.
646    fn div(self, rhs: f64) -> Da {
647        if rhs == 0.0 {
648            dace_panic(codes::DIVIDING_BY_ZERO, "Dividing by zero");
649        }
650        weighted_sum(&self, 1.0 / rhs, &self, 0.0)
651    }
652}
653
654impl Add<Da> for f64 {
655    type Output = Da;
656    fn add(self, rhs: Da) -> Da {
657        weighted_sum(&Da::constant(self), 1.0, &rhs, 1.0)
658    }
659}
660
661impl Sub<Da> for f64 {
662    type Output = Da;
663    fn sub(self, rhs: Da) -> Da {
664        weighted_sum(&Da::constant(self), 1.0, &rhs, -1.0)
665    }
666}
667
668impl Mul<Da> for f64 {
669    type Output = Da;
670    fn mul(self, rhs: Da) -> Da {
671        weighted_sum(&rhs, self, &rhs, 0.0)
672    }
673}
674
675impl Div<Da> for f64 {
676    type Output = Da;
677    fn div(self, rhs: Da) -> Da {
678        Da::constant(self) / rhs
679    }
680}
681
682impl AddAssign for Da {
683    fn add_assign(&mut self, rhs: Da) {
684        *self = self.clone() + rhs;
685    }
686}
687
688impl SubAssign for Da {
689    fn sub_assign(&mut self, rhs: Da) {
690        *self = self.clone() - rhs;
691    }
692}
693
694impl MulAssign for Da {
695    fn mul_assign(&mut self, rhs: Da) {
696        *self = self.clone() * rhs;
697    }
698}
699
700impl DivAssign for Da {
701    fn div_assign(&mut self, rhs: Da) {
702        *self = self.clone() / rhs;
703    }
704}
705
706impl AddAssign<f64> for Da {
707    fn add_assign(&mut self, rhs: f64) {
708        *self = self.clone() + rhs;
709    }
710}
711
712impl SubAssign<f64> for Da {
713    fn sub_assign(&mut self, rhs: f64) {
714        *self = self.clone() - rhs;
715    }
716}
717
718impl MulAssign<f64> for Da {
719    fn mul_assign(&mut self, rhs: f64) {
720        *self = self.clone() * rhs;
721    }
722}
723
724impl DivAssign<f64> for Da {
725    fn div_assign(&mut self, rhs: f64) {
726        *self = self.clone() / rhs;
727    }
728}
729
730// ---------------------------------------------------------------------------
731// Deterministic PRNG for Da::random (documented divergence: C uses libc rand)
732// ---------------------------------------------------------------------------
733
734thread_local! {
735    static RAND_STATE: Cell<u64> = const { Cell::new(0x9E3779B97F4A7C15) };
736}
737
738/// Deterministic pseudo-random number in `[0, 1)` from a per-thread 64-bit LCG.
739pub(crate) fn dace_random() -> f64 {
740    RAND_STATE.with(|s| {
741        let state = s
742            .get()
743            .wrapping_mul(6364136223846793005)
744            .wrapping_add(1442695040888963407);
745        s.set(state);
746        (state >> 11) as f64 / (1u64 << 53) as f64
747    })
748}
749
750#[cfg(test)]
751mod tests {
752    use super::*;
753    use crate::test_support::CONTEXT_LOCK;
754
755    #[test]
756    fn arithmetic_and_calculus_basics() {
757        let _g = CONTEXT_LOCK.lock();
758        crate::context::init(3, 2).unwrap();
759        let x = Da::variable(1);
760        let y = Da::variable(2);
761
762        // (x+y) + (x-y) == 2x
763        let s = (x.clone() + y.clone()) + (x.clone() - y.clone());
764        assert_eq!(s.size(), 1);
765        assert!((s.get_coefficient(&[1, 0]) - 2.0).abs() == 0.0);
766        assert_eq!(s.get_coefficient(&[0, 1]), 0.0);
767
768        // deriv / integ roundtrip
769        assert_eq!(x.clone().deriv(1).cons(), 1.0);
770        assert_eq!(x.clone().deriv(1).size(), 1);
771        let xi = x.clone().integ(1);
772        assert!((xi.get_coefficient(&[2, 0]) - 0.5).abs() < 1e-15);
773        let xd = xi.deriv(1);
774        assert_eq!(xd.size(), 1);
775        assert!((xd.get_coefficient(&[1, 0]) - 1.0).abs() < 1e-15);
776
777        // multiplication
778        let xx = x.clone() * x.clone();
779        assert_eq!(xx.size(), 1);
780        assert!((xx.get_coefficient(&[2, 0]) - 1.0).abs() < 1e-15);
781        let xy = x.clone() * y.clone();
782        assert!((xy.get_coefficient(&[1, 1]) - 1.0).abs() < 1e-15);
783
784        // division by scalar and inverse
785        let d = (x.clone() * 2.0) / 2.0;
786        assert!((d.get_coefficient(&[1, 0]) - 1.0).abs() < 1e-15);
787        let inv = (1.0 + x.clone()).minv();
788        // 1/(1+x) = 1 - x + x^2 - x^3 at order 3
789        assert!((inv.cons() - 1.0).abs() < 1e-15);
790        assert!((inv.get_coefficient(&[1, 0]) + 1.0).abs() < 1e-15);
791        assert!((inv.get_coefficient(&[2, 0]) - 1.0).abs() < 1e-15);
792        assert!((inv.get_coefficient(&[3, 0]) + 1.0).abs() < 1e-15);
793
794        // division a/b*b ~ a
795        let b = 2.0 + x.clone() * y.clone();
796        let q = (1.0 + x.clone()) / b.clone();
797        let r = q * b;
798        for m in r.iter_monomials() {
799            let expect = if m.jj == vec![0, 0] || m.jj == vec![1, 0] {
800                1.0
801            } else {
802                0.0
803            };
804            assert!(
805                (m.c - expect).abs() <= 1e-13 * expect.abs().max(1.0),
806                "coefficient of {:?} = {}",
807                m.jj,
808                m.c
809            );
810        }
811
812        // trim keeps only constant/linear part
813        let f = (1.0 + x.clone() + y.clone()) * (x.clone() - y.clone());
814        let ft = f.trim(0, 1);
815        assert_eq!(ft.size(), 2);
816        assert!((ft.get_coefficient(&[1, 0]) - 1.0).abs() < 1e-15);
817        assert!((ft.get_coefficient(&[0, 1]) + 1.0).abs() < 1e-15);
818        assert_eq!(f.trim(2, 3).size(), 2); // x^2 - y^2
819
820        // constructors
821        assert_eq!(Da::constant(5.0).cons(), 5.0);
822        assert_eq!(Da::new().size(), 0);
823        assert_eq!(Da::default().size(), 0);
824        assert_eq!(Da::identity(2).get_coefficient(&[0, 1]), 1.0);
825        assert_eq!(Da::monomial(&[2, 1], 3.0).get_coefficient(&[2, 1]), 3.0);
826        assert_eq!(Da::variable(3).size(), 0); // out of range -> warn + zero
827
828        // inspection
829        let f = 1.0 + 2.0 * x.clone() + 3.0 * y.clone();
830        assert_eq!(f.cons(), 1.0);
831        assert_eq!(f.linear(), vec![2.0, 3.0]);
832        assert_eq!(f.size(), 3);
833        let mut g = f.clone();
834        g.set_coefficient(&[1, 1], 7.0);
835        assert_eq!(g.get_coefficient(&[1, 1]), 7.0);
836        g.set_coefficient(&[1, 1], 0.0); // |0| <= eps -> removed
837        assert_eq!(g.get_coefficient(&[1, 1]), 0.0);
838        assert_eq!(g.size(), 3);
839        assert_eq!(f.get_monomial(1).unwrap().jj, vec![0, 0]);
840        assert!(f.get_monomial(4).is_none());
841        assert_eq!(f.iter_monomials().count(), 3);
842        assert!(!f.is_nan());
843        assert!(!f.is_inf());
844        assert!(Da::constant(f64::NAN).is_nan());
845        assert!(Da::constant(f64::INFINITY).is_inf());
846
847        // gradient
848        let grad = f.gradient();
849        assert_eq!(grad.len(), 2);
850        assert_eq!(grad[0].cons(), 2.0);
851        assert_eq!(grad[1].cons(), 3.0);
852    }
853
854    #[test]
855    fn eps_flush_and_operators() {
856        let _g = CONTEXT_LOCK.lock();
857        crate::context::init(3, 2).unwrap();
858
859        let old = crate::context::set_epsilon(0.5);
860        let z = Da::constant(0.5) + Da::constant(0.25); // both <= eps -> flushed
861        assert_eq!(z.size(), 0);
862        crate::context::set_epsilon(old);
863        assert_eq!(Da::constant(0.5).cons(), 0.5);
864
865        // scalar operators
866        let f = Da::variable(1);
867        assert!(((2.0 * f.clone()).get_coefficient(&[1, 0]) - 2.0).abs() < 1e-15);
868        assert!(((f.clone() + 1.0).cons() - 1.0).abs() < 1e-15);
869        assert!(((f.clone() - 1.0).cons() + 1.0).abs() < 1e-15);
870        assert!(((1.0 + f.clone()).cons() - 1.0).abs() < 1e-15);
871        assert!(((1.0 - f.clone()).cons() - 1.0).abs() < 1e-15);
872        assert_eq!((1.0 / (1.0 + f.clone())).cons(), 1.0);
873
874        // assign operators
875        let mut a = Da::variable(1);
876        a += 1.0;
877        a *= 2.0;
878        a -= Da::constant(1.0);
879        a /= 2.0;
880        assert_eq!(a.cons(), 0.5);
881        assert!((a.get_coefficient(&[1, 0]) - 1.0).abs() < 1e-15);
882
883        // neg
884        assert_eq!((-f.clone()).get_coefficient(&[1, 0]), -1.0);
885    }
886
887    #[test]
888    fn multiplication_truncation_and_division() {
889        let _g = CONTEXT_LOCK.lock();
890        crate::context::init(6, 3).unwrap();
891        let x = Da::variable(1);
892        let y = Da::variable(2);
893        let z = Da::variable(3);
894
895        // (x*y)*(x*y) == x^2 y^2
896        let xy = x.clone() * y.clone();
897        let r = xy.clone() * xy.clone();
898        assert!((r.get_coefficient(&[2, 2, 0]) - 1.0).abs() < 1e-15);
899        assert_eq!(r.size(), 1);
900
901        // x*x*x == x^3 at order 3
902        let xxx = x.clone() * x.clone() * x.clone();
903        assert!((xxx.get_coefficient(&[3, 0, 0]) - 1.0).abs() < 1e-15);
904        assert_eq!(xxx.size(), 1);
905
906        // associativity on random Das
907        for trial in 0..5 {
908            let a = Da::random(-0.4);
909            let b = Da::random(-0.4);
910            let c = Da::random(-0.4);
911            let ab_c = (a.clone() * b.clone()) * c.clone();
912            let a_bc = a.clone() * (b.clone() * c.clone());
913            for (m1, m2) in ab_c.iter_monomials().zip(a_bc.iter_monomials()) {
914                assert_eq!(m1.jj, m2.jj);
915                let denom = m1.c.abs().max(1.0);
916                assert!(
917                    (m1.c - m2.c).abs() <= 1e-12 * denom,
918                    "trial {trial}: {:?} {} vs {}",
919                    m1.jj,
920                    m1.c,
921                    m2.c
922                );
923            }
924            assert_eq!(ab_c.size(), a_bc.size());
925        }
926
927        // truncation: at nocut=2, (x+y)^3 has no terms
928        crate::context::set_truncation_order(2);
929        let s = x.clone() + y.clone();
930        let cube = s.clone() * s.clone() * s.clone();
931        assert_eq!(cube.size(), 0);
932        crate::context::set_truncation_order(3);
933
934        // a/b*b ~ a (rtol 1e-13) for b with cons != 0
935        let a = 1.0 + x.clone() + 0.5 * z.clone();
936        let b = 2.0 + x.clone() * y.clone() - 0.3 * z.clone() * z.clone();
937        let q = a.clone() / b.clone();
938        let r = q * b;
939        for m in r.iter_monomials() {
940            let expect = a.get_coefficient(&m.jj);
941            let denom = expect.abs().max(1.0);
942            assert!(
943                (m.c - expect).abs() <= 1e-13 * denom,
944                "{:?}: {} vs {}",
945                m.jj,
946                m.c,
947                expect
948            );
949        }
950        assert_eq!(r.size(), a.size());
951
952        // divide_variable on x^2 y by (1,1) == xy; by (1,3) panics
953        let x2y = x.clone() * x.clone() * y.clone();
954        let d = x2y.clone().divide_variable(1, 1);
955        assert!((d.get_coefficient(&[1, 1, 0]) - 1.0).abs() < 1e-15);
956        assert_eq!(d.size(), 1);
957        assert_eq!(x2y.divide_variable(1, 2).get_coefficient(&[0, 1, 0]), 1.0);
958        let result = std::panic::catch_unwind(|| x2y.divide_variable(1, 3));
959        assert!(result.is_err());
960
961        // multiply_monomials: coefficient-wise product on matching indices
962        let p = (1.0 + x.clone() + y.clone()).multiply_monomials(&(2.0 + 3.0 * y.clone()));
963        assert_eq!(p.size(), 2);
964        assert!((p.cons() - 2.0).abs() < 1e-15);
965        assert!((p.get_coefficient(&[0, 1]) - 3.0).abs() < 1e-15);
966
967        // fma == weighted sum
968        let f = crate::fma(&x.clone(), 2.0, &y.clone(), -1.0);
969        assert!((f.get_coefficient(&[1, 0, 0]) - 2.0).abs() < 1e-15);
970        assert!((f.get_coefficient(&[0, 1, 0]) + 1.0).abs() < 1e-15);
971    }
972}