1use crate::analytic::{self, Analytic};
10use crate::builtins;
11use crate::chain::{Chain, Solution};
12use crate::conjugate::{self, Seen};
13use crate::continuous::{Family, Rng};
14use crate::dist::{Budget, Counts, Dist};
15use crate::error::{OpError, OpResult, Result, RuntimeError};
16use crate::ops::{self, Truth};
17use crate::report::Sink;
18use crate::value::{Closure, Delayed, Value, fmt_prob};
19use crate::weight::Weight;
20use crate::world::Returned;
21use crate::world::{Flow, World, clear, clear_dead, live_slots, merge, merge_values, state_hash, total_weight};
22use probl_sema::builtins::Lifting;
23use probl_sema::conjugate::{Conjugacy, Likelihood, Update};
24use probl_sema::ir::*;
25use probl_sema::{Builtin, Liveness};
26use probl_syntax::Span;
27use probl_syntax::ast::BinOp;
28use rustc_hash::{FxHashMap, FxHashSet};
29use std::collections::BTreeMap;
30use std::sync::Arc;
31use std::sync::atomic::{AtomicBool, Ordering};
32
33#[derive(Clone, Debug, Default)]
34pub struct Stats {
35 pub peak_worlds: usize,
37 pub world_steps: u64,
39 pub calls: u64,
40 pub memo_hits: u64,
41 pub solved_loops: u64,
43 pub chain_states: u64,
44 pub solved_calls: u64,
47 pub call_rounds: u64,
48 pub updates: Vec<Updates>,
51}
52
53#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
56pub struct Updates {
57 pub delayed: u64,
59 pub exact: u64,
61 pub drawn: u64,
63}
64
65impl Stats {
66 pub fn absorb(&mut self, other: &Stats) {
68 self.peak_worlds = self.peak_worlds.max(other.peak_worlds);
69 self.world_steps += other.world_steps;
70 self.calls += other.calls;
71 self.memo_hits += other.memo_hits;
72 self.solved_loops += other.solved_loops;
73 self.chain_states += other.chain_states;
74 self.solved_calls += other.solved_calls;
75 self.call_rounds += other.call_rounds;
76 if self.updates.len() < other.updates.len() {
77 self.updates.resize(other.updates.len(), Updates::default());
78 }
79 for (mine, theirs) in self.updates.iter_mut().zip(&other.updates) {
80 mine.delayed += theirs.delayed;
81 mine.exact += theirs.exact;
82 mine.drawn += theirs.drawn;
83 }
84 }
85}
86
87#[derive(Clone, Debug)]
89pub struct Config {
90 pub today: Option<i32>,
91 pub epsilon: f64,
92 pub merging: bool,
93 pub memoizing: bool,
94 pub max_worlds: usize,
95 pub max_iterations: u64,
96 pub max_call_depth: usize,
97 pub max_cached_calls: usize,
98 pub max_output: usize,
99 pub budget: Budget,
100 pub cancel: Option<Arc<AtomicBool>>,
101 pub sample_seed: Option<u64>,
104 pub conjugate: bool,
107 pub solving: bool,
109 pub max_chain_states: usize,
111}
112
113const MAX_ELIMINATION: u64 = 200_000_000;
116
117pub const BATCH: u64 = 1000;
121
122#[derive(Debug)]
125pub struct Batch {
126 pub sinks: Vec<Sink>,
127 pub totals: SampleTotals,
128 pub observed: bool,
129 pub densities: bool,
130 pub unresolved: Weight,
131 pub last_ruling_out: Option<Span>,
132 pub stats: Stats,
133}
134
135pub type Printed = Vec<(Span, String)>;
137
138#[derive(Clone, Copy, Debug)]
140pub struct SampleTotals {
141 pub weight: Weight,
142 pub squares: Weight,
144}
145
146#[derive(Clone, Debug)]
148pub struct CallResult {
149 pub outcomes: Vec<Returned>,
150 pub unresolved: Weight,
152 pub lost: Weight,
154 pub observed: bool,
156 pub pending: Vec<(usize, Weight)>,
160}
161
162impl CallResult {
163 fn waiting(depth: usize) -> CallResult {
166 CallResult {
167 outcomes: Vec::new(),
168 unresolved: Weight::ONE,
169 lost: Weight::ZERO,
170 observed: false,
171 pending: vec![(depth, Weight::ONE)],
172 }
173 }
174
175 fn take_pending(&mut self, depth: usize) -> Weight {
177 match self.pending.iter().position(|&(d, _)| d == depth) {
178 Some(i) => self.pending.remove(i).1,
179 None => Weight::ZERO,
180 }
181 }
182}
183
184struct Approx {
187 result: Arc<CallResult>,
188 head: usize,
192 frame: u64,
193 round: u64,
194}
195
196fn add_pending(into: &mut Vec<(usize, Weight)>, from: &[(usize, Weight)], weight: Weight) {
198 for &(depth, w) in from {
199 match into.iter_mut().find(|(d, _)| *d == depth) {
200 Some((_, total)) => *total += weight * w,
201 None => into.push((depth, weight * w)),
202 }
203 }
204}
205
206type CallKey = (FnId, Vec<Value>);
207
208pub struct Engine<'p> {
209 prog: &'p Program,
210 live: &'p Liveness,
211 conj: &'p Conjugacy,
213 config: Config,
214 budget: Budget,
215 memo: FxHashMap<CallKey, Arc<CallResult>>,
216 active: FxHashMap<CallKey, usize>,
218 approx: FxHashMap<CallKey, Approx>,
221 frames_at: Vec<u64>,
224 rounds_at: Vec<u64>,
225 last_number: u64,
226 pending: Vec<(usize, Weight)>,
228 depth: usize,
229 pub unresolved: Weight,
231 lost: Weight,
234 solvable: Vec<bool>,
237 too_large: FxHashSet<StmtId>,
240 pub observed: bool,
242 pub densities: bool,
245 pub last_ruling_out: Option<Span>,
247 pub sinks: Vec<Sink>,
248 dice: FxHashMap<(u32, u32), Value>,
249 pools: FxHashMap<(u32, Value), Value>,
250 record_names: Vec<Arc<str>>,
251 enum_values: Vec<Vec<Value>>,
252 print: &'p mut (dyn FnMut(&str) + Send),
253 inputs: &'p [Value],
255 printed: usize,
257 lines: Option<Printed>,
260 pub stats: Stats,
261 sampler: Option<Rng>,
264 nested: usize,
266 next_latent: u64,
267 callback: Option<&'static str>,
270}
271
272impl<'p> Engine<'p> {
273 pub fn new(
274 prog: &'p Program,
275 live: &'p Liveness,
276 conj: &'p Conjugacy,
277 config: Config,
278 inputs: &'p [Value],
279 print: &'p mut (dyn FnMut(&str) + Send),
280 ) -> Engine<'p> {
281 let sampler = config.sample_seed.map(Rng::new);
282 Engine {
283 prog,
284 live,
285 conj,
286 budget: config.budget.clone(),
287 config,
288 memo: FxHashMap::default(),
289 active: FxHashMap::default(),
290 approx: FxHashMap::default(),
291 frames_at: Vec::new(),
292 rounds_at: Vec::new(),
293 last_number: 0,
294 pending: Vec::new(),
295 depth: 0,
296 unresolved: Weight::ZERO,
297 lost: Weight::ZERO,
298 solvable: probl_sema::effects::solvable_loops(prog),
299 too_large: FxHashSet::default(),
300 observed: false,
301 densities: false,
302 last_ruling_out: None,
303 sinks: vec![Sink::default(); prog.reports.len()],
304 dice: FxHashMap::default(),
305 pools: FxHashMap::default(),
306 record_names: prog.records.iter().map(|r| Arc::from(r.name.as_str())).collect(),
307 enum_values: prog
308 .enums
309 .iter()
310 .enumerate()
311 .map(|(t, e)| {
312 e.variants
313 .iter()
314 .enumerate()
315 .map(|(v, name)| ops::enum_value(t as u32, v as u32, name))
316 .collect()
317 })
318 .collect(),
319 print,
320 inputs,
321 printed: 0,
322 lines: None,
323 stats: Stats::default(),
324 sampler,
325 nested: 0,
326 next_latent: 0,
327 callback: None,
328 }
329 }
330
331 pub fn run_main(&mut self) -> Result<Weight> {
333 let prog = self.prog;
334 let main = prog.main();
335 let world = World {
336 slots: vec![Value::Dead; main.n_slots()],
337 constraints: Default::default(),
338 inherited: Default::default(),
339 weight: Weight::ONE,
340 run: 0,
341 };
342 let flow = self.exec_block(MAIN, &main.body, vec![world])?;
343 Ok(total_weight(&flow.next))
344 }
345
346 pub fn run_batch(&mut self, rng: Rng, first: u64, n: u64, printed: usize) -> (Result<Batch>, Printed) {
352 self.sampler = Some(rng);
353 self.sinks = vec![Sink::default(); self.prog.reports.len()];
354 self.observed = false;
355 self.densities = false;
356 self.unresolved = Weight::ZERO;
357 self.last_ruling_out = None;
358 self.stats = Stats::default();
359 self.printed = printed;
360 self.lines = Some(Vec::new());
361 self.budget = self.config.budget.clone();
362 let result = self.sample_runs(first, n);
363 self.budget.give_back();
364 (result, self.lines.take().unwrap_or_default())
365 }
366
367 fn sample_runs(&mut self, first: u64, n: u64) -> Result<Batch> {
368 let main = self.prog.main();
369 let worlds = (first..first + n)
370 .map(|run| World {
371 slots: vec![Value::Dead; main.n_slots()],
372 constraints: Default::default(),
373 inherited: Default::default(),
374 weight: Weight::ONE,
375 run: run as u32,
376 })
377 .collect();
378 let flow = self.exec_block(MAIN, &main.body, worlds)?;
379 let mut totals = SampleTotals {
380 weight: Weight::ZERO,
381 squares: Weight::ZERO,
382 };
383 for w in &flow.next {
384 totals.weight += w.weight;
385 totals.squares += w.weight * w.weight;
386 }
387 for sink in &mut self.sinks {
388 sink.end_batch();
389 }
390 Ok(Batch {
391 sinks: std::mem::take(&mut self.sinks),
392 totals,
393 observed: self.observed,
394 densities: self.densities,
395 unresolved: self.unresolved,
396 last_ruling_out: self.last_ruling_out,
397 stats: std::mem::take(&mut self.stats),
398 })
399 }
400
401 fn sample(&mut self, v: &Value) -> Value {
405 let rng = self.sampler.as_mut().expect("only called when sampling");
406 match v {
407 Value::Dist(d) if !d.outcomes.is_empty() => {
408 let i = rng.choose(d.outcomes.iter().map(|(_, p)| *p)).unwrap_or(0);
409 match &d.outcomes[i].0 {
410 Value::Continuous(f) => Value::Float(f.sample(rng)),
411 x => x.clone(),
412 }
413 }
414 Value::Continuous(f) => Value::Float(f.sample(rng)),
415 other => other.clone(),
416 }
417 }
418
419 fn sample_continuous(&mut self, v: Value) -> Value {
423 match &v {
424 Value::Continuous(_) => self.sample(&v),
425 Value::Dist(d) if d.outcomes.iter().any(|(x, _)| matches!(x, Value::Continuous(_))) => {
426 let pairs = d.outcomes.iter().map(|(x, p)| (self.sample(x), *p)).collect();
427 Dist::from_pairs(pairs, d.missing).into_value()
428 }
429 _ => v,
430 }
431 }
432
433 fn direct_counts(&mut self, f: FnId, e: &'p Expr, w: &World) -> Result<Option<Counts>> {
437 let ExprKind::Builtin {
438 func: b @ (Builtin::Binomial | Builtin::Poisson | Builtin::Geometric),
439 args,
440 ..
441 } = &e.kind
442 else {
443 return Ok(None);
444 };
445 if self.sampler.is_none() {
446 return Ok(None);
447 }
448 let mut values = Vec::with_capacity(args.len());
449 for a in args {
450 values.push(self.eval(f, a, w)?);
451 }
452 let counts = builtins::counts(*b, &values, &mut self.budget).map_err(|err| err.at(e.span))?;
453 Ok(counts.filter(Counts::direct))
454 }
455
456 fn continuous_draw(&self, span: Span) -> RuntimeError {
457 analytic::unsupported("a continuous draw inside `simulate`")
458 .at(span)
459 .with_note("`simulate` is computed by enumeration, even in sample mode")
460 .with_help("draw the value outside `simulate`; analytic joint distribution recipes aren't supported yet")
461 }
462
463 fn delaying(&self) -> bool {
466 self.config.conjugate && self.sampler.is_some()
467 }
468
469 fn updates(&mut self, variable: u32) -> &mut Updates {
470 let i = variable as usize;
471 if self.stats.updates.len() <= i {
472 let n = self.conj.variables.len().max(i + 1);
473 self.stats.updates.resize(n, Updates::default());
474 }
475 &mut self.stats.updates[i]
476 }
477
478 fn draw_delayed(&mut self, w: &mut World, slot: SlotId) {
481 if let Value::Delayed(d) = &w.slots[slot as usize] {
482 let d = **d;
483 let x = d
484 .family
485 .sample(self.sampler.as_mut().expect("delayed only when sampling"));
486 w.slots[slot as usize] = Value::Float(x);
487 self.updates(d.variable).drawn += 1;
488 }
489 }
490
491 fn observe_exactly(&mut self, f: FnId, u: &Update<'p>, w: &mut World, from: Span) -> Result<Option<f64>> {
500 let Value::Delayed(d) = &w.slots[u.slot as usize] else {
501 return Ok(None);
502 };
503 let d = **d;
504 let pair = matches!(
505 (d.family, u.likelihood),
506 (
507 Family::Beta { .. },
508 Likelihood::Binomial { .. } | Likelihood::Bernoulli { .. }
509 ) | (Family::Gamma { .. }, Likelihood::Poisson { .. })
510 | (Family::Normal { .. }, Likelihood::Normal { .. })
511 );
512 let seen = if pair { self.seen(f, u, w, from)? } else { None };
513 let Some(seen) = seen else {
514 self.draw_delayed(w, u.slot);
515 return Ok(None);
516 };
517 let (ln, posterior) = conjugate::update(&d.family, seen).expect("a conjugate pair");
518 w.slots[u.slot as usize] = Value::Delayed(Arc::new(Delayed {
519 family: posterior,
520 variable: d.variable,
521 }));
522 self.updates(d.variable).exact += 1;
523 self.densities |= matches!(seen, Seen::Normal { .. });
524 Ok(Some(ln))
525 }
526
527 fn seen(&mut self, f: FnId, u: &Update<'p>, w: &World, from: Span) -> Result<Option<Seen>> {
530 let count = |v: &Value| match v {
533 Value::Bool(_) => f64::NAN,
534 _ => v.as_f64().unwrap_or(f64::NAN),
535 };
536 Ok(match u.likelihood {
537 Likelihood::Binomial { value, trials } => {
538 let v = self.eval(f, value, w)?;
539 let n = self.eval(f, trials, w)?;
540 match builtins::counts(Builtin::Binomial, &[n, Value::Prob(0.5)], &mut self.budget)
541 .map_err(|e| e.at(from))?
542 {
543 Some(Counts::Binomial { n, .. }) => Some(Seen::Binomial {
544 trials: n,
545 k: count(&v),
546 }),
547 _ => None,
548 }
549 }
550 Likelihood::Bernoulli { value } => match self.eval(f, value, w)? {
551 Value::Bool(b) => Some(Seen::Bernoulli(b)),
552 _ => None,
553 },
554 Likelihood::Poisson { value } => {
555 let v = self.eval(f, value, w)?;
556 Some(Seen::Poisson { k: count(&v) })
557 }
558 Likelihood::Normal { value, sd } => {
559 let v = self.eval(f, value, w)?;
560 let sd = self.eval(f, sd, w)?;
561 if sd.is_uncertain() {
562 return Ok(None);
563 }
564 let checked = builtins::call_plain(Builtin::Normal, &[Value::Float(0.0), sd], &mut self.budget)
565 .map_err(|e| e.at(from))?;
566 let Value::Continuous(family) = checked else {
567 unreachable!("`normal` gives a continuous distribution")
568 };
569 let Family::Normal { sd, .. } = *family else {
570 unreachable!("`normal` gives a normal distribution")
571 };
572 match v {
573 Value::Int(_) | Value::Float(_) | Value::Prob(_) if v.as_f64().is_some() => Some(Seen::Normal {
574 y: v.as_f64().expect("a number"),
575 sd,
576 }),
577 _ => None,
578 }
579 }
580 })
581 }
582
583 fn merge(&self, worlds: Vec<World>, stmt: StmtId) -> Vec<World> {
584 merge(worlds, &self.live.after[stmt as usize], self.merging())
585 }
586
587 fn merging(&self) -> bool {
590 self.config.merging && self.sampler.is_none()
591 }
592
593 fn check_worlds(&self, n: usize, span: Span) -> Result<()> {
596 if n > self.config.max_worlds && self.sampler.is_none() {
597 return Err(RuntimeError::limit(
598 span,
599 format!("more than {} worlds are too many to follow", self.config.max_worlds),
600 )
601 .with_help("simplify the model, raise the limit, or sample it with `@mode sample(runs: 10_000)`"));
602 }
603 Ok(())
604 }
605
606 fn spend(&mut self, n: u64, span: Span) -> Result<()> {
607 self.budget.work(n).map_err(|e| e.at(span))?;
608 if let Some(cancel) = &self.config.cancel {
609 if cancel.load(Ordering::Relaxed) {
610 return Err(RuntimeError::limit(span, "the run was cancelled"));
611 }
612 }
613 Ok(())
614 }
615
616 fn exec_block(&mut self, f: FnId, block: &'p Block, worlds: Vec<World>) -> Result<Flow> {
619 let mut flow = Flow::next(worlds);
620 for stmt in &block.stmts {
621 if flow.next.is_empty() {
622 break;
623 }
624 let input = std::mem::take(&mut flow.next);
625 let out = self.exec_stmt(f, stmt, input)?;
626 flow.next = out.next;
627 clear(&mut flow.next, &self.live.dies[stmt.id as usize]);
628 flow.broke.extend(out.broke);
629 flow.continued.extend(out.continued);
630 flow.returned.extend(out.returned);
631 }
632 Ok(flow)
633 }
634
635 fn restrict_event(&mut self, event: &analytic::Event, yes: bool, w: &mut World, span: Span) -> Result<f64> {
636 self.budget
637 .collection((w.constraints.len() + 1) as u128)
638 .map_err(|e| e.at(span))?;
639 self.budget
640 .work((w.constraints.len() + event.yes.0.len() + event.draw.domain.0.len()) as u64)
641 .map_err(|e| e.at(span))?;
642 Ok(event.restrict(yes, &mut w.constraints))
643 }
644
645 fn exec_stmt(&mut self, f: FnId, stmt: &'p Stmt, worlds: Vec<World>) -> Result<Flow> {
646 let span = stmt.span;
647 if matches!(
648 stmt.kind,
649 StmtKind::Draw { .. } | StmtKind::Take { .. } | StmtKind::Observe { .. } | StmtKind::Chance { .. }
650 ) {
651 self.check_callback_effect(span)?;
652 }
653 let n = worlds.len();
654 self.stats.world_steps += n as u64;
655 self.stats.peak_worlds = self.stats.peak_worlds.max(n);
656 self.spend(n as u64, span)?;
657 self.check_worlds(n, span)?;
658 let mut worlds = worlds;
659 if self.delaying() {
660 let conj = self.conj;
662 let first = &conj.draws_first[stmt.id as usize];
663 if !first.is_empty() {
664 for w in &mut worlds {
665 for &slot in first {
666 self.draw_delayed(w, slot);
667 }
668 }
669 }
670 }
671 match &stmt.kind {
672 StmtKind::Set { place, value } => {
673 let mut worlds = worlds;
674 for w in &mut worlds {
675 let v = match self.record_place_type(place, w) {
676 Some(ty) => self.eval_expected(f, value, w, &ty)?,
677 None => self.eval(f, value, w)?,
678 };
679 self.assign(f, place, v, w, span)?;
680 }
681 if self.live.narrows[stmt.id as usize] && self.merging() {
686 return Ok(Flow::next(self.merge(worlds, stmt.id)));
687 }
688 Ok(Flow::next(worlds))
689 }
690 StmtKind::Draw { place, dist } => {
691 let delay = if self.delaying() {
693 self.conj.delays[stmt.id as usize]
694 } else {
695 None
696 };
697 let mut out = Vec::with_capacity(worlds.len());
698 for w in worlds {
699 if self.sampler.is_some() {
700 if let Some(counts) = self.direct_counts(f, dist, &w)? {
701 let k = counts.sample(self.sampler.as_mut().expect("sampling"));
702 let mut w = w;
703 self.assign(f, place, Value::Int(k.into()), &mut w, dist.span)?;
704 out.push(w);
705 continue;
706 }
707 }
708 let d = self.eval(f, dist, &w)?;
709 if let (Some(variable), Value::Continuous(family)) = (delay, &d) {
710 if conjugate::is_prior(family) {
711 let mut w = w;
712 w.slots[place.slot as usize] = Value::Delayed(Arc::new(Delayed {
713 family: **family,
714 variable,
715 }));
716 self.updates(variable).delayed += 1;
717 out.push(w);
718 continue;
719 }
720 }
721 self.split_by(f, place, d, w, dist.span, &mut out)?;
722 self.check_worlds(out.len(), span)?;
723 }
724 Ok(Flow::next(self.merge(out, stmt.id)))
725 }
726 StmtKind::Take { place, bag } => {
727 let mut out = Vec::new();
728 for w in worlds {
729 let current = self.read_place(f, bag, &w, span)?;
730 let Value::Bag(cards) = ¤t else {
731 return Err(RuntimeError::new(
732 span,
733 format!("`take` needs a bag, found {}", ops::article(¤t.kind())),
734 )
735 .with_help("make one with `bag([card: count, …])`"));
736 };
737 let total: u128 = cards.values().map(|n| *n as u128).sum();
738 if total == 0 {
739 return Err(RuntimeError::new(span, "can't take a card from an empty bag"));
740 }
741 let chosen = match &mut self.sampler {
742 Some(rng) => rng.choose(cards.values().map(|n| *n as f64)),
743 None => None,
744 };
745 for (i, (card, count)) in cards.iter().enumerate() {
746 if self.sampler.is_some() && chosen != Some(i) {
747 continue;
748 }
749 let rest = cards.without_nth(i);
750 let mut nw = if self.sampler.is_some() {
751 w.clone()
752 } else {
753 w.clone().scaled(*count as f64 / total as f64)
754 };
755 self.assign(f, bag, Value::multiset(rest), &mut nw, span)?;
756 self.assign(f, place, card.clone(), &mut nw, span)?;
757 out.push(nw);
758 }
759 self.check_worlds(out.len(), span)?;
760 }
761 Ok(Flow::next(self.merge(out, stmt.id)))
762 }
763 StmtKind::Call { dest, callee, args } => {
764 let mut out = Vec::with_capacity(worlds.len());
765 for w in worlds {
766 let mut key = Vec::with_capacity(args.len() + 4);
767 for a in args {
768 key.push(self.eval(f, a, &w)?);
769 }
770 let func = match callee {
771 Callee::Fn { func, capture_args } => {
772 for &s in capture_args {
773 key.push(self.slot(f, s, &w, span)?);
774 }
775 *func
776 }
777 Callee::Value(e, named) => match self.eval(f, e, &w)? {
778 Value::Builtin(b) => {
779 let value = self.builtin_values(b, &key, named, w.weight, span)?;
780 let mut nw = w;
781 self.assign(f, dest, value, &mut nw, span)?;
782 out.push(nw);
783 continue;
784 }
785 Value::Closure(c) => {
786 if !named.is_empty() {
787 return Err(RuntimeError::new(
788 span,
789 "only builtin minimum/maximum accept named defaults",
790 ));
791 }
792 self.check_arity(&c, args.len(), span)?;
793 key.extend(c.captured.iter().cloned());
794 c.func
795 }
796 other => {
797 return Err(RuntimeError::new(
798 e.span,
799 format!("can't call {}", ops::article(&other.kind())),
800 ));
801 }
802 },
803 };
804 let result = self.call(func, key, span)?;
805 self.unresolved += w.weight * result.unresolved;
806 self.lost += w.weight * result.lost;
807 add_pending(&mut self.pending, &result.pending, w.weight);
808 let last = result.outcomes.len().saturating_sub(1);
809 let mut w = Some(w);
810 for (i, (v, p, restrictions)) in result.outcomes.iter().enumerate() {
811 let mut nw = if i == last {
812 w.take().unwrap()
813 } else {
814 w.clone().unwrap()
815 };
816 nw.weight = nw.weight * *p;
817 if !restrictions.is_empty() {
820 Arc::make_mut(&mut nw.constraints)
821 .extend(restrictions.iter().map(|(id, d)| (*id, d.clone())));
822 }
823 self.assign(f, dest, v.clone(), &mut nw, span)?;
824 out.push(nw);
825 }
826 self.check_worlds(out.len(), span)?;
827 }
828 Ok(Flow::next(self.merge(out, stmt.id)))
829 }
830 StmtKind::If { cond, then, otherwise } => {
831 let (mut yes, mut no) = (Vec::new(), Vec::new());
832 for w in worlds {
833 let condition = self.eval(f, cond, &w)?;
834 if let Value::Event(event) = condition {
835 self.check_callback_effect(cond.span)?;
836 let mut y = w.clone();
837 let mut n = w;
838 let p = self.restrict_event(&event, true, &mut y, cond.span)?;
839 let q = self.restrict_event(&event, false, &mut n, cond.span)?;
840 if p > 0.0 {
841 yes.push(y.scaled(p));
842 }
843 if q > 0.0 {
844 no.push(n.scaled(q));
845 }
846 continue;
847 }
848 let c = ops::condition(&condition).map_err(|e| e.at(cond.span))?;
849 if (c.yes > 0.0 && c.no > 0.0) || c.missing > 0.0 {
850 self.check_callback_effect(cond.span)?;
851 }
852 if let Some(rng) = &mut self.sampler {
853 match rng.choose([c.yes, c.no].into_iter()) {
855 Some(0) => yes.push(w),
856 Some(_) => no.push(w),
857 None => self.unresolved += w.weight,
858 }
859 continue;
860 }
861 let (p, q) = (c.yes, c.no);
862 if c.missing > 0.0 {
863 self.unresolved += w.weight.scale(c.missing);
864 }
865 if q <= 0.0 {
866 yes.push(w.scaled(p));
867 } else if p <= 0.0 {
868 no.push(w.scaled(q));
869 } else {
870 no.push(w.clone().scaled(q));
871 yes.push(w.scaled(p));
872 }
873 }
874 self.check_worlds(yes.len() + no.len(), span)?;
875 let mut flow = Flow::default();
876 if !yes.is_empty() {
877 flow.join(self.exec_block(f, then, yes)?);
878 }
879 if !no.is_empty() {
880 flow.join(self.exec_block(f, otherwise, no)?);
881 }
882 flow.next = self.merge(flow.next, stmt.id);
883 Ok(flow)
884 }
885 StmtKind::Chance {
886 arms,
887 otherwise,
888 exhaustive,
889 } => {
890 let mut buckets: Vec<Vec<World>> = vec![Vec::new(); arms.len()];
891 let mut rest = Vec::new();
892 for w in worlds {
893 let mut chances = Vec::with_capacity(arms.len());
895 for (weight, _) in arms {
896 let value = self.eval(f, weight, &w)?;
897 chances.push(ops::to_prob(&value).map_err(|e| e.at(weight.span))?);
898 }
899 let sum: f64 = chances.iter().sum();
900 if sum > 1.0 + 1e-9 {
901 return Err(RuntimeError::new(
902 span,
903 format!(
904 "the chances in this `chance` add up to {}, more than 100%",
905 fmt_prob(sum)
906 ),
907 ));
908 }
909 let remainder = (1.0 - sum).max(0.0);
910 if remainder > 1e-9 && *exhaustive && otherwise.is_none() {
911 return Err(RuntimeError::new(
912 span,
913 format!("the chances add up to {} and there's no `else` arm", fmt_prob(sum)),
914 )
915 .with_help("a `chance` used as a value needs its chances to add up to 100%, or an `else`"));
916 }
917 if let Some(rng) = &mut self.sampler {
918 let options = chances.iter().copied().chain(std::iter::once(remainder));
920 match rng.choose(options) {
921 Some(i) if i < arms.len() => buckets[i].push(w),
922 Some(_) => rest.push(w),
923 None => self.unresolved += w.weight,
924 }
925 continue;
926 }
927 for (bucket, p) in buckets.iter_mut().zip(&chances) {
928 if *p > 0.0 {
929 bucket.push(w.clone().scaled(*p));
930 }
931 }
932 if remainder > 1e-12 {
933 rest.push(w.scaled(remainder));
934 }
935 }
936 self.check_worlds(buckets.iter().map(Vec::len).sum::<usize>() + rest.len(), span)?;
937 let mut flow = Flow::default();
938 for ((_, body), bucket) in arms.iter().zip(buckets) {
939 if !bucket.is_empty() {
940 flow.join(self.exec_block(f, body, bucket)?);
941 }
942 }
943 if !rest.is_empty() {
944 match otherwise {
945 Some(body) => flow.join(self.exec_block(f, body, rest)?),
946 None if *exhaustive => {}
947 None => flow.next.extend(rest),
948 }
949 }
950 flow.next = self.merge(flow.next, stmt.id);
951 Ok(flow)
952 }
953 StmtKind::Loop { body, bounded } => self.exec_loop(f, stmt, body, *bounded, worlds),
954 StmtKind::Break => Ok(Flow {
955 broke: worlds,
956 ..Flow::default()
957 }),
958 StmtKind::Continue => Ok(Flow {
959 continued: worlds,
960 ..Flow::default()
961 }),
962 StmtKind::Return(value) => {
963 let mut flow = Flow::default();
964 for w in worlds {
965 let v = self.eval(f, value, &w)?;
966 let mut constraints = w.constraints;
967 if constraints.keys().any(|id| !w.inherited.contains(id)) {
968 Arc::make_mut(&mut constraints).retain(|id, _| w.inherited.contains(id));
969 }
970 flow.returned.push((v, w.weight, constraints));
971 }
972 Ok(flow)
973 }
974 StmtKind::Observe { value, from } => {
975 self.observed = true;
976 let exact = if self.delaying() {
977 probl_sema::conjugate::update(value, from.as_ref())
978 } else {
979 None
980 };
981 let mut out = Vec::with_capacity(worlds.len());
982 for mut w in worlds {
983 if let Some(u) = &exact {
984 let at = from.as_ref().map_or(span, |d| d.span);
985 if let Some(ln) = self.observe_exactly(f, u, &mut w, at)? {
986 w.weight = w.weight * Weight::from_ln(ln);
987 if w.weight.is_zero() {
988 self.last_ruling_out = Some(span);
989 } else {
990 out.push(w);
991 }
992 continue;
993 }
994 }
995 let (factor, missing, ruled_out) = match from {
996 None => {
997 let v = self.eval(f, value, &w)?;
998 if let Value::Event(event) = v {
999 let p = self.restrict_event(&event, true, &mut w, value.span)?;
1000 (p, 0.0, 1.0 - p)
1001 } else {
1002 let b = ops::fact(&v, "observe").map_err(|e| e.at(value.span))?;
1003 if b { (1.0, 0.0, 0.0) } else { (0.0, 0.0, 1.0) }
1004 }
1005 }
1006 Some(d) => {
1007 let v = self.eval(f, value, &w)?;
1008 match self.direct_counts(f, d, &w)? {
1009 Some(counts) if self.sampler.is_some() => {
1010 let x = match v {
1011 Value::Bool(_) => None,
1012 _ => v.as_f64(),
1013 };
1014 (x.map_or(0.0, |x| counts.pmf(x)), 0.0, 0.0)
1015 }
1016 _ => {
1017 let dist = self.eval(f, d, &w)?;
1018 if let (Value::Bool(b), Value::Event(event)) = (&v, &dist) {
1019 let p = self.restrict_event(event, *b, &mut w, span)?;
1020 (p, 0.0, 1.0 - p)
1021 } else {
1022 if analytic::contains(&v) || analytic::contains(&dist) {
1023 return Err(analytic::unsupported("this likelihood observation").at(span));
1024 }
1025 self.densities |= is_density(&dist);
1026 likelihood(&dist, &v, self.sampler.is_some()).map_err(|e| e.at(span))?
1027 }
1028 }
1029 }
1030 }
1031 };
1032 if missing > 0.0 && self.sampler.is_none() {
1033 self.unresolved += w.weight.scale(missing);
1034 }
1035 if ruled_out > 0.0 && self.sampler.is_none() {
1036 self.lost += w.weight.scale(ruled_out);
1037 }
1038 if factor > 0.0 {
1039 w.weight = w.weight.scale(factor);
1040 out.push(w);
1041 } else {
1042 self.last_ruling_out = Some(span);
1043 }
1044 }
1045 Ok(Flow::next(out))
1046 }
1047 StmtKind::Report { site, value, key } => {
1048 for w in &worlds {
1049 let v = self.eval(f, value, w)?;
1050 let k = match key {
1051 Some(k) => {
1052 let k_value = self.eval(f, k, w)?;
1053 if analytic::contains(&k_value) {
1054 return Err(analytic::unsupported("grouping by a continuous outcome").at(k.span));
1055 }
1056 if k_value.is_uncertain() {
1057 return Err(RuntimeError::new(
1058 k.span,
1059 "a `by` key must be a plain value, not a distribution",
1060 )
1061 .with_help("draw a value first with `~`"));
1062 }
1063 k_value
1064 }
1065 None => Value::Unit,
1066 };
1067 let (v, run) = if self.sampler.is_some() {
1068 (self.sample_continuous(v), Some(w.run))
1069 } else {
1070 if analytic::contains(&v) && !matches!(v, Value::Analytic(_) | Value::Event(_)) {
1071 return Err(analytic::unsupported(
1072 "reporting an aggregate containing analytic outcomes; report its fields individually",
1073 )
1074 .at(value.span));
1075 }
1076 (v, None)
1077 };
1078 self.sinks[*site as usize]
1079 .validate_analytic(&k, &v)
1080 .map_err(|e| e.at(value.span))?;
1081 self.sinks[*site as usize].add(k, &v, w.weight, run);
1082 }
1083 Ok(Flow::next(worlds))
1084 }
1085 StmtKind::Fail { message } => Err(RuntimeError::new(span, message.clone())),
1086 StmtKind::Check { slot, ty } => {
1087 for w in &mut worlds {
1088 if let Value::Delayed(_) = &w.slots[*slot as usize] {
1091 let always = *ty == TypeSpec::Float;
1092 if always {
1093 continue;
1094 }
1095 self.draw_delayed(w, *slot);
1096 }
1097 w.slots[*slot as usize] = self.coerce(w.slots[*slot as usize].clone(), ty, span)?;
1098 let v = &w.slots[*slot as usize];
1099 if !self.conforms(v, ty) {
1100 let name = &self.prog.functions[f as usize].slots[*slot as usize].name;
1101 let what = if name == TEMP {
1102 "the result".to_string()
1103 } else {
1104 format!("`{name}`")
1105 };
1106 return Err(RuntimeError::new(
1107 span,
1108 format!(
1109 "{what} should be {}, but it's {}",
1110 ops::article(&ty.describe(self.prog)),
1111 ops::article(&v.kind())
1112 ),
1113 ));
1114 }
1115 }
1116 Ok(Flow::next(worlds))
1117 }
1118 }
1119 }
1120
1121 fn exec_loop(
1122 &mut self,
1123 f: FnId,
1124 stmt: &'p Stmt,
1125 body: &'p Block,
1126 bounded: bool,
1127 worlds: Vec<World>,
1128 ) -> Result<Flow> {
1129 let live = self.live;
1130 let span = stmt.span;
1131 let entered = total_weight(&worlds);
1132 let cutoff = entered.scale(self.config.epsilon);
1133 let mut inside = worlds;
1134 let mut out = Flow::default();
1135 let mut iterations: u64 = 0;
1136 let mut met = self.may_solve(stmt, bounded).then(FxHashSet::<u64>::default);
1140 while !inside.is_empty() {
1141 let cycles = match &mut met {
1142 Some(met) => {
1143 let head = live_slots(&live.loop_head[stmt.id as usize], inside[0].slots.len());
1144 let hashes: Vec<u64> = inside.iter().map(|w| state_hash(w, &head)).collect();
1145 let again = hashes.iter().any(|h| met.contains(h));
1146 met.extend(hashes);
1147 again
1148 }
1149 None => false,
1150 };
1151 if cycles {
1152 met = None;
1153 if let Some(solved) = self.solve_loop(f, stmt, body, &inside)? {
1154 out.join(solved);
1155 break;
1156 }
1157 }
1158 let mass = total_weight(&inside);
1159 if !bounded && mass < cutoff && self.sampler.is_none() {
1160 self.unresolved += mass;
1161 break;
1162 }
1163 iterations += 1;
1164 if iterations > self.config.max_iterations {
1165 return Err(RuntimeError::limit(
1166 span,
1167 format!("this loop ran {} times without finishing", self.config.max_iterations),
1168 )
1169 .with_note(format!(
1170 "worlds still inside the loop weigh {} of what entered it",
1171 fmt_prob(mass.ratio(entered))
1172 ))
1173 .with_help("check that every world can leave the loop"));
1174 }
1175 let mut flow = self.exec_block(f, body, inside)?;
1176 clear_dead(&mut flow.broke, &live.after[stmt.id as usize]);
1177 clear_dead(&mut flow.continued, &live.loop_head[stmt.id as usize]);
1178 out.next.extend(flow.broke);
1179 out.returned.extend(flow.returned);
1180 let mut again = flow.next;
1181 again.extend(flow.continued);
1182 inside = merge(again, &live.loop_head[stmt.id as usize], self.merging());
1183 }
1184 out.next = self.merge(out.next, stmt.id);
1185 Ok(out)
1186 }
1187
1188 fn may_solve(&self, stmt: &Stmt, bounded: bool) -> bool {
1192 !bounded
1193 && self.config.solving
1194 && self.sampler.is_none()
1195 && self.solvable[stmt.id as usize]
1196 && !self.too_large.contains(&stmt.id)
1197 }
1198
1199 fn solve_loop(&mut self, f: FnId, stmt: &'p Stmt, body: &'p Block, inside: &[World]) -> Result<Option<Flow>> {
1207 if inside
1208 .iter()
1209 .any(|w| !w.constraints.is_empty() || w.slots.iter().any(analytic::contains))
1210 {
1211 self.too_large.insert(stmt.id);
1212 return Ok(None);
1213 }
1214 let head_set = &self.live.loop_head[stmt.id as usize];
1215 let after = &self.live.after[stmt.id as usize];
1216 let n = inside[0].slots.len();
1217 let head = live_slots(head_set, n);
1218 let dead: Vec<SlotId> = head_set.iter_missing(n).collect();
1219 let mut states: Vec<Vec<Value>> = Vec::new();
1222 let mut index: FxHashMap<Vec<Value>, usize> = FxHashMap::default();
1223 let mut intern = |mut slots: Vec<Value>, states: &mut Vec<Vec<Value>>| -> usize {
1224 for &d in &dead {
1225 slots[d as usize] = Value::Dead;
1226 }
1227 let key: Vec<Value> = head.iter().map(|&i| slots[i].clone()).collect();
1228 *index.entry(key).or_insert_with(|| {
1229 states.push(slots);
1230 states.len() - 1
1231 })
1232 };
1233 let total = total_weight(inside);
1234 let mut start: Vec<f64> = Vec::new();
1235 for w in inside {
1236 let i = intern(w.slots.clone(), &mut states);
1237 if start.len() <= i {
1238 start.resize(i + 1, 0.0);
1239 }
1240 start[i] += w.weight.ratio(total);
1241 }
1242 let mut chain = Chain::default();
1243 let mut exits: Vec<Vec<World>> = Vec::new();
1244 let mut returns: Vec<Vec<Returned>> = Vec::new();
1245 let mut unresolved: Vec<Weight> = Vec::new();
1246 let saved_pending = std::mem::take(&mut self.pending);
1249 let mut k = 0;
1250 while k < states.len() {
1251 if !self.pending.is_empty() {
1252 self.pending = saved_pending;
1253 return Ok(None);
1254 }
1255 if states.len() > self.config.max_chain_states {
1256 self.pending = saved_pending;
1257 self.too_large.insert(stmt.id);
1258 return Ok(None);
1259 }
1260 let world = World {
1261 slots: states[k].clone(),
1262 constraints: Default::default(),
1263 inherited: Default::default(),
1264 weight: Weight::ONE,
1265 run: 0,
1266 };
1267 let saved_unresolved = std::mem::replace(&mut self.unresolved, Weight::ZERO);
1268 let saved_lost = std::mem::replace(&mut self.lost, Weight::ZERO);
1269 let flow = self.exec_block(f, body, vec![world]);
1270 let left = std::mem::replace(&mut self.unresolved, saved_unresolved);
1271 let lost = std::mem::replace(&mut self.lost, saved_lost);
1272 let mut flow = flow?;
1273 if flow
1274 .next
1275 .iter()
1276 .chain(&flow.continued)
1277 .chain(&flow.broke)
1278 .any(|w| !w.constraints.is_empty() || w.slots.iter().any(analytic::contains))
1279 || flow
1280 .returned
1281 .iter()
1282 .any(|(v, _, c)| !c.is_empty() || analytic::contains(v))
1283 {
1284 self.pending = saved_pending;
1285 self.too_large.insert(stmt.id);
1286 return Ok(None);
1287 }
1288 clear_dead(&mut flow.broke, after);
1289 let mut leave = left + lost;
1291 for w in &flow.broke {
1292 leave += w.weight;
1293 }
1294 for (_, w, _) in &flow.returned {
1295 leave += *w;
1296 }
1297 let mut next: FxHashMap<usize, f64> = FxHashMap::default();
1298 for w in flow.next.into_iter().chain(flow.continued) {
1299 let p = w.weight.to_f64();
1300 if p == 0.0 {
1301 self.too_large.insert(stmt.id);
1303 self.pending = saved_pending;
1304 return Ok(None);
1305 }
1306 let j = intern(w.slots, &mut states);
1307 *next.entry(j).or_insert(0.0) += p;
1308 }
1309 let mut next: Vec<(usize, f64)> = next.into_iter().collect();
1310 next.sort_by_key(|&(j, _)| j);
1311 chain.next.push(next);
1312 chain.leave.push(leave.to_f64());
1313 exits.push(flow.broke);
1314 returns.push(flow.returned);
1315 unresolved.push(left);
1316 k += 1;
1317 }
1318 let waiting = !self.pending.is_empty();
1319 self.pending = saved_pending;
1320 if waiting {
1321 return Ok(None);
1322 }
1323 start.resize(states.len(), 0.0);
1324 let mut steps = MAX_ELIMINATION;
1325 let solution = chain.visits(&start, &mut steps);
1326 self.spend(MAX_ELIMINATION - steps, stmt.span)?;
1327 let visits = match solution {
1328 Solution::Visits(v) => v,
1329 Solution::TooBig => {
1330 self.too_large.insert(stmt.id);
1331 return Ok(None);
1332 }
1333 Solution::Stuck(s) => return Err(self.never_leaves(f, stmt, &states[s], &head)),
1334 };
1335 self.stats.solved_loops += 1;
1336 self.stats.chain_states += states.len() as u64;
1337 let mut flow = Flow::default();
1338 let mut left = Weight::ZERO;
1339 for (k, v) in visits.into_iter().enumerate() {
1340 if v == 0.0 {
1341 continue;
1342 }
1343 let times = total.scale(v);
1344 for mut w in std::mem::take(&mut exits[k]) {
1345 w.weight = w.weight * times;
1346 flow.next.push(w);
1347 }
1348 for (value, w, c) in std::mem::take(&mut returns[k]) {
1349 flow.returned.push((value, w * times, c));
1350 }
1351 left += unresolved[k] * times;
1352 }
1353 self.unresolved += left;
1354 Ok(Some(flow))
1355 }
1356
1357 fn never_leaves(&self, f: FnId, stmt: &Stmt, slots: &[Value], head: &[usize]) -> RuntimeError {
1360 let names = &self.prog.functions[f as usize].slots;
1361 let parts: Vec<String> = head
1362 .iter()
1363 .filter(|&&i| names[i].name != TEMP)
1364 .map(|&i| format!("`{}` is {:?}", names[i].name, slots[i]))
1365 .collect();
1366 let err = RuntimeError::new(stmt.span, "some worlds can never leave this loop")
1367 .with_help("check that every world can leave the loop");
1368 if parts.is_empty() {
1369 err
1370 } else {
1371 err.with_note(format!("for example, the worlds where {}", parts.join(" and ")))
1372 }
1373 }
1374
1375 fn split_by(&mut self, f: FnId, place: &Place, d: Value, w: World, span: Span, out: &mut Vec<World>) -> Result<()> {
1377 match d {
1378 Value::Dist(_) | Value::Continuous(_) if self.sampler.is_some() => {
1379 let v = self.sample(&d);
1380 let mut w = w;
1381 self.assign(f, place, v, &mut w, span)?;
1382 out.push(w);
1383 }
1384 Value::Continuous(family) => {
1385 if self.nested > 0 {
1386 return Err(self.continuous_draw(span));
1387 }
1388 self.next_latent += 1;
1389 let v = Analytic::new(self.next_latent, *family)
1390 .value()
1391 .map_err(|e| e.at(span))?;
1392 let mut w = w;
1393 self.assign(f, place, v, &mut w, span)?;
1394 out.push(w);
1395 }
1396 Value::Dist(dist) => {
1397 self.unresolved += w.weight.scale(dist.missing);
1398 let last = dist.outcomes.len().saturating_sub(1);
1399 let mut w = Some(w);
1400 for (i, (v, p)) in dist.outcomes.iter().enumerate() {
1401 let base = if i == last {
1402 w.take().unwrap()
1403 } else {
1404 w.clone().unwrap()
1405 };
1406 let mut nw = base.scaled(*p);
1407 if matches!(v, Value::Continuous(_)) {
1408 self.split_by(f, place, v.clone(), nw, span, out)?;
1409 } else {
1410 self.assign(f, place, v.clone(), &mut nw, span)?;
1411 out.push(nw);
1412 }
1413 }
1414 }
1415 Value::Prob(p) => {
1416 return self.split_by(f, place, Dist::bernoulli(p).into_value(), w, span, out);
1417 }
1418 other => {
1419 let mut w = w;
1420 self.assign(f, place, other, &mut w, span)?;
1421 out.push(w);
1422 }
1423 }
1424 Ok(())
1425 }
1426
1427 fn check_arity(&self, c: &Closure, given: usize, span: Span) -> Result<()> {
1428 let expected = self.prog.functions[c.func as usize].n_params as usize;
1429 if expected != given {
1430 return Err(RuntimeError::new(
1431 span,
1432 format!("this function takes {expected} argument(s), but {given} were given"),
1433 ));
1434 }
1435 Ok(())
1436 }
1437
1438 fn coerce(&mut self, v: Value, ty: &TypeSpec, span: Span) -> Result<Value> {
1441 self.coerce_at(v, ty, span, 0)
1442 }
1443
1444 fn coerce_at(&mut self, v: Value, ty: &TypeSpec, span: Span, depth: usize) -> Result<Value> {
1445 self.budget.work(1).map_err(|e| e.at(span))?;
1446 if depth > 64 {
1447 return Err(RuntimeError::new(
1448 span,
1449 "type conversion nesting exceeds the limit of 64",
1450 ));
1451 }
1452 Ok(match (v, ty) {
1453 (Value::Analytic(_) | Value::Event(_), TypeSpec::Prob) => {
1454 return Err(analytic::unsupported("converting this outcome to prob").at(span));
1455 }
1456 (v @ Value::Float(_), TypeSpec::Int) => Value::Int(
1457 ops::integer(&v, "int conversion", &mut self.budget)
1458 .map_err(|e| e.at(span))?
1459 .into_owned(),
1460 ),
1461 (Value::Analytic(_), TypeSpec::Int) => {
1462 return Err(analytic::unsupported("converting this outcome to int").at(span));
1463 }
1464 (Value::Prob(p), TypeSpec::Float) => Value::Float(p),
1465 (v @ (Value::Int(_) | Value::Float(_)), TypeSpec::Prob) => ops::make_prob(&v).map_err(|e| e.at(span))?,
1466 (Value::List(xs), TypeSpec::List(t)) => {
1467 self.budget.collection(xs.len() as u128).map_err(|e| e.at(span))?;
1468 let mut out = Vec::with_capacity(xs.len());
1469 for x in xs.iter() {
1470 out.push(self.coerce_at(x.clone(), t, span, depth + 1)?);
1471 }
1472 Value::list(out)
1473 }
1474 (Value::Map(xs), TypeSpec::Map(kt, vt)) => {
1475 let mut out = BTreeMap::new();
1476 for (k, v) in xs.iter() {
1477 let key = self.coerce_at(k.clone(), kt, span, depth + 1)?;
1478 let value = self.coerce_at(v.clone(), vt, span, depth + 1)?;
1479 if out.insert(key, value).is_some() {
1480 return Err(RuntimeError::new(span, "map key collision during type conversion"));
1481 }
1482 }
1483 Value::map(out)
1484 }
1485 (Value::Bag(xs), TypeSpec::Bag(t)) => {
1486 let mut out = BTreeMap::<Value, u64>::new();
1487 for (x, n) in xs.iter() {
1488 let key = self.coerce_at(x.clone(), t, span, depth + 1)?;
1489 let count = out.entry(key).or_default();
1490 *count = count
1491 .checked_add(*n)
1492 .ok_or_else(|| RuntimeError::new(span, "bag count overflow during type conversion"))?;
1493 }
1494 Value::multiset(crate::value::Multiset::new(out))
1495 }
1496 (Value::Dist(d), TypeSpec::Dist(t)) => {
1497 let mut out = Vec::with_capacity(d.outcomes.len());
1498 for (x, p) in &d.outcomes {
1499 out.push((self.coerce_at(x.clone(), t, span, depth + 1)?, *p));
1500 }
1501 ops::combine(out, d.missing, &mut self.budget).map_err(|e| e.at(span))?
1502 }
1503 (Value::Record(r), TypeSpec::Record(t)) => {
1504 let types: Vec<_> = self.prog.records[*t as usize]
1505 .fields
1506 .iter()
1507 .map(|f| (f.name.clone(), f.ty.clone()))
1508 .collect();
1509 let mut fields = Vec::with_capacity(r.fields.len());
1510 for (name, value) in &r.fields {
1511 let value = match types.iter().find(|(n, _)| n == &**name) {
1512 Some((_, t)) => self.coerce_at(value.clone(), t, span, depth + 1)?,
1513 None => value.clone(),
1514 };
1515 fields.push((name.clone(), value));
1516 }
1517 ops::make_record(r.ty.clone(), fields)
1518 }
1519 (Value::Record(r), TypeSpec::AnonRecord(types)) => {
1520 let mut fields = Vec::with_capacity(r.fields.len());
1521 for (name, value) in &r.fields {
1522 let value = match types.iter().find(|(n, _)| n == &**name) {
1523 Some((_, t)) => self.coerce_at(value.clone(), t, span, depth + 1)?,
1524 None => value.clone(),
1525 };
1526 fields.push((name.clone(), value));
1527 }
1528 ops::make_record(r.ty.clone(), fields)
1529 }
1530 (v, _) => v,
1531 })
1532 }
1533
1534 fn conforms(&self, v: &Value, ty: &TypeSpec) -> bool {
1536 match (ty, v) {
1537 (TypeSpec::Int, Value::Int(_)) => true,
1538 (TypeSpec::Float, Value::Float(_) | Value::Int(_) | Value::Analytic(_)) => true,
1539 (TypeSpec::Complex, Value::Complex(_)) => true,
1540 (TypeSpec::Prob, Value::Prob(_)) => true,
1541 (TypeSpec::Bool, Value::Bool(_) | Value::Event(_)) => true,
1542 (TypeSpec::Str, Value::Str(_)) => true,
1543 (TypeSpec::Date, Value::Date(_)) => true,
1544 (TypeSpec::Unit, Value::Unit) => true,
1545 (TypeSpec::Function, Value::Closure(_) | Value::Builtin(_)) => true,
1546 (TypeSpec::List(t), Value::List(items)) => items.iter().all(|x| self.conforms(x, t)),
1547 (TypeSpec::List(t), Value::Range(..)) => **t == TypeSpec::Int,
1548 (TypeSpec::Map(k, t), Value::Map(m)) => m.iter().all(|(a, b)| self.conforms(a, k) && self.conforms(b, t)),
1549 (TypeSpec::Bag(t), Value::Bag(b)) => b.keys().all(|x| self.conforms(x, t)),
1550 (TypeSpec::Dist(t), Value::Dist(d)) => d.outcomes.iter().all(|(x, _)| match x {
1551 Value::Continuous(_) => **t == TypeSpec::Float,
1552 _ => self.conforms(x, t),
1553 }),
1554 (TypeSpec::Dist(t), Value::Continuous(_)) => **t == TypeSpec::Float,
1555 (TypeSpec::Record(r), Value::Record(rec)) => {
1556 let declared = &self.prog.records[*r as usize];
1557 rec.ty.as_deref() == Some(declared.name.as_str())
1558 && rec.fields.len() == declared.fields.len()
1559 && declared
1560 .fields
1561 .iter()
1562 .all(|f| rec.get(&f.name).is_some_and(|x| self.conforms(x, &f.ty)))
1563 }
1564 (TypeSpec::Enum(e), Value::Enum(x)) => x.ty == *e,
1565 (TypeSpec::AnonRecord(fields), Value::Record(rec)) => {
1566 rec.fields.len() == fields.len()
1567 && rec
1568 .fields
1569 .iter()
1570 .zip(fields)
1571 .all(|((n, x), (m, t))| &**n == m.as_str() && self.conforms(x, t))
1572 }
1573 _ => false,
1574 }
1575 }
1576
1577 fn call(&mut self, func: FnId, key: Vec<Value>, span: Span) -> Result<Arc<CallResult>> {
1591 self.stats.calls += 1;
1592 let prog = self.prog;
1593 let fun = &prog.functions[func as usize];
1594 let sampling = self.sampler.is_some();
1596 let memoizable = self.config.memoizing && !sampling && !fun.effects.prints && self.callback.is_none();
1599 let call_key = (func, key);
1600 if memoizable {
1601 if let Some(r) = self.memo.get(&call_key) {
1602 self.stats.memo_hits += 1;
1603 self.observed |= r.observed;
1604 return Ok(r.clone());
1605 }
1606 }
1607 if !sampling {
1608 if let Some(&depth) = self.active.get(&call_key) {
1609 if fun.effects.prints {
1610 return Err(RuntimeError::new(
1611 span,
1612 format!("`{}` comes back to itself, and prints", describe_call(fun, &call_key.1)),
1613 )
1614 .with_note("a call that comes back to itself runs several times, until its result settles, and would print each time")
1615 .with_help("remove the `print`, or print the result where the call is made"));
1616 }
1617 let so_far = match self.approx.get(&call_key) {
1618 Some(a) if self.frames_at.get(a.head) == Some(&a.frame) => a.result.clone(),
1619 _ => {
1620 let result = Arc::new(CallResult::waiting(depth));
1621 let approx = Approx {
1622 result: result.clone(),
1623 head: depth,
1624 frame: self.frames_at[depth],
1625 round: self.rounds_at[depth],
1626 };
1627 self.approx.insert(call_key, approx);
1628 result
1629 }
1630 };
1631 return Ok(so_far);
1632 }
1633 if let Some(a) = self.approx.get(&call_key) {
1635 if self.frames_at.get(a.head) == Some(&a.frame) && self.rounds_at.get(a.head) == Some(&a.round) {
1636 self.stats.memo_hits += 1;
1637 return Ok(a.result.clone());
1638 }
1639 }
1640 }
1641 if self.depth >= self.config.max_call_depth {
1642 return Err(RuntimeError::limit(
1643 span,
1644 format!("calls are nested more than {} deep", self.config.max_call_depth),
1645 )
1646 .with_help("check for recursion that doesn't stop, or write it as a loop"));
1647 }
1648 let mut slots = vec![Value::Dead; fun.n_slots()];
1649 let n = fun.n_params as usize;
1650 for (i, v) in call_key.1[..n].iter().enumerate() {
1651 slots[i] = v.clone();
1652 }
1653 for (cap, v) in fun.captures.iter().zip(&call_key.1[n..]) {
1654 slots[cap.slot as usize] = v.clone();
1655 }
1656 let depth = self.depth + 1;
1657 if !sampling {
1658 self.active.insert(call_key.clone(), depth);
1659 if self.frames_at.len() <= depth {
1660 self.frames_at.resize(depth + 1, 0);
1661 self.rounds_at.resize(depth + 1, 0);
1662 }
1663 self.last_number += 1;
1664 self.frames_at[depth] = self.last_number;
1665 }
1666 let result = self.run_call(func, &call_key, slots, depth, span);
1667 if !sampling {
1668 self.active.remove(&call_key);
1669 self.frames_at[depth] = 0;
1671 }
1672 let result = result?;
1673 if memoizable
1675 && result.pending.is_empty()
1676 && result
1677 .outcomes
1678 .iter()
1679 .all(|(v, _, c)| c.is_empty() && !analytic::contains(v))
1680 {
1681 if self.memo.len() >= self.config.max_cached_calls {
1682 self.memo.clear();
1683 }
1684 self.memo.insert(call_key, result.clone());
1685 }
1686 Ok(result)
1687 }
1688
1689 fn run_call(
1699 &mut self,
1700 func: FnId,
1701 call_key: &CallKey,
1702 slots: Vec<Value>,
1703 depth: usize,
1704 span: Span,
1705 ) -> Result<Arc<CallResult>> {
1706 let fun = &self.prog.functions[func as usize];
1707 let sampling = self.sampler.is_some();
1708 let mut inherited = std::collections::BTreeSet::new();
1709 for v in &slots {
1710 analytic::collect_ids(v, &mut inherited);
1711 }
1712 let inherited = Arc::new(inherited);
1713 let mut rounds: u64 = 0;
1714 let mut before = f64::INFINITY;
1715 loop {
1716 if !sampling {
1717 self.last_number += 1;
1718 self.rounds_at[depth] = self.last_number;
1719 }
1720 let saved_unresolved = std::mem::replace(&mut self.unresolved, Weight::ZERO);
1721 let saved_lost = std::mem::replace(&mut self.lost, Weight::ZERO);
1722 let saved_observed = std::mem::replace(&mut self.observed, false);
1723 let saved_pending = std::mem::take(&mut self.pending);
1724 self.depth += 1;
1725 let flow = self.exec_block(
1726 func,
1727 &fun.body,
1728 vec![World {
1729 slots: slots.clone(),
1730 constraints: Default::default(),
1731 inherited: inherited.clone(),
1732 weight: Weight::ONE,
1733 run: 0,
1734 }],
1735 );
1736 self.depth -= 1;
1737 let unresolved = std::mem::replace(&mut self.unresolved, saved_unresolved);
1738 let lost = std::mem::replace(&mut self.lost, saved_lost);
1739 let observed = std::mem::replace(&mut self.observed, saved_observed);
1740 let pending = std::mem::replace(&mut self.pending, saved_pending);
1741 self.observed |= observed;
1742 let flow = flow.map_err(|e| match fun.kind {
1743 FnKind::Named => e.with_note(format!("in a call to `{}`", fun.name)),
1744 FnKind::Simulate => e.with_note("inside a `simulate` block"),
1745 FnKind::Lambda => e.with_note("inside a lambda"),
1746 FnKind::Main => e,
1747 })?;
1748 let mut result = CallResult {
1749 outcomes: merge_values(flow.returned),
1750 unresolved,
1751 lost,
1752 observed,
1753 pending,
1754 };
1755 if let Some(head) = result.pending.iter().map(|&(d, _)| d).filter(|&d| d < depth).min() {
1756 let waiting = Weight::sum(result.pending.iter().map(|&(_, w)| w));
1758 result.pending = vec![(head, waiting)];
1759 let result = Arc::new(result);
1760 let approx = Approx {
1761 result: result.clone(),
1762 head,
1763 frame: self.frames_at[head],
1764 round: self.rounds_at[head],
1765 };
1766 self.approx.insert(call_key.clone(), approx);
1767 return Ok(result);
1768 }
1769 let waiting = result.take_pending(depth);
1770 if waiting.is_zero() {
1771 return Ok(Arc::new(result));
1772 }
1773 if rounds == 0 {
1775 self.stats.solved_calls += 1;
1776 }
1777 let w = waiting.to_f64();
1778 if w <= self.config.epsilon {
1779 self.settle(depth);
1780 return Ok(Arc::new(result));
1781 }
1782 rounds += 1;
1783 self.stats.call_rounds += 1;
1784 if w >= before {
1785 return Err(RuntimeError::new(
1786 span,
1787 format!(
1788 "`{}` never returns for some of its worlds",
1789 describe_call(fun, &call_key.1)
1790 ),
1791 )
1792 .with_note(format!("{} of its weight keeps coming back to it", fmt_prob(w)))
1793 .with_help("check that the recursion can end"));
1794 }
1795 if rounds >= self.config.max_iterations {
1796 return Err(RuntimeError::limit(
1797 span,
1798 format!(
1799 "`{}` came back to itself {} times without settling",
1800 describe_call(fun, &call_key.1),
1801 self.config.max_iterations
1802 ),
1803 )
1804 .with_note(format!("{} of its weight is still waiting on it", fmt_prob(w)))
1805 .with_help("check that the recursion can end, or raise `@epsilon`"));
1806 }
1807 before = w;
1808 result.pending.push((depth, waiting));
1809 let approx = Approx {
1810 result: Arc::new(result),
1811 head: depth,
1812 frame: self.frames_at[depth],
1813 round: self.rounds_at[depth],
1814 };
1815 self.approx.insert(call_key.clone(), approx);
1816 }
1817 }
1818
1819 fn settle(&mut self, depth: usize) {
1823 let frame = self.frames_at[depth];
1824 let epsilon = self.config.epsilon;
1825 let settled: Vec<(CallKey, CallResult)> = self
1826 .approx
1827 .iter()
1828 .filter(|(_, a)| a.head == depth && a.frame == frame)
1829 .filter(|(_, a)| a.result.pending.iter().all(|&(_, w)| w.to_f64() <= epsilon))
1830 .map(|(key, a)| {
1831 let mut result = (*a.result).clone();
1832 result.pending.clear();
1834 (key.clone(), result)
1835 })
1836 .collect();
1837 self.approx.retain(|_, a| a.head != depth);
1838 if !self.config.memoizing {
1839 return;
1840 }
1841 for (key, result) in settled {
1842 if !self.active.contains_key(&key)
1843 && !self.prog.functions[key.0 as usize].effects.prints
1844 && result
1845 .outcomes
1846 .iter()
1847 .all(|(v, _, c)| c.is_empty() && !analytic::contains(v))
1848 {
1849 self.memo.insert(key, Arc::new(result));
1850 }
1851 }
1852 }
1853
1854 fn simulate(&mut self, func: FnId, key: Vec<Value>, span: Span) -> Result<Value> {
1857 if key.iter().any(analytic::contains) {
1858 return Err(analytic::unsupported("capturing an analytic draw inside `simulate`").at(span));
1859 }
1860 let saved_observed = self.observed;
1861 let saved_densities = self.densities;
1862 let sampler = self.sampler.take();
1864 let callback = self.callback.take();
1865 self.nested += 1;
1866 let result = self.call(func, key, span);
1867 self.nested -= 1;
1868 self.sampler = sampler;
1869 self.callback = callback;
1870 self.observed = saved_observed;
1872 self.densities = saved_densities;
1873 let result = result?;
1874 if !result.pending.is_empty() {
1875 return Err(OpError::unsupported(
1876 "a `simulate` block that comes back to a call still running isn't supported yet",
1877 )
1878 .at(span));
1879 }
1880 let resolved = Weight::sum(result.outcomes.iter().map(|(_, w, _)| *w));
1881 if resolved.is_zero() {
1882 let message = if result.unresolved.is_zero() {
1883 "every world in this `simulate` was ruled out by `observe`"
1884 } else {
1885 "this `simulate` left all of its weight unresolved"
1886 };
1887 return Err(RuntimeError::new(span, message));
1888 }
1889 let denom = resolved + result.unresolved;
1890 let pairs = result
1891 .outcomes
1892 .iter()
1893 .map(|(v, w, _)| (v.clone(), w.ratio(denom)))
1894 .collect();
1895 ops::combine(pairs, result.unresolved.ratio(denom), &mut self.budget).map_err(|e| e.at(span))
1896 }
1897
1898 fn call_pure(&mut self, closure: &Value, args: Vec<Value>, what: &'static str, span: Span) -> Result<Value> {
1900 if let Value::Builtin(b) = closure {
1901 return self.builtin_values(*b, &args, &[], Weight::ONE, span);
1902 }
1903 let Value::Closure(c) = closure else {
1904 return Err(RuntimeError::new(
1905 span,
1906 format!(
1907 "`{what}` needs a function, like `x -> x * 2`, found {}",
1908 ops::article(&closure.kind())
1909 ),
1910 ));
1911 };
1912 self.check_arity(c, args.len(), span)?;
1913 let mut key = args;
1914 key.extend(c.captured.iter().cloned());
1915 let previous = self.callback.replace(what);
1916 let result = self.call(c.func, key, span);
1917 self.callback = previous;
1918 let result = result?;
1919 if !result.pending.is_empty() {
1920 return Err(OpError::unsupported(format!(
1921 "a function given to `{what}` that comes back to a call still running isn't supported yet"
1922 ))
1923 .at(span));
1924 }
1925 match result.outcomes.as_slice() {
1926 [(v, p, c)] if c.is_empty() && (p.to_f64() - 1.0).abs() < 1e-12 && result.unresolved.is_zero() => {
1927 Ok(v.clone())
1928 }
1929 _ => Err(RuntimeError::new(
1930 span,
1931 format!("the function given to `{what}` can't branch on chances, draw values or observe"),
1932 )),
1933 }
1934 }
1935
1936 fn check_callback_effect(&self, span: Span) -> Result<()> {
1939 match self.callback {
1940 Some(what) => Err(RuntimeError::new(
1941 span,
1942 format!("the function given to `{what}` can't branch on chances, draw values or observe"),
1943 )
1944 .with_help("use a loop for probabilistic traversal, or `simulate` for a local distribution")),
1945 None => Ok(()),
1946 }
1947 }
1948
1949 fn record_place_type(&self, place: &Place, w: &World) -> Option<TypeSpec> {
1950 if place.path.is_empty() {
1951 return None;
1952 }
1953 let Value::Record(r) = &w.slots[place.slot as usize] else {
1954 return None;
1955 };
1956 let id = self
1957 .prog
1958 .records
1959 .iter()
1960 .position(|t| Some(t.name.as_str()) == r.ty.as_deref())?;
1961 let mut ty = TypeSpec::Record(id as u32);
1962 for element in &place.path {
1963 ty = match (element, ty) {
1964 (PathElem::Field(name), TypeSpec::Record(id)) => self.prog.records[id as usize]
1965 .fields
1966 .iter()
1967 .find(|f| f.name == *name)?
1968 .ty
1969 .clone(),
1970 (PathElem::Field(name), TypeSpec::AnonRecord(fields)) => fields.into_iter().find(|(n, _)| n == name)?.1,
1971 (PathElem::Index(_), TypeSpec::List(t) | TypeSpec::Map(_, t)) => *t,
1972 _ => return None,
1973 };
1974 }
1975 Some(ty)
1976 }
1977
1978 fn assign(&mut self, f: FnId, place: &Place, v: Value, w: &mut World, span: Span) -> Result<()> {
1979 if place.path.is_empty() {
1980 w.slots[place.slot as usize] = v;
1981 return Ok(());
1982 }
1983 let v = match self.record_place_type(place, w) {
1984 Some(ty) => {
1985 let v = self.coerce(v, &ty, span)?;
1986 if !self.conforms(&v, &ty) {
1987 return Err(RuntimeError::new(
1988 span,
1989 format!(
1990 "assigned field should be {}, found {}",
1991 ty.describe(self.prog),
1992 v.kind()
1993 ),
1994 ));
1995 }
1996 v
1997 }
1998 None => v,
1999 };
2000 let mut keys = Vec::with_capacity(place.path.len());
2001 for elem in &place.path {
2002 keys.push(match elem {
2003 PathElem::Field(name) => PathKey::Field(name.as_str()),
2004 PathElem::Index(e) => {
2005 let key = self.eval(f, e, w)?;
2006 if analytic::contains(&key) {
2007 return Err(analytic::unsupported("an analytic assignment index").at(e.span));
2008 }
2009 PathKey::Index(key)
2010 }
2011 });
2012 }
2013 let target = &mut w.slots[place.slot as usize];
2014 if matches!(target, Value::Dead) {
2015 return Err(self.no_value(f, place.slot, span));
2016 }
2017 update(target, &keys, v, &mut self.budget).map_err(|e| e.at(span))
2018 }
2019
2020 fn read_place(&mut self, f: FnId, place: &Place, w: &World, span: Span) -> Result<Value> {
2021 let mut v = self.slot(f, place.slot, w, span)?;
2022 for elem in &place.path {
2023 v = match elem {
2024 PathElem::Field(name) => ops::field(&v, name, &mut self.budget),
2025 PathElem::Index(e) => {
2026 let i = self.eval(f, e, w)?;
2027 ops::index(&v, &i, &mut self.budget)
2028 }
2029 }
2030 .map_err(|e| e.at(span))?;
2031 }
2032 Ok(v)
2033 }
2034
2035 fn slot(&mut self, f: FnId, s: SlotId, w: &World, span: Span) -> Result<Value> {
2036 match &w.slots[s as usize] {
2037 Value::Dead => Err(self.no_value(f, s, span)),
2038 Value::Delayed(_) => {
2039 let name = &self.prog.functions[f as usize].slots[s as usize].name;
2040 Err(
2041 OpError::internal("internal error: a variable was read before it was drawn")
2042 .at(span)
2043 .with_note(format!(
2044 "`{name}`'s draw was delayed for an exact update (docs/semantics.md, section 14)"
2045 )),
2046 )
2047 }
2048 v => analytic::resolve(v, &w.constraints, &mut self.budget).map_err(|e| e.at(span)),
2049 }
2050 }
2051
2052 fn no_value(&self, f: FnId, s: SlotId, span: Span) -> RuntimeError {
2053 let name = &self.prog.functions[f as usize].slots[s as usize].name;
2054 RuntimeError::new(span, format!("`{name}` has no value here")).with_help(
2055 "it's used before it's given a value; if a function reads it, call the function after the variable is set",
2056 )
2057 }
2058
2059 fn eval_expected(&mut self, f: FnId, e: &Expr, w: &World, ty: &TypeSpec) -> Result<Value> {
2062 match (&e.kind, ty) {
2063 (ExprKind::List(xs), TypeSpec::List(t)) => {
2064 self.budget.collection(xs.len() as u128).map_err(|err| err.at(e.span))?;
2065 let mut out = Vec::with_capacity(xs.len());
2066 for x in xs {
2067 out.push(self.eval_expected(f, x, w, t)?);
2068 }
2069 Ok(Value::list(out))
2070 }
2071 (ExprKind::Map(xs), TypeSpec::Map(kt, vt)) => {
2072 let mut out = BTreeMap::new();
2073 for (k, v) in xs {
2074 let key = self.eval_expected(f, k, w, kt)?;
2075 if analytic::contains(&key) {
2076 return Err(analytic::unsupported("analytic map keys").at(k.span));
2077 }
2078 if key.is_uncertain() {
2079 return Err(RuntimeError::new(k.span, "map keys can't be distributions"));
2080 }
2081 out.insert(key, self.eval_expected(f, v, w, vt)?);
2082 }
2083 Ok(Value::map(out))
2084 }
2085 (ExprKind::Record { ty: None, fields }, TypeSpec::AnonRecord(types)) => {
2086 let mut out = Vec::with_capacity(fields.len());
2087 for (n, e) in fields {
2088 let v = match types.iter().find(|(name, _)| name == n) {
2089 Some((_, ty)) => self.eval_expected(f, e, w, ty)?,
2090 None => self.eval(f, e, w)?,
2091 };
2092 out.push((Arc::from(n.as_str()), v));
2093 }
2094 Ok(ops::make_record(None, out))
2095 }
2096 _ => {
2097 let v = self.eval(f, e, w)?;
2098 self.coerce(v, ty, e.span)
2099 }
2100 }
2101 }
2102
2103 fn record_field_type(&self, value: &Value, name: &str) -> Option<TypeSpec> {
2107 match value {
2108 Value::Record(r) => self
2109 .prog
2110 .records
2111 .iter()
2112 .find(|t| Some(t.name.as_str()) == r.ty.as_deref())?
2113 .fields
2114 .iter()
2115 .find(|f| f.name == name)
2116 .map(|f| f.ty.clone()),
2117 Value::Dist(d) => {
2118 let first = self.record_field_type(&d.outcomes.first()?.0, name)?;
2119 d.outcomes
2120 .iter()
2121 .all(|(v, _)| self.record_field_type(v, name).as_ref() == Some(&first))
2122 .then_some(first)
2123 }
2124 _ => None,
2125 }
2126 }
2127
2128 fn check_updated_records(&mut self, value: Value, span: Span) -> Result<Value> {
2129 match value {
2130 Value::Record(ref r) => {
2131 let Some(id) = self
2132 .prog
2133 .records
2134 .iter()
2135 .position(|t| Some(t.name.as_str()) == r.ty.as_deref())
2136 else {
2137 return Ok(value);
2138 };
2139 let ty = TypeSpec::Record(id as u32);
2140 let value = self.coerce(value, &ty, span)?;
2141 if !self.conforms(&value, &ty) {
2142 return Err(RuntimeError::new(
2143 span,
2144 format!("updated record doesn't conform to {}", ty.describe(self.prog)),
2145 ));
2146 }
2147 Ok(value)
2148 }
2149 Value::Dist(d) => {
2150 let mut outcomes = Vec::with_capacity(d.outcomes.len());
2151 for (v, p) in &d.outcomes {
2152 outcomes.push((self.check_updated_records(v.clone(), span)?, *p));
2153 }
2154 ops::combine(outcomes, d.missing, &mut self.budget).map_err(|e| e.at(span))
2155 }
2156 v => Ok(v),
2157 }
2158 }
2159
2160 fn eval(&mut self, f: FnId, e: &Expr, w: &World) -> Result<Value> {
2161 let span = e.span;
2162 let at = |err: OpError| err.at(span);
2163 match &e.kind {
2164 ExprKind::Lit(l) => self.literal(l).map_err(at),
2165 ExprKind::Slot(s) => self.slot(f, *s, w, span),
2166 ExprKind::Unary(op, x) => {
2167 let v = self.eval(f, x, w)?;
2168 ops::unary(*op, &v, &mut self.budget).map_err(at)
2169 }
2170 ExprKind::Binary(op @ (BinOp::And | BinOp::Or), a, b) => {
2171 let and = *op == BinOp::And;
2172 let word = if and { "and" } else { "or" };
2173 let va = self.eval(f, a, w)?;
2174 let ta = ops::truth(&va, word).map_err(|err| err.at(a.span))?;
2175 if let Truth::Fact(x) = ta {
2176 if x != and {
2177 return Ok(Value::Bool(x));
2179 }
2180 }
2181 let vb = self.eval(f, b, w)?;
2182 let tb = ops::truth(&vb, word).map_err(|err| err.at(b.span))?;
2183 ops::logic(and, ta, tb, &mut self.budget).map_err(at)
2184 }
2185 ExprKind::Binary(op, a, b) => {
2186 let va = self.eval(f, a, w)?;
2187 let vb = self.eval(f, b, w)?;
2188 ops::binary(*op, &va, &vb, &mut self.budget).map_err(at)
2189 }
2190 ExprKind::List(items) => {
2191 self.budget.collection(items.len() as u128).map_err(at)?;
2192 let mut values = Vec::with_capacity(items.len());
2193 for item in items {
2194 values.push(self.eval(f, item, w)?);
2195 }
2196 Ok(Value::list(values))
2197 }
2198 ExprKind::Map(entries) => {
2199 let mut map = BTreeMap::new();
2200 for (k, v) in entries {
2201 let key = self.eval(f, k, w)?;
2202 if analytic::contains(&key) {
2203 return Err(analytic::unsupported("analytic map keys").at(k.span));
2204 }
2205 if key.is_uncertain() {
2206 return Err(RuntimeError::new(k.span, "map keys can't be distributions"));
2207 }
2208 let value = self.eval(f, v, w)?;
2209 map.insert(key, value);
2210 }
2211 Ok(Value::map(map))
2212 }
2213 ExprKind::Record { ty, fields } => {
2214 let mut values = Vec::with_capacity(fields.len());
2215 for (name, v) in fields {
2216 let field_ty = ty
2217 .and_then(|t| self.prog.records[t as usize].fields.iter().find(|d| d.name == *name))
2218 .map(|d| d.ty.clone());
2219 let value = match field_ty {
2220 Some(t) => {
2221 let value = self.eval(f, v, w)?;
2222 let value = self.coerce(value, &t, v.span)?;
2223 if !self.conforms(&value, &t) {
2224 return Err(RuntimeError::new(
2225 v.span,
2226 format!(
2227 "field `{name}` should be {}, found {}",
2228 t.describe(self.prog),
2229 value.kind()
2230 ),
2231 ));
2232 }
2233 value
2234 }
2235 None => self.eval(f, v, w)?,
2236 };
2237 values.push((Arc::from(name.as_str()), value));
2238 }
2239 Ok(ops::make_record(
2240 ty.map(|t| self.record_names[t as usize].clone()),
2241 values,
2242 ))
2243 }
2244 ExprKind::Field(x, name) => {
2245 let v = self.eval(f, x, w)?;
2246 ops::field(&v, name, &mut self.budget).map_err(at)
2247 }
2248 ExprKind::Index(x, i) => {
2249 let v = self.eval(f, x, w)?;
2250 let i = self.eval(f, i, w)?;
2251 ops::index(&v, &i, &mut self.budget).map_err(at)
2252 }
2253 ExprKind::With(x, fields) => {
2254 let base = self.eval(f, x, w)?;
2255 let mut updates = Vec::with_capacity(fields.len());
2256 for (name, v) in fields {
2257 let field_ty = self.record_field_type(&base, name);
2258 let value = match field_ty {
2259 Some(ty) => {
2260 let value = self.eval_expected(f, v, w, &ty)?;
2261 if !self.conforms(&value, &ty) {
2262 return Err(RuntimeError::new(
2263 v.span,
2264 format!(
2265 "field `{name}` should be {}, found {}",
2266 ty.describe(self.prog),
2267 value.kind()
2268 ),
2269 ));
2270 }
2271 value
2272 }
2273 None => self.eval(f, v, w)?,
2274 };
2275 updates.push((Arc::from(name.as_str()), value));
2276 }
2277 let value = ops::lift1(&base, &mut self.budget, |b, _| ops::with_fields(b, &updates)).map_err(at)?;
2278 self.check_updated_records(value, span)
2279 }
2280 ExprKind::Builtin { func, args, named } => self.builtin(f, *func, args, named, w, span),
2281 ExprKind::Closure { func, capture_args } => Ok(Value::Closure(Arc::new(Closure {
2282 func: *func,
2283 captured: capture_args
2284 .iter()
2285 .map(|&s| self.slot(f, s, w, span))
2286 .collect::<Result<_>>()?,
2287 }))),
2288 ExprKind::Simulate { func, capture_args } => {
2289 let key = capture_args
2290 .iter()
2291 .map(|&s| self.slot(f, s, w, span))
2292 .collect::<Result<_>>()?;
2293 self.simulate(*func, key, span)
2294 }
2295 ExprKind::Interp(parts) => {
2296 let mut text = String::new();
2297 for part in parts {
2298 match part {
2299 InterpPart::Lit(s) => crate::text::push(&mut text, s, &mut self.budget).map_err(at)?,
2300 InterpPart::Expr(e) => {
2301 let value = self.eval(f, e, w)?;
2302 crate::text::push_value(&mut text, &value, &mut self.budget).map_err(at)?;
2303 }
2304 }
2305 }
2306 Ok(Value::str(&text))
2307 }
2308 ExprKind::Input(i) => Ok(self.inputs[*i as usize].clone()),
2309 }
2310 }
2311
2312 fn literal(&mut self, l: &Lit) -> OpResult<Value> {
2313 Ok(match l {
2314 Lit::Builtin(b) => Value::Builtin(*b),
2315 Lit::Unit => Value::Unit,
2316 Lit::Bool(b) => Value::Bool(*b),
2317 Lit::Int(i) => {
2318 self.budget.integer_bits(i.bits())?;
2319 self.budget.work(i.bits().div_ceil(64).max(1))?;
2320 Value::Int(i.clone())
2321 }
2322 Lit::Float(x) | Lit::FloatConstant(x) => Value::Float(*x),
2323 Lit::Prob(p) => Value::Prob(*p),
2324 Lit::Str(s) => crate::text::value(s, &mut self.budget)?,
2325 Lit::Dice { count, sides } => {
2326 if let Some(d) = self.dice.get(&(*count, *sides)) {
2327 return Ok(d.clone());
2328 }
2329 let d = Dist::dice(*count, *sides, &mut self.budget)?.into_value();
2330 self.dice.insert((*count, *sides), d.clone());
2331 d
2332 }
2333 Lit::Enum { ty, variant } => self.enum_values[*ty as usize][*variant as usize].clone(),
2334 })
2335 }
2336
2337 fn builtin(
2338 &mut self,
2339 f: FnId,
2340 b: Builtin,
2341 args: &[Expr],
2342 named: &[(String, Expr)],
2343 w: &World,
2344 span: Span,
2345 ) -> Result<Value> {
2346 if b == Builtin::Typeof {
2347 let v = match args[0].kind {
2350 ExprKind::Slot(s) if matches!(w.slots[s as usize], Value::Delayed(_)) => w.slots[s as usize].clone(),
2351 _ => self.eval(f, &args[0], w)?,
2352 };
2353 return crate::type_name::of(&v, &self.prog.enums, &mut self.budget).map_err(|e| e.at(span));
2354 }
2355 let mut values = Vec::with_capacity(args.len());
2356 for a in args {
2357 values.push(self.eval(f, a, w)?);
2358 }
2359 let mut names = Vec::with_capacity(named.len());
2360 for (name, value) in named {
2361 names.push(name.clone());
2362 values.push(self.eval(f, value, w)?);
2363 }
2364 self.builtin_values(b, &values, &names, w.weight, span)
2365 }
2366
2367 fn check_function_arity(&self, value: &Value, given: usize, span: Span) -> Result<()> {
2368 match value {
2369 Value::Closure(c) => self.check_arity(c, given, span),
2370 Value::Builtin(b) => {
2371 let (min, max) = b.arity();
2372 if given < min || given > max || (*b == Builtin::Date && given == 2) {
2373 return Err(RuntimeError::new(
2374 span,
2375 format!(
2376 "`{}` takes {}, but {given} were given",
2377 b.name(),
2378 if *b == Builtin::Date {
2379 "one ISO string or three integers".into()
2380 } else if min == max {
2381 format!("{min} argument(s)")
2382 } else if max == usize::MAX {
2383 format!("at least {min} arguments")
2384 } else {
2385 format!("{min} to {max} arguments")
2386 }
2387 ),
2388 ));
2389 }
2390 Ok(())
2391 }
2392 _ => Err(RuntimeError::new(
2393 span,
2394 "expected a comparator function, like `(a, b) -> a - b`",
2395 )),
2396 }
2397 }
2398
2399 fn builtin_values(
2401 &mut self,
2402 b: Builtin,
2403 values: &[Value],
2404 named: &[String],
2405 weight: Weight,
2406 span: Span,
2407 ) -> Result<Value> {
2408 let extrema = matches!(b, Builtin::Minimum | Builtin::Maximum);
2409 if !named.is_empty() && (!extrema || named != ["default"]) {
2410 return Err(RuntimeError::new(
2411 span,
2412 format!(
2413 "`{}` does not accept these named arguments; only minimum/maximum support `default`",
2414 b.name()
2415 ),
2416 ));
2417 }
2418 let positional = values.len() - named.len();
2419 if !named.is_empty() && positional == 0 {
2420 return Err(RuntimeError::new(
2421 span,
2422 "a collection argument is required before `default`",
2423 ));
2424 }
2425 if !named.is_empty() && positional >= 3 {
2426 return Err(RuntimeError::new(
2427 span,
2428 "the default was supplied both positionally and by name",
2429 ));
2430 }
2431 self.check_function_arity(&Value::Builtin(b), values.len(), span)?;
2432 let (values, default) = if extrema && (!named.is_empty() || values.len() == 3) {
2433 (&values[..values.len() - 1], values.last())
2434 } else {
2435 (values, None)
2436 };
2437 let at = |err: OpError| err.at(span);
2438 builtins::check_query_input(b, values).map_err(at)?;
2439 if values.iter().any(analytic::contains)
2440 && !matches!(
2441 b,
2442 Builtin::Map
2443 | Builtin::Filter
2444 | Builtin::Reduce
2445 | Builtin::Len
2446 | Builtin::IterItems
2447 | Builtin::Settled
2448 | Builtin::BooleanLaw
2449 )
2450 {
2451 return Err(analytic::unsupported(&format!("`{}` on this outcome", b.name())).at(span));
2452 }
2453 match b {
2454 Builtin::Typeof => crate::type_name::of(&values[0], &self.prog.enums, &mut self.budget).map_err(at),
2455 Builtin::RunDate => self.config.today.map(Value::Date).ok_or_else(|| {
2456 RuntimeError::new(span, "the host did not supply an execution date for `today`")
2457 .with_help("set Options.today, or supply today in the WASM request")
2458 }),
2459 Builtin::Print => {
2460 let mut text = String::new();
2461 for (i, value) in values.iter().enumerate() {
2462 if i != 0 {
2463 crate::text::push(&mut text, " ", &mut self.budget).map_err(at)?;
2464 }
2465 crate::text::push_value(&mut text, value, &mut self.budget).map_err(at)?;
2466 }
2467 let line = if weight == Weight::ONE {
2468 text
2469 } else {
2470 format!("[{}] {text}", fmt_weight(weight))
2471 };
2472 self.printed += line.len() + 1;
2473 if self.printed > self.config.max_output {
2474 return Err(too_much_output(span));
2475 }
2476 match &mut self.lines {
2477 Some(lines) => lines.push((span, line)),
2478 None => (self.print)(&line),
2479 }
2480 Ok(Value::Unit)
2481 }
2482 Builtin::Map | Builtin::Filter | Builtin::Reduce => self.higher_order(b, values, span),
2483 Builtin::Count if values.len() == 2 => self.higher_order(b, values, span),
2484 Builtin::Minimum | Builtin::Maximum => {
2485 if values.len() == 2 {
2486 self.collection_extreme(b, values, default, span)
2487 } else {
2488 builtins::population_extreme(&values[0], b == Builtin::Maximum, default, &mut self.budget)
2489 .map_err(at)
2490 }
2491 }
2492 Builtin::Sort | Builtin::SortDesc if values.len() == 2 => self.higher_order(b, values, span),
2493 Builtin::Highest | Builtin::Lowest if values.len() == 3 => self.higher_order(b, values, span),
2494 Builtin::Min | Builtin::Max => builtins::call_plain(b, values, &mut self.budget).map_err(at),
2495 Builtin::Roll => self.roll(values).map_err(at),
2496 Builtin::Take => Err(RuntimeError::new(
2497 span,
2498 "use `deck.take()` to draw and remove an item from a mutable bag",
2499 )),
2500 _ if b.lifting() == Lifting::Raw => builtins::call_raw(b, values, &mut self.budget).map_err(at),
2501 _ if values.iter().any(|v| matches!(v, Value::Continuous(_))) => {
2502 let kind = values
2503 .iter()
2504 .find(|v| matches!(v, Value::Continuous(_)))
2505 .unwrap()
2506 .kind();
2507 Err(
2508 RuntimeError::new(span, format!("`{}` needs a value, not a {kind}", b.name()))
2509 .with_help("draw a value first, like `let x ~ normal(0, 1)`"),
2510 )
2511 }
2512 _ => ops::lift_n(values, &mut self.budget, &|a, budget| {
2513 builtins::call_plain(b, a, budget)
2514 })
2515 .map_err(at),
2516 }
2517 }
2518
2519 fn roll(&mut self, values: &[Value]) -> OpResult<Value> {
2520 if values[0].is_uncertain() {
2521 return Err(OpError::new("the number of dice to roll must be a plain number")
2522 .help("draw it first, like `let n ~ d4`"));
2523 }
2524 let count = ops::integer(&values[0], "roll count", &mut self.budget)?;
2525 let count = count
2526 .to_u64()
2527 .filter(|n| *n <= 1000)
2528 .ok_or_else(|| OpError::new("roll needs between 0 and 1000 dice"))? as u32;
2529 let die = match &values[1] {
2530 v @ (Value::Int(_) | Value::Float(_)) => {
2531 let sides = ops::integer(v, "roll sides", &mut self.budget)?;
2532 let sides = sides
2533 .to_u64()
2534 .filter(|n| *n >= 1 && *n <= u32::MAX as u64)
2535 .ok_or_else(|| OpError::new("roll needs between 1 and 4294967295 sides"))?;
2536 Dist::dice(1, sides as u32, &mut self.budget)?
2537 }
2538 Value::Dist(d) => (**d).clone(),
2539 v => {
2540 return Err(OpError::new(format!(
2541 "roll needs a die, like `d6`, found {}",
2542 ops::article(&v.kind())
2543 )));
2544 }
2545 };
2546 let key = (count, die.clone().into_value());
2547 if let Some(pool) = self.pools.get(&key) {
2548 return Ok(pool.clone());
2549 }
2550 let pool = Dist::pool(count, &die, &mut self.budget)?.into_value();
2551 self.pools.insert(key, pool.clone());
2552 Ok(pool)
2553 }
2554
2555 fn compare_callback(
2556 &mut self,
2557 comparator: &Value,
2558 a: &Value,
2559 other: &Value,
2560 b: Builtin,
2561 span: Span,
2562 ) -> Result<std::cmp::Ordering> {
2563 self.budget.work(1).map_err(|e| e.at(span))?;
2564 match self.call_pure(comparator, vec![a.clone(), other.clone()], b.name(), span)? {
2565 Value::Int(n) => Ok(n.cmp(&probl_number::Integer::ZERO)),
2566 Value::Float(x) if x.is_finite() => Ok(x.partial_cmp(&0.0).expect("finite comparator result")),
2567 other => Err(RuntimeError::new(
2568 span,
2569 format!(
2570 "the comparator given to `{}` must return a finite int or float (negative, zero, or positive), found {}",
2571 b.name(),
2572 ops::article(&other.kind())
2573 ),
2574 )),
2575 }
2576 }
2577
2578 fn collection_extreme(
2579 &mut self,
2580 b: Builtin,
2581 values: &[Value],
2582 default: Option<&Value>,
2583 span: Span,
2584 ) -> Result<Value> {
2585 if matches!(values[0], Value::Dist(_) | Value::Continuous(_)) {
2586 return Err(RuntimeError::new(
2587 span,
2588 "a comparator is supported only for collection extrema, not distribution support bounds",
2589 ));
2590 }
2591 let comparator = &values[1];
2592 self.check_function_arity(comparator, 2, span)?;
2593 let items = builtins::items(&values[0], b.name(), &mut self.budget).map_err(|e| e.at(span))?;
2594 let mut iter = items.into_iter();
2595 let Some(mut best) = iter.next() else {
2596 return builtins::empty_extreme(b.name(), default).map_err(|e| e.at(span));
2597 };
2598 for x in iter {
2599 let order = self.compare_callback(comparator, &x, &best, b, span)?;
2600 if (b == Builtin::Maximum && order.is_gt()) || (b == Builtin::Minimum && order.is_lt()) {
2601 best = x;
2602 }
2603 }
2604 Ok(best)
2605 }
2606
2607 fn higher_order(&mut self, b: Builtin, values: &[Value], span: Span) -> Result<Value> {
2609 if let Value::Dist(d) = &values[0] {
2610 let mut results = Vec::with_capacity(d.outcomes.len());
2611 for (coll, p) in &d.outcomes {
2612 let mut args = values.to_vec();
2613 args[0] = coll.clone();
2614 results.push((self.higher_order(b, &args, span)?, *p));
2615 }
2616 return ops::combine(results, d.missing, &mut self.budget).map_err(|e| e.at(span));
2617 }
2618 if matches!(b, Builtin::Highest | Builtin::Lowest) {
2619 if let Value::Dist(d) = &values[1] {
2620 let mut results = Vec::with_capacity(d.outcomes.len());
2621 for (count, p) in &d.outcomes {
2622 let mut args = values.to_vec();
2623 args[1] = count.clone();
2624 results.push((self.higher_order(b, &args, span)?, *p));
2625 }
2626 return ops::combine(results, d.missing, &mut self.budget).map_err(|e| e.at(span));
2627 }
2628 }
2629 let items = builtins::items(&values[0], b.name(), &mut self.budget).map_err(|e| e.at(span))?;
2630 match b {
2631 Builtin::Sort | Builtin::SortDesc | Builtin::Highest | Builtin::Lowest => {
2632 let comparator = values.last().unwrap();
2633 self.check_function_arity(comparator, 2, span)?;
2635 let count = if matches!(b, Builtin::Highest | Builtin::Lowest) {
2636 Some(builtins::extreme_count(&values[1], &mut self.budget).map_err(|e| e.at(span))?)
2637 } else {
2638 None
2639 };
2640 crate::ordering::reserve_sort(items.len(), &mut self.budget).map_err(|e| e.at(span))?;
2641 let mut items = items;
2642 crate::ordering::try_sort_by(&mut items, |a, other| {
2643 let order = self.compare_callback(comparator, a, other, b, span)?;
2644 Ok(if matches!(b, Builtin::SortDesc | Builtin::Highest) {
2645 order.reverse()
2646 } else {
2647 order
2648 })
2649 })?;
2650 if let Some(n) = count {
2651 items.truncate(n);
2652 }
2653 Ok(Value::list(items))
2654 }
2655 Builtin::Map => {
2656 let mut out = Vec::with_capacity(items.len());
2657 for x in items {
2658 out.push(self.call_pure(&values[1], vec![x], "map", span)?);
2659 }
2660 Ok(Value::list(out))
2661 }
2662 Builtin::Filter | Builtin::Count => {
2663 let mut kept = Vec::new();
2664 for x in items {
2665 let test = self.call_pure(&values[1], vec![x.clone()], b.name(), span)?;
2666 match test {
2667 Value::Bool(true) => kept.push(x),
2668 Value::Bool(false) => {}
2669 other => {
2670 return Err(RuntimeError::new(
2671 span,
2672 format!("the test given to `{}` must give a fact (true or false)", b.name()),
2673 )
2674 .with_help(format!("it gave {}", ops::article(&other.kind()))));
2675 }
2676 }
2677 }
2678 Ok(if b == Builtin::Count {
2679 Value::Int((kept.len() as i64).into())
2680 } else {
2681 Value::list(kept)
2682 })
2683 }
2684 Builtin::Reduce => {
2685 let callback = &values[1];
2686 if !matches!(callback, Value::Closure(_) | Value::Builtin(_)) {
2687 return Err(
2688 RuntimeError::new(span, "`reduce` needs a function as its second argument")
2689 .with_help("write `reduce(xs, f)` or `reduce(xs, f, initial)`"),
2690 );
2691 }
2692 self.check_function_arity(callback, 2, span)?;
2693 let mut items = items.into_iter();
2694 let mut acc = match values.get(2) {
2695 Some(initial) => initial.clone(),
2696 None => items.next().ok_or_else(|| {
2697 RuntimeError::new(span, "`reduce` of an empty collection needs an initial value")
2698 .with_help("supply an initial value: `reduce(xs, f, initial)`")
2699 })?,
2700 };
2701 for x in items {
2702 acc = self.call_pure(callback, vec![acc, x], "reduce", span)?;
2703 }
2704 Ok(acc)
2705 }
2706 _ => unreachable!(),
2707 }
2708 }
2709}
2710
2711fn describe_call(fun: &Function, key: &[Value]) -> String {
2713 let args: Vec<String> = key[..fun.n_params as usize].iter().map(|v| format!("{v:?}")).collect();
2714 format!("{}({})", fun.name, args.join(", "))
2715}
2716
2717pub fn too_much_output(span: Span) -> RuntimeError {
2719 RuntimeError::limit(span, "the program printed more than the output limit")
2720}
2721
2722pub fn fmt_weight(w: Weight) -> String {
2724 let x = w.to_f64();
2725 if x == 0.0 && !w.is_zero() {
2726 let l = w.log10() + 2.0;
2727 return format!("{:.1}e{}%", libm::pow(10.0, l - l.floor()), l.floor());
2728 }
2729 fmt_prob(x)
2730}
2731
2732enum PathKey<'a> {
2733 Field(&'a str),
2734 Index(Value),
2735}
2736
2737fn update(target: &mut Value, keys: &[PathKey], v: Value, budget: &mut Budget) -> OpResult<()> {
2738 let Some((first, rest)) = keys.split_first() else {
2739 *target = v;
2740 return Ok(());
2741 };
2742 match (first, target) {
2743 (PathKey::Field(name), Value::Record(r)) => {
2744 let r = Arc::make_mut(r);
2745 let slot = r
2746 .get_mut(name)
2747 .ok_or_else(|| OpError::new(format!("this record has no field `{name}`")))?;
2748 update(slot, rest, v, budget)
2749 }
2750 (PathKey::Index(i), Value::List(items)) => {
2751 let items = Arc::make_mut(items);
2752 let k = ops::as_index(i, items.len() as u128, budget)? as usize;
2753 update(&mut items[k], rest, v, budget)
2754 }
2755 (PathKey::Index(k), Value::Map(m)) => {
2756 let m = Arc::make_mut(m);
2757 if rest.is_empty() {
2758 m.insert(k.clone(), v);
2759 return Ok(());
2760 }
2761 let slot = m
2762 .get_mut(k)
2763 .ok_or_else(|| OpError::new(format!("the key {k:?} isn't in the map")))?;
2764 update(slot, rest, v, budget)
2765 }
2766 (PathKey::Field(name), other) => Err(OpError::new(format!(
2767 "can't set the field `{name}` of {}",
2768 ops::article(&other.kind())
2769 ))),
2770 (PathKey::Index(_), other) => Err(OpError::new(format!(
2771 "can't set an element of {}",
2772 ops::article(&other.kind())
2773 ))),
2774 }
2775}
2776
2777fn is_density(d: &Value) -> bool {
2779 match d {
2780 Value::Continuous(_) => true,
2781 Value::Dist(dist) => dist.outcomes.iter().any(|(x, _)| matches!(x, Value::Continuous(_))),
2782 _ => false,
2783 }
2784}
2785
2786fn likelihood(d: &Value, v: &Value, sampling: bool) -> OpResult<(f64, f64, f64)> {
2791 let continuous = match d {
2792 Value::Continuous(f) => Some(vec![(**f, 1.0)]),
2793 Value::Dist(dist) if dist.outcomes.iter().any(|(x, _)| matches!(x, Value::Continuous(_))) => {
2794 let parts: Option<Vec<_>> = dist
2795 .outcomes
2796 .iter()
2797 .map(|(x, p)| match x {
2798 Value::Continuous(f) => Some((**f, *p)),
2799 _ => None,
2800 })
2801 .collect();
2802 let parts = parts.ok_or_else(|| {
2803 OpError::new("`observe … from` can't mix a density with the probabilities of single values")
2804 })?;
2805 Some(parts)
2806 }
2807 _ => None,
2808 };
2809 if let Some(parts) = continuous {
2810 if !sampling {
2811 return Err(OpError::new("observing a value from a continuous distribution needs sample mode")
2812 .help("its density isn't a probability, so enumeration can't use it; sample with `@mode sample(runs: 10_000)`"));
2813 }
2814 let x = match v {
2815 Value::Bool(_) => None,
2816 _ => v.as_f64(),
2817 }
2818 .ok_or_else(|| {
2819 OpError::new(format!(
2820 "a continuous distribution can't produce {}",
2821 ops::article(&v.kind())
2822 ))
2823 })?;
2824 return Ok((parts.iter().map(|(f, p)| p * f.pdf(x)).sum(), 0.0, 0.0));
2825 }
2826 match d {
2827 Value::Dist(dist) => {
2828 let (mut seen, mut other) = (0.0, 0.0);
2829 for (x, p) in &dist.outcomes {
2830 if ops::equals(x, v) {
2831 seen += p;
2832 } else {
2833 other += p;
2834 }
2835 }
2836 Ok((seen, dist.missing, other))
2837 }
2838 Value::Prob(_) => Err(OpError::new("`observe … from` needs a distribution, not a probability")
2839 .help("to observe that a fact with probability p is true, write `observe true from bernoulli(p)`")),
2840 other => Ok(if ops::equals(other, v) {
2841 (1.0, 0.0, 0.0)
2842 } else {
2843 (0.0, 0.0, 1.0)
2844 }),
2845 }
2846}