Skip to main content

probl_engine/
analytic.rs

1//! Identity-preserving continuous outcomes during enumeration.
2//!
3//! A draw is an affine function of one latent variable. Evidence is a set of
4//! intervals in that variable's CDF coordinates. Worlds own their restrictions;
5//! values and closures remain immutable, and calls return restrictions alongside
6//! their results. Independent latent combinations require sampling for now.
7
8use crate::continuous::Family;
9use crate::dist::Budget;
10use crate::error::{Fault, OpError, OpResult};
11use crate::value::{Closure, Value, family_key, float_key};
12use probl_syntax::ast::BinOp;
13use std::collections::BTreeMap;
14use std::sync::Arc;
15
16pub type Constraints = Arc<BTreeMap<u64, Domain>>;
17
18/// Sorted, disjoint CDF intervals. Endpoints have zero probability.
19#[derive(Clone, Debug, PartialEq)]
20pub struct Domain(pub Vec<(f64, f64)>);
21impl Domain {
22    pub fn full() -> Self {
23        Self(vec![(0.0, 1.0)])
24    }
25    pub fn mass(&self) -> f64 {
26        self.0.iter().map(|(a, b)| b - a).sum()
27    }
28    pub fn intersect(&self, other: &Self) -> Self {
29        let (mut i, mut j) = (0, 0);
30        let mut out = Vec::new();
31        while i < self.0.len() && j < other.0.len() {
32            let (a, b) = self.0[i];
33            let (c, d) = other.0[j];
34            if a.max(c) < b.min(d) {
35                out.push((a.max(c), b.min(d)));
36            }
37            if b < d {
38                i += 1;
39            } else {
40                j += 1;
41            }
42        }
43        Self(out)
44    }
45    pub fn complement(&self) -> Self {
46        let mut out = Vec::new();
47        let mut lo = 0.0;
48        for &(a, b) in &self.0 {
49            if lo < a {
50                out.push((lo, a));
51            }
52            lo = b;
53        }
54        if lo < 1.0 {
55            out.push((lo, 1.0));
56        }
57        Self(out)
58    }
59    fn below(p: f64) -> Self {
60        if p > 0.0 { Self(vec![(0.0, p)]) } else { Self(vec![]) }
61    }
62    fn key(&self) -> Vec<(u64, u64)> {
63        self.0.iter().map(|(a, b)| (float_key(*a), float_key(*b))).collect()
64    }
65}
66impl Eq for Domain {}
67impl std::hash::Hash for Domain {
68    fn hash<H: std::hash::Hasher>(&self, h: &mut H) {
69        self.key().hash(h);
70    }
71}
72impl Ord for Domain {
73    fn cmp(&self, b: &Self) -> std::cmp::Ordering {
74        self.key().cmp(&b.key())
75    }
76}
77impl PartialOrd for Domain {
78    fn partial_cmp(&self, b: &Self) -> Option<std::cmp::Ordering> {
79        Some(self.cmp(b))
80    }
81}
82
83#[derive(Clone, Debug)]
84pub struct Analytic {
85    pub id: u64,
86    pub family: Family,
87    pub domain: Domain,
88    pub scale: f64,
89    pub offset: f64,
90}
91impl Analytic {
92    pub fn new(id: u64, family: Family) -> Self {
93        Self {
94            id,
95            family,
96            domain: Domain::full(),
97            scale: 1.0,
98            offset: 0.0,
99        }
100    }
101    pub fn key(&self) -> impl Ord + std::hash::Hash + use<> {
102        (
103            self.id,
104            family_key(&self.family),
105            self.domain.key(),
106            float_key(self.scale),
107            float_key(self.offset),
108        )
109    }
110    pub fn value(self) -> OpResult<Value> {
111        if !self.scale.is_finite() || !self.offset.is_finite() {
112            return Err(unsupported("non-finite affine coefficients"));
113        }
114        if self.scale == 0.0 {
115            Ok(Value::Float(self.offset))
116        } else {
117            Ok(Value::Analytic(Arc::new(self)))
118        }
119    }
120    pub fn cdf(&self, x: f64) -> f64 {
121        let p = self.family.cdf((x - self.offset) / self.scale);
122        let below = self.domain.intersect(&Domain::below(p)).mass() / self.domain.mass();
123        (if self.scale > 0.0 { below } else { 1.0 - below }).clamp(0.0, 1.0)
124    }
125    pub fn pdf(&self, x: f64) -> f64 {
126        let x = (x - self.offset) / self.scale;
127        let u = self.family.cdf(x);
128        if self.domain.0.iter().any(|(a, b)| *a <= u && u <= *b) {
129            self.family.pdf(x) / (self.scale.abs() * self.domain.mass())
130        } else {
131            0.0
132        }
133    }
134    pub fn quantile(&self, q: f64) -> f64 {
135        let mut left = q * self.domain.mass();
136        let n = self.domain.0.len();
137        for i in 0..n {
138            let (a, b) = self.domain.0[if self.scale > 0.0 { i } else { n - 1 - i }];
139            if left <= b - a || i + 1 == n {
140                let u = if self.scale > 0.0 { a + left } else { b - left };
141                return self.scale * self.family.quantile(u.clamp(a, b)) + self.offset;
142            }
143            left -= b - a;
144        }
145        f64::NAN
146    }
147    pub fn moments(&self) -> (f64, f64) {
148        let (mean, variance) = if self.domain == Domain::full() {
149            (self.family.mean(), self.family.variance())
150        } else {
151            let parts: Vec<_> = self
152                .domain
153                .0
154                .iter()
155                .map(|&(a, b)| {
156                    let (m, v) = self
157                        .family
158                        .interval_moments(self.family.quantile(a), self.family.quantile(b));
159                    (m, v, (b - a) / self.domain.mass())
160                })
161                .collect();
162            let mean = parts.iter().map(|(m, _, p)| m * p).sum::<f64>();
163            let variance = parts.iter().map(|(m, v, p)| p * (v + (m - mean).powi(2))).sum();
164            (mean, variance)
165        };
166        (self.scale * mean + self.offset, self.scale * self.scale * variance)
167    }
168
169    pub fn sd(&self) -> f64 {
170        if self.domain == Domain::full() {
171            self.scale.abs() * self.family.sd()
172        } else {
173            libm::sqrt(self.moments().1)
174        }
175    }
176}
177
178#[derive(Clone, Debug)]
179pub struct Event {
180    pub draw: Analytic,
181    pub yes: Domain,
182}
183impl Event {
184    pub fn key(&self) -> impl Ord + std::hash::Hash + use<> {
185        (self.draw.key(), self.yes.key())
186    }
187    pub fn probability(&self) -> f64 {
188        (self.draw.domain.intersect(&self.yes).mass() / self.draw.domain.mass()).clamp(0.0, 1.0)
189    }
190    pub fn value(self) -> Value {
191        let p = self.probability();
192        if p == 0.0 || p == 1.0 {
193            Value::Bool(p == 1.0)
194        } else {
195            Value::Event(Arc::new(self))
196        }
197    }
198    pub fn restrict(&self, yes: bool, context: &mut Constraints) -> f64 {
199        let prior = &self.draw.domain;
200        let domain = prior.intersect(&if yes { self.yes.clone() } else { self.yes.complement() });
201        let p = domain.mass() / prior.mass();
202        Arc::make_mut(context).insert(self.draw.id, domain);
203        p.clamp(0.0, 1.0)
204    }
205}
206
207pub fn unsupported(what: &str) -> OpError {
208    OpError::unsupported(format!("{what} isn't supported for analytic continuous draws yet"))
209        .help("use `@mode sample(runs: 10_000)` for this operation")
210}
211
212pub fn binary(op: BinOp, a: &Value, b: &Value) -> OpResult<Value> {
213    use BinOp::*;
214    let (mut x, reverse, other) = match (a, b) {
215        (Value::Analytic(x), b) => ((**x).clone(), false, b),
216        (a, Value::Analytic(x)) => ((**x).clone(), true, a),
217        _ => return Err(unsupported("this operation")),
218    };
219    let (scale, offset) = match other {
220        Value::Analytic(y) if x.id == y.id => (y.scale, y.offset),
221        Value::Analytic(_) => return Err(unsupported("combining independent continuous draws")),
222        v => (
223            0.0,
224            v.as_f64()
225                .filter(|v| v.is_finite())
226                .ok_or_else(|| unsupported("this operand"))?,
227        ),
228    };
229    match op {
230        Add => {
231            x.scale += scale;
232            x.offset += offset;
233        }
234        Sub => {
235            x.scale -= scale;
236            x.offset -= offset;
237            if reverse {
238                x.scale = -x.scale;
239                x.offset = -x.offset;
240            }
241        }
242        Mul if scale == 0.0 => {
243            x.scale *= offset;
244            x.offset *= offset;
245        }
246        Div if !reverse && scale == 0.0 => {
247            if offset == 0.0 {
248                return Err(OpError::fault(Fault::DivisionByZero, "division by zero"));
249            }
250            x.scale /= offset;
251            x.offset /= offset;
252        }
253        Eq | Ne | Lt | Le | Gt | Ge => {
254            x.scale -= scale;
255            x.offset -= offset;
256            if reverse {
257                x.scale = -x.scale;
258                x.offset = -x.offset;
259            }
260            if !x.scale.is_finite() || !x.offset.is_finite() {
261                return Err(unsupported("non-finite affine coefficients"));
262            }
263            if x.scale == 0.0 {
264                return Ok(Value::Bool(match op {
265                    Eq => x.offset == 0.0,
266                    Ne => x.offset != 0.0,
267                    Lt => x.offset < 0.0,
268                    Le => x.offset <= 0.0,
269                    Gt => x.offset > 0.0,
270                    _ => x.offset >= 0.0,
271                }));
272            }
273            if op == Eq || op == Ne {
274                return Ok(Value::Bool(op == Ne));
275            }
276            let p = x.family.cdf(-x.offset / x.scale);
277            if !p.is_finite() {
278                return Err(unsupported("this numerically unstable comparison"));
279            }
280            let below = Domain::below(p);
281            let yes = if matches!(op, Lt | Le) == (x.scale > 0.0) {
282                below
283            } else {
284                below.complement()
285            };
286            return Ok(Event { draw: x, yes }.value());
287        }
288        _ => return Err(unsupported("nonlinear arithmetic")),
289    }
290    x.value()
291}
292
293pub fn logic(and: bool, a: &Value, b: &Value) -> OpResult<Value> {
294    let (event, other) = match (a, b) {
295        (Value::Event(e), b) | (b, Value::Event(e)) => (e, b),
296        _ => return Err(unsupported("this logical operation")),
297    };
298    let mut e = (**event).clone();
299    match other {
300        Value::Bool(x) => {
301            if *x != and {
302                return Ok(Value::Bool(*x));
303            }
304        }
305        Value::Event(other) if e.draw.id == other.draw.id => {
306            e.yes = if and {
307                e.yes.intersect(&other.yes)
308            } else {
309                e.yes.complement().intersect(&other.yes.complement()).complement()
310            };
311        }
312        _ => {
313            return Err(unsupported(
314                "combining this event with another probability or independent draw",
315            ));
316        }
317    }
318    Ok(e.value())
319}
320
321/// Scan through aggregate values too: symbolic numbers must never become
322/// categorical keys or be compared by their internal identity as user data.
323pub fn contains(v: &Value) -> bool {
324    match v {
325        Value::Analytic(_) | Value::Event(_) => return true,
326        Value::List(_) | Value::Map(_) | Value::Bag(_) | Value::Record(_) | Value::Dist(_) | Value::Closure(_) => {}
327        _ => return false,
328    }
329    let mut pending = vec![v];
330    while let Some(v) = pending.pop() {
331        match v {
332            Value::Analytic(_) | Value::Event(_) => return true,
333            Value::List(v) => pending.extend(v.iter()),
334            Value::Map(v) => pending.extend(v.iter().flat_map(|(k, v)| [k, v])),
335            Value::Bag(v) => pending.extend(v.keys()),
336            Value::Record(v) => pending.extend(v.fields.iter().map(|(_, v)| v)),
337            Value::Dist(d) => pending.extend(d.outcomes.iter().map(|(v, _)| v)),
338            Value::Closure(c) => pending.extend(c.captured.iter()),
339            _ => {}
340        }
341    }
342    false
343}
344
345/// Latents reachable through an immutable value (including closure captures).
346pub fn collect_ids(v: &Value, ids: &mut std::collections::BTreeSet<u64>) {
347    let mut pending = vec![v];
348    while let Some(v) = pending.pop() {
349        match v {
350            Value::Analytic(a) => {
351                ids.insert(a.id);
352            }
353            Value::Event(e) => {
354                ids.insert(e.draw.id);
355            }
356            Value::List(v) => pending.extend(v.iter()),
357            Value::Map(v) => pending.extend(v.iter().flat_map(|(k, v)| [k, v])),
358            Value::Bag(v) => pending.extend(v.keys()),
359            Value::Record(v) => pending.extend(v.fields.iter().map(|(_, v)| v)),
360            Value::Dist(d) => pending.extend(d.outcomes.iter().map(|(v, _)| v)),
361            Value::Closure(c) => pending.extend(c.captured.iter()),
362            _ => {}
363        }
364    }
365}
366
367/// Snapshot a value under this world's posterior without mutating aliases.
368pub fn resolve(v: &Value, context: &Constraints, budget: &mut Budget) -> OpResult<Value> {
369    if context.is_empty() || !contains(v) {
370        return Ok(v.clone());
371    }
372    resolve_at(v, context, budget, 0)
373}
374fn resolve_at(v: &Value, c: &Constraints, b: &mut Budget, depth: usize) -> OpResult<Value> {
375    b.work(1)?;
376    if depth > 64 {
377        return Err(OpError::limit("analytic value nesting exceeds the limit of 64"));
378    }
379    let draw = |x: &Analytic| {
380        let mut x = x.clone();
381        if let Some(d) = c.get(&x.id) {
382            x.domain = x.domain.intersect(d);
383        }
384        x
385    };
386    let mut child = |v: &Value| resolve_at(v, c, b, depth + 1);
387    Ok(match v {
388        Value::Analytic(x) => draw(x).value()?,
389        Value::Event(e) => Event {
390            draw: draw(&e.draw),
391            yes: e.yes.clone(),
392        }
393        .value(),
394        Value::List(xs) => Value::list(xs.iter().map(&mut child).collect::<OpResult<_>>()?),
395        Value::Record(r) => crate::ops::make_record(
396            r.ty.clone(),
397            r.fields
398                .iter()
399                .map(|(k, v)| Ok((k.clone(), child(v)?)))
400                .collect::<OpResult<_>>()?,
401        ),
402        Value::Map(xs) => Value::map(
403            xs.iter()
404                .map(|(k, v)| Ok((k.clone(), child(v)?)))
405                .collect::<OpResult<_>>()?,
406        ),
407        Value::Closure(f) => Value::Closure(Arc::new(Closure {
408            func: f.func,
409            captured: f.captured.iter().map(&mut child).collect::<OpResult<_>>()?,
410        })),
411        Value::Dist(d) => {
412            let pairs = d
413                .outcomes
414                .iter()
415                .map(|(v, p)| Ok((child(v)?, *p)))
416                .collect::<OpResult<_>>()?;
417            crate::ops::combine(pairs, d.missing, b)?
418        }
419        // Analytic keys and bags are rejected when constructed.
420        _ => v.clone(),
421    })
422}