Skip to main content

dace_rs/
context.rs

1//! DACE computation context: initialization, monomial index encoding, and
2//! thread-local computation settings.
3//!
4//! A DACE computation is parameterized by a maximum computation order `no` and
5//! a number of variables `nv`. Monomials are stored as sparse lists of
6//! `{coefficient, index}` pairs, where the index is a canonical packed
7//! position `0..nmmax-1` built from a base-`(no+1)` digit encoding of the
8//! exponent vector, split into two halves of `nv1 = (nv+1)/2` and
9//! `nv2 = nv - nv1` variables (this is what makes the reverse lookup tables
10//! fit in 32 bits for useful `(no, nv)` combinations).
11//!
12//! This module ports the C library's `daceInitialize` (`core/daceinit.c`) and
13//! the encoding machinery of `core/daceaux.c`. Divergences from the C library:
14//!
15//! - [`init`] may be called again at any time; previously created [`Da`][crate::Da]
16//!   values keep working against their original context (in C,
17//!   `daceInitialize` purges all existing DA objects).
18//! - after any re-initialization, every thread re-derives its computation
19//!   settings lazily on next use (in C, only the thread that called
20//!   `daceInitialize` is reset; other threads keep stale settings).
21//! - informational messages go through the [`log`] crate instead of stderr.
22
23use std::cell::{Cell, RefCell};
24use std::sync::atomic::{AtomicU64, Ordering};
25use std::sync::{Arc, LazyLock};
26
27use parking_lot::RwLock;
28
29use crate::error::{DaceError, codes};
30
31/// Immutable computation context: lookup tables and limits.
32///
33/// Built by [`init`]; shared (via `Arc`) by every [`Da`][crate::Da] created
34/// while it is the active context.
35// Encoding/decoding machinery is wired into `Da` in Phase 2 and the
36// computation kernels in Phase 3; until then the fields are only exercised
37// by the tests below.
38#[derive(Debug)]
39#[allow(dead_code)]
40pub(crate) struct Context {
41    pub nomax: u32,
42    pub nvmax: u32,
43    pub nv1: u32,
44    pub nv2: u32,
45    pub nmmax: u32,
46    pub epsmac: f64,
47    /// Encoded exponents of the first `nv1` variables, by monomial index
48    /// (length `nmmax`).
49    pub ie1: Vec<u32>,
50    /// Encoded exponents of the last `nv2` variables, by monomial index
51    /// (length `nmmax`).
52    pub ie2: Vec<u32>,
53    /// Total order of each monomial index (length `nmmax`).
54    pub ieo: Vec<u32>,
55    /// Base monomial index by encoded exponents of the first half
56    /// (length `lia+1`).
57    pub ia1: Vec<u32>,
58    /// Monomial offset by encoded exponents of the second half
59    /// (length `lia+1`).
60    pub ia2: Vec<u32>,
61    /// Generation of the global context this snapshot belongs to; used to
62    /// invalidate thread-local scratch buffers after re-initialization.
63    pub generation: u64,
64}
65
66#[allow(dead_code)] // encode/decode/order_of are wired into Da in Phase 2
67impl Context {
68    /// Encode an exponent vector (length `nv`) into its monomial index.
69    ///
70    /// Returns `None` if the slice length differs from `nv`, any single
71    /// exponent exceeds `nomax` (unrepresentable in base `nomax+1`), or the
72    /// total order exceeds `nomax` (C error 622).
73    pub(crate) fn encode(&self, jj: &[u32]) -> Option<u32> {
74        if jj.len() != self.nvmax as usize {
75            return None;
76        }
77        let base = self.nomax + 1;
78        let mut io: u32 = 0;
79        let mut ic1: u32 = 0;
80        let mut ic2: u32 = 0;
81        for &e in jj[self.nv1 as usize..].iter().rev() {
82            if e > self.nomax {
83                return None;
84            }
85            ic2 = ic2 * base + e;
86            io += e;
87        }
88        for &e in jj[..self.nv1 as usize].iter().rev() {
89            if e > self.nomax {
90                return None;
91            }
92            ic1 = ic1 * base + e;
93            io += e;
94        }
95        if io > self.nomax {
96            return None;
97        }
98        Some(self.ia1[ic1 as usize] + self.ia2[ic2 as usize])
99    }
100
101    /// Decode a monomial index into its exponent vector (length `nv`).
102    ///
103    /// # Panics
104    ///
105    /// Panics if `ii >= nmmax` (C error 626, invalid encoded exponent).
106    pub(crate) fn decode(&self, ii: u32) -> Vec<u32> {
107        let mut jj = vec![0u32; self.nvmax as usize];
108        self.decode_into(ii, &mut jj);
109        jj
110    }
111
112    /// Decode a monomial index into a caller-provided exponent vector.
113    pub(crate) fn decode_into(&self, ii: u32, jj: &mut [u32]) {
114        assert!(jj.len() >= self.nvmax as usize, "decode buffer too short");
115        if ii >= self.nmmax {
116            crate::error::dace_panic(codes::INVALID_ENCODED_EXPONENT, "Invalid encoded exponent");
117        }
118        let base = self.nomax + 1;
119        let mut ic = self.ie1[ii as usize];
120        for slot in jj[..self.nv1 as usize].iter_mut() {
121            *slot = ic % base;
122            ic /= base;
123        }
124        let mut ic = self.ie2[ii as usize];
125        for slot in jj[self.nv1 as usize..self.nvmax as usize].iter_mut() {
126            *slot = ic % base;
127            ic /= base;
128        }
129    }
130
131    /// Total order of the monomial with the given index.
132    pub(crate) fn order_of(&self, ii: u32) -> u32 {
133        self.ieo[ii as usize]
134    }
135}
136
137static CONTEXT: LazyLock<RwLock<Option<Arc<Context>>>> = LazyLock::new(|| RwLock::new(None));
138static GENERATION: AtomicU64 = AtomicU64::new(0);
139
140impl Context {
141    /// The active context.
142    ///
143    /// # Panics
144    ///
145    /// Panics with [`DaceError`] code 1003 if DACE has not been initialized.
146    pub(crate) fn current() -> Arc<Context> {
147        match CONTEXT.read().clone() {
148            Some(ctx) => ctx,
149            None => {
150                crate::error::dace_panic(codes::NOT_INITIALIZED, "DACE has not been initialized")
151            }
152        }
153    }
154}
155
156/// Whether DACE has been initialized on this process.
157pub fn initialized() -> bool {
158    CONTEXT.read().is_some()
159}
160
161/// DACE version this crate reproduces (`"2.1.0-rs"`).
162pub fn version() -> &'static str {
163    "2.1.0-rs"
164}
165
166/// Initialize DACE with computation order `order` and `nvars` variables,
167/// replacing any previous context.
168///
169/// Values of `order`/`nvars` below 1 are clamped to 1 with a warning (C
170/// informational messages 167/168). Returns an error if the required lookup
171/// tables do not fit in a 32-bit index space (C error 911).
172///
173/// Unlike the C library, re-initializing does **not** invalidate existing
174/// [`Da`][crate::Da] values: they keep operating on their original context.
175/// After any re-initialization, every thread re-derives its computation
176/// settings (epsilon cutoff, truncation order) lazily on next use, as if the
177/// thread had never been used; user-set settings therefore do not survive
178/// re-initialization on any thread.
179pub fn init(order: u32, nvars: u32) -> Result<(), DaceError> {
180    let mut no = order;
181    let mut nv = nvars;
182    if no < 1 {
183        log::warn!("DACE info 167: computation order increased to 1");
184        no = 1;
185    }
186    if nv < 1 {
187        log::warn!("DACE info 168: number of variables increased to 1");
188        nv = 1;
189    }
190
191    // Machine epsilon, computed as in the C library.
192    let mut epsmac = 1.0f64;
193    while 1.0 + epsmac > 1.0 {
194        epsmac /= 2.0;
195    }
196    epsmac *= 2.0;
197
198    // Length of the reverse lookup arrays must fit the 32 bit index space.
199    let nv1 = nv.div_ceil(2);
200    let clia = pown(f64::from(no + 1), nv1);
201    if clia >= pown(2.0, 32) {
202        return Err(DaceError::new(
203            codes::ORDER_VARIABLE_TOO_LARGE,
204            "Order and/or variable too large",
205        ));
206    }
207    let lia = clia as u32;
208    let nmmax = count_monomials(no, nv);
209
210    let mut ie1 = vec![0u32; nmmax as usize];
211    let mut ie2 = vec![0u32; nmmax as usize];
212    let mut ieo = vec![0u32; nmmax as usize];
213    let mut ia1 = vec![0u32; lia as usize + 1];
214    let mut ia2 = vec![0u32; lia as usize + 1];
215
216    // Fill the addressing arrays, enumerating ordered monomials exactly like
217    // the C implementation (core/daceinit.c:139-152).
218    let nv2 = nv - nv1;
219    let mut p1 = vec![0u32; nv1 as usize];
220    let mut p2 = vec![0u32; nv2 as usize];
221    let mut i: u32 = 0;
222    let mut no1: u32;
223    let mut no2: u32;
224    loop {
225        let exp1 = encode_exponents(&p1, no);
226        let i0 = i;
227        ia1[exp1 as usize] = i0;
228        no1 = p1.iter().sum();
229        loop {
230            ie1[i as usize] = exp1;
231            let exp2 = encode_exponents(&p2, no);
232            ie2[i as usize] = exp2;
233            ieo[i as usize] = no1 + p2.iter().sum::<u32>();
234            ia2[exp2 as usize] = i - i0;
235            i += 1;
236            no2 = next_ordered_monomial(&mut p2, no - no1);
237            if no2 == 0 {
238                break;
239            }
240        }
241        no1 = next_ordered_monomial(&mut p1, no);
242        if no1 == 0 {
243            break;
244        }
245    }
246
247    // Cross-checks mirroring the C PANIC 5/6 internal invariants.
248    if i != nmmax {
249        crate::error::dace_panic(1005, "Incorrect number of monomials");
250    }
251    for i in 0..nmmax as usize {
252        let nn = ia1[ie1[i] as usize] + ia2[ie2[i] as usize];
253        if nn != i as u32 {
254            crate::error::dace_panic(1006, "Incorrect DA coding arrays");
255        }
256    }
257
258    let generation = GENERATION.fetch_add(1, Ordering::SeqCst) + 1;
259    let ctx = Arc::new(Context {
260        nomax: no,
261        nvmax: nv,
262        nv1,
263        nv2,
264        nmmax,
265        epsmac,
266        ie1,
267        ie2,
268        ieo,
269        ia1,
270        ia2,
271        generation,
272    });
273    *CONTEXT.write() = Some(ctx);
274
275    Ok(())
276}
277
278/// Raise `a` to the positive integer power `b` (binary exponentiation, as in
279/// the C library's `pown`).
280pub(crate) fn pown(a: f64, b: u32) -> f64 {
281    let mut res = 1.0;
282    let mut a = a;
283    let mut b = b;
284    while b > 0 {
285        if b & 1 != 0 {
286            res *= a;
287        }
288        a *= a;
289        b >>= 1;
290    }
291    res
292}
293
294/// Raise integer `a` to the positive integer power `b` (the C library's
295/// `npown`), computed in `u64`. Callers guarantee the result fits the
296/// relevant table bound (it indexes `ia1`/`ia2`, both of length `lia+1`).
297pub(crate) fn npown_i64(a: u32, b: u32) -> u32 {
298    let mut res: u64 = 1;
299    let mut a: u64 = u64::from(a);
300    let mut b = b;
301    while b > 0 {
302        if b & 1 != 0 {
303            res *= a;
304        }
305        a *= a;
306        b >>= 1;
307    }
308    res as u32
309}
310
311/// Number of monomials of maximum order `no` in `nv` variables, i.e.
312/// `C(no+nv, min(no,nv))` (the C library's `daceCountMonomials`).
313pub(crate) fn count_monomials(no: u32, nv: u32) -> u32 {
314    let mut dnumda = 1.0f64;
315    let mm = nv.max(no);
316    for i in 1..=nv.min(no) {
317        dnumda = dnumda * f64::from(mm + i) / f64::from(i);
318    }
319    dnumda as u32
320}
321
322/// Encode `nv` exponents (each at most `no`) into one base-`(no+1)` integer
323/// (the C library's `daceEncodeExponents`).
324fn encode_exponents(p: &[u32], no: u32) -> u32 {
325    if p.is_empty() {
326        return 0;
327    }
328    let base = no + 1;
329    let mut res = p[p.len() - 1];
330    for &e in p[..p.len() - 1].iter().rev() {
331        res = res * base + e;
332    }
333    res
334}
335
336/// Advance `p` to the next monomial of `nv` variables in arbitrary order
337/// (the C library's `daceNextMonomial`); returns the new order, or 0 when
338/// wrapping back to the constant monomial.
339fn next_monomial(p: &mut [u32], no: u32) -> u32 {
340    let mut o: u32 = p.iter().sum();
341    for e in p.iter_mut() {
342        if o < no {
343            *e += 1;
344            return o + 1;
345        }
346        o -= *e;
347        *e = 0;
348    }
349    0
350}
351
352/// Advance `p` to the next monomial in order-sorted enumeration (the C
353/// library's `daceNextOrderedMonomial`).
354fn next_ordered_monomial(p: &mut [u32], no: u32) -> u32 {
355    if p.is_empty() || no == 0 {
356        return 0;
357    }
358    let mut o: u32 = p.iter().sum();
359    let oo = next_monomial(&mut p[1..], o);
360    if oo == 0 {
361        o = (o + 1) % (no + 1); // jump to next order
362    }
363    p[0] = o - oo; // complete the monomial up to order o
364    o
365}
366
367// ---------------------------------------------------------------------------
368// Thread-local computation settings (C DACECom_t)
369// ---------------------------------------------------------------------------
370
371struct Settings {
372    eps: Cell<f64>,
373    nocut: Cell<u32>,
374    ready: Cell<bool>,
375    generation: Cell<u64>,
376    stack: RefCell<Vec<u32>>,
377}
378
379thread_local! {
380    static SETTINGS: Settings = const {
381        Settings {
382            eps: Cell::new(0.0),
383            nocut: Cell::new(0),
384            ready: Cell::new(false),
385            generation: Cell::new(0),
386            stack: RefCell::new(Vec::new()),
387        }
388    };
389}
390
391/// Lazily initialize this thread's settings from the active context on first
392/// use, and re-derive them whenever the context generation advances (a new
393/// [`init`]) — on every thread, not just the initializing one (C
394/// `daceInitializeThread0`: eps = 0, nocut = nomax).
395fn with_settings<R>(f: impl FnOnce(&Settings) -> R) -> R {
396    SETTINGS.with(|s| {
397        let current = GENERATION.load(Ordering::Relaxed);
398        if !s.ready.get() || s.generation.get() != current {
399            let ctx = Context::current(); // panics (1003) when never initialized
400            s.eps.set(0.0);
401            s.nocut.set(ctx.nomax);
402            s.stack.borrow_mut().clear();
403            s.generation.set(ctx.generation);
404            s.ready.set(true);
405        }
406        f(s)
407    })
408}
409
410/// Current coefficient cutoff epsilon (coefficients with `|c| <= eps` are
411/// flushed to zero).
412///
413/// Initialized to `0.0` (cutoff disabled).
414pub fn epsilon() -> f64 {
415    with_settings(|s| s.eps.get())
416}
417
418/// Set the coefficient cutoff epsilon to `eps` (its absolute value is used)
419/// and return the previous value.
420///
421/// # Warning
422///
423/// Flushing occurs for any intermediate result also within the engine, and
424/// can produce wrong results whenever DA coefficients become very small
425/// relative to epsilon (e.g. a division by a large DA divisor can flush the
426/// internally computed inverse entirely to zero).
427pub fn set_epsilon(eps: f64) -> f64 {
428    with_settings(|s| {
429        let old = s.eps.get();
430        s.eps.set(eps.abs());
431        old
432    })
433}
434
435/// The experimentally determined machine epsilon of the active context.
436///
437/// # Panics
438///
439/// Panics if DACE has not been initialized.
440pub fn machine_epsilon() -> f64 {
441    Context::current().epsmac
442}
443
444/// The maximum computation order of the active context.
445///
446/// # Panics
447///
448/// Panics if DACE has not been initialized.
449pub fn max_order() -> u32 {
450    Context::current().nomax
451}
452
453/// The number of variables of the active context.
454///
455/// # Panics
456///
457/// Panics if DACE has not been initialized.
458pub fn max_variables() -> u32 {
459    Context::current().nvmax
460}
461
462/// The total number of monomials of the active context.
463///
464/// # Panics
465///
466/// Panics if DACE has not been initialized.
467pub fn max_monomials() -> u32 {
468    Context::current().nmmax
469}
470
471/// The current truncation order (order above which computed terms are dropped).
472pub fn truncation_order() -> u32 {
473    with_settings(|s| s.nocut.get())
474}
475
476/// Set the truncation order, clamped to `[1, nomax]` (with a warning when
477/// clamping, C informational message 162), and return the previous value.
478pub fn set_truncation_order(order: u32) -> u32 {
479    with_settings(|s| {
480        let ctx = Context::current();
481        if order > ctx.nomax {
482            log::warn!(
483                "DACE info 162: truncation order too high, clamping to {}",
484                ctx.nomax
485            );
486        }
487        let old = s.nocut.get();
488        s.nocut.set(order.min(ctx.nomax).max(1));
489        old
490    })
491}
492
493/// Push the current truncation order on this thread's stack and set a new one
494/// (clamped to `[1, nomax]`).
495pub fn push_truncation_order(order: u32) {
496    with_settings(|s| {
497        let ctx = Context::current();
498        if order > ctx.nomax {
499            log::warn!(
500                "DACE info 162: truncation order too high, clamping to {}",
501                ctx.nomax
502            );
503        }
504        s.stack.borrow_mut().push(s.nocut.get());
505        s.nocut.set(order.min(ctx.nomax).max(1));
506    });
507}
508
509/// Pop the truncation order stack, restoring the value saved by the matching
510/// [`push_truncation_order`].
511///
512/// # Panics
513///
514/// Panics (C error 161) if the stack is empty.
515pub fn pop_truncation_order() {
516    with_settings(|s| match s.stack.borrow_mut().pop() {
517        Some(nocut) => s.nocut.set(nocut),
518        None => crate::error::dace_panic(161, "Free or invalid variable"),
519    });
520}
521
522/// Read the active settings (eps, nocut) in one go; internal helper for the
523/// computation kernels.
524#[allow(dead_code)] // used by the computation kernels from Phase 3
525pub(crate) fn eps_nocut() -> (f64, u32) {
526    with_settings(|s| (s.eps.get(), s.nocut.get()))
527}
528
529/// Current context generation; internal helper for scratch invalidation.
530#[allow(dead_code)] // used by thread-local scratch invalidation from Phase 3
531pub(crate) fn generation() -> u64 {
532    GENERATION.load(Ordering::Relaxed)
533}
534
535#[cfg(test)]
536mod tests {
537    use super::*;
538    use crate::test_support::CONTEXT_LOCK;
539    fn binom(n: u64, k: u64) -> u64 {
540        let mut r = 1u64;
541        for i in 1..=k {
542            r = r * (n - k + i) / i;
543        }
544        r
545    }
546
547    #[test]
548    fn encoding_roundtrip_all_indices() {
549        let _g = CONTEXT_LOCK.lock();
550        for &(no, nv) in &[(3u32, 2u32), (5, 3), (10, 7), (1, 1)] {
551            init(no, nv).unwrap();
552            let ctx = Context::current();
553            assert_eq!(ctx.nv1 + ctx.nv2, nv);
554            assert_eq!(ctx.nv1, nv.div_ceil(2));
555            assert_eq!(ctx.nmmax as u64, binom(u64::from(no + nv), u64::from(nv)));
556            for ii in 0..ctx.nmmax {
557                let jj = ctx.decode(ii);
558                assert_eq!(jj.len(), nv as usize);
559                let re = ctx.encode(&jj).expect("valid monomial must encode");
560                assert_eq!(
561                    re, ii,
562                    "encode(decode({ii})) mismatch at (no,nv)=({no},{nv})"
563                );
564                let order: u32 = jj.iter().sum();
565                assert_eq!(ctx.order_of(ii), order);
566                assert!(order <= no);
567                // decode_into agrees with decode
568                let mut buf = vec![0u32; nv as usize];
569                ctx.decode_into(ii, &mut buf);
570                assert_eq!(buf, jj);
571            }
572        }
573    }
574
575    #[test]
576    fn encode_rejects_invalid() {
577        let _g = CONTEXT_LOCK.lock();
578        init(3, 2).unwrap();
579        let ctx = Context::current();
580        assert_eq!(ctx.encode(&[0, 0]), Some(0));
581        assert_eq!(ctx.encode(&[4, 0]), None); // exponent above nomax
582        assert_eq!(ctx.encode(&[2, 2]), None); // total order above nomax
583        assert_eq!(ctx.encode(&[1]), None); // wrong length
584    }
585
586    #[test]
587    fn init_clamps_and_errors() {
588        let _g = CONTEXT_LOCK.lock();
589        init(0, 0).unwrap();
590        assert_eq!(max_order(), 1);
591        assert_eq!(max_variables(), 1);
592        let err = init(100, 21).unwrap_err();
593        assert_eq!(err.code, codes::ORDER_VARIABLE_TOO_LARGE);
594        // previous context still active after failed init
595        assert_eq!(max_order(), 1);
596        assert!(initialized());
597    }
598
599    #[test]
600    fn settings_roundtrip() {
601        let _g = CONTEXT_LOCK.lock();
602        init(5, 2).unwrap();
603        assert_eq!(epsilon(), 0.0);
604        assert_eq!(truncation_order(), 5);
605        assert_eq!(set_epsilon(-0.5), 0.0);
606        assert_eq!(epsilon(), 0.5);
607        assert_eq!(set_epsilon(0.0), 0.5);
608        assert_eq!(set_truncation_order(3), 5);
609        assert_eq!(truncation_order(), 3);
610        assert_eq!(set_truncation_order(99), 3); // clamps to nomax=5
611        assert_eq!(truncation_order(), 5);
612        assert_eq!(set_truncation_order(0), 5); // clamps to 1
613        assert_eq!(truncation_order(), 1);
614        push_truncation_order(2);
615        assert_eq!(truncation_order(), 2);
616        pop_truncation_order();
617        assert_eq!(truncation_order(), 1);
618        // init resets settings
619        init(4, 3).unwrap();
620        assert_eq!(truncation_order(), 4);
621        assert_eq!(epsilon(), 0.0);
622        assert_eq!(version(), "2.1.0-rs");
623        assert!(machine_epsilon() > 0.0 && machine_epsilon() <= f64::EPSILON * 2.0);
624    }
625
626    #[test]
627    fn truncation_stack_empty_pop_panics() {
628        let _g = CONTEXT_LOCK.lock();
629        init(5, 2).unwrap();
630        let result = std::panic::catch_unwind(|| {
631            pop_truncation_order();
632        });
633        assert!(result.is_err());
634    }
635}