1use 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#[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
321pub 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
345pub 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
367pub 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 _ => v.clone(),
421 })
422}