Skip to main content

probl_engine/
world.rs

1//! Worlds, and merging the ones that have become identical.
2
3use crate::analytic::Constraints;
4use crate::error::RuntimeError;
5use crate::value::Value;
6use crate::weight::Weight;
7use probl_sema::SlotSet;
8use probl_sema::ir::SlotId;
9use rustc_hash::{FxHashMap, FxHasher};
10use std::collections::BTreeSet;
11use std::collections::hash_map::Entry;
12use std::hash::{Hash, Hasher};
13use std::sync::Arc;
14
15/// A returned value, the weight of the world that returned it, and the
16/// restrictions on analytic draws that its caller can see.
17pub type Returned = (Value, Weight, Constraints);
18
19/// One possible state of the program: the variables of the current frame,
20/// and how likely it is (including every observation so far).
21#[derive(Clone, Debug)]
22pub struct World {
23    pub slots: Vec<Value>,
24    pub constraints: Constraints,
25    /// Latents passed into this call: restrictions on these remain observable
26    /// by its caller even after every local alias dies. Constant per frame.
27    pub inherited: Arc<BTreeSet<u64>>,
28    pub weight: Weight,
29    /// When sampling, the run this world is (docs/semantics.md, section 14).
30    pub run: u32,
31}
32
33impl World {
34    pub fn scaled(mut self, factor: f64) -> World {
35        self.weight = self.weight.scale(factor);
36        self
37    }
38}
39
40pub fn total_weight(worlds: &[World]) -> Weight {
41    Weight::sum(worlds.iter().map(|w| w.weight))
42}
43
44/// Worlds leaving a statement, grouped by how they left it.
45#[derive(Debug, Default)]
46pub struct Flow {
47    /// Continue with the next statement.
48    pub next: Vec<World>,
49    pub broke: Vec<World>,
50    pub continued: Vec<World>,
51    /// Left the function: the returned value and the world's weight.
52    pub returned: Vec<Returned>,
53    /// Faulted inside a `try` that may catch the fault: each world as it
54    /// was when its statement began, with its fault, on its way to the
55    /// `catch`.
56    pub faulted: Vec<(World, RuntimeError)>,
57}
58
59impl Flow {
60    pub fn next(worlds: Vec<World>) -> Flow {
61        Flow {
62            next: worlds,
63            ..Flow::default()
64        }
65    }
66
67    pub fn join(&mut self, other: Flow) {
68        self.next.extend(other.next);
69        self.broke.extend(other.broke);
70        self.continued.extend(other.continued);
71        self.returned.extend(other.returned);
72        self.faulted.extend(other.faulted);
73    }
74}
75
76/// Clear these slots: nobody will read them again.
77pub fn clear(worlds: &mut [World], slots: &[SlotId]) {
78    if slots.is_empty() {
79        return;
80    }
81    for w in worlds {
82        for &s in slots {
83            w.slots[s as usize] = Value::Dead;
84        }
85    }
86}
87
88/// Clear every slot that isn't live, in worlds that jumped (with `break` or
89/// `continue`) past statements that would have cleared some of them.
90pub fn clear_dead(worlds: &mut [World], live: &SlotSet) {
91    let Some(n) = worlds.first().map(|w| w.slots.len()) else {
92        return;
93    };
94    let dead: Vec<SlotId> = live.iter_missing(n).collect();
95    clear(worlds, &dead);
96}
97
98/// The live slots among a frame's `n`.
99pub fn live_slots(live: &SlotSet, n: usize) -> Vec<usize> {
100    live.iter().map(|s| s as usize).filter(|&s| s < n).collect()
101}
102
103/// A hash of a world's live slots, as merging computes it.
104pub fn state_hash(w: &World, live: &[usize]) -> u64 {
105    let mut hasher = FxHasher::default();
106    for &i in live {
107        w.slots[i].hash(&mut hasher);
108    }
109    w.constraints.hash(&mut hasher);
110    hasher.finish()
111}
112
113/// If `enabled`, merge the worlds that agree on the live slots, adding up
114/// their weights, and keeping the order of first appearance. Dead slots
115/// don't matter, cleared or not.
116pub fn merge(mut worlds: Vec<World>, live: &SlotSet, enabled: bool) -> Vec<World> {
117    for w in &mut worlds {
118        if w.constraints.keys().all(|id| w.inherited.contains(id)) {
119            continue;
120        }
121        let mut ids = (*w.inherited).clone();
122        for slot in live.iter() {
123            crate::analytic::collect_ids(&w.slots[slot as usize], &mut ids);
124        }
125        Arc::make_mut(&mut w.constraints).retain(|id, _| ids.contains(id));
126    }
127    if !enabled || worlds.len() < 2 {
128        return worlds;
129    }
130    let n = worlds[0].slots.len();
131    let live = live_slots(live, n);
132    let mut out: Vec<World> = Vec::with_capacity(worlds.len());
133    // The first world kept with each hash, and after each world, the next
134    // one with the same hash (only when different worlds' hashes collide).
135    let mut first: FxHashMap<u64, usize> = FxHashMap::with_capacity_and_hasher(worlds.len(), Default::default());
136    let mut next: Vec<usize> = Vec::with_capacity(worlds.len());
137    const NONE: usize = usize::MAX;
138    for w in worlds {
139        let same =
140            |kept: &World| kept.constraints == w.constraints && live.iter().all(|&i| kept.slots[i] == w.slots[i]);
141        let found = match first.entry(state_hash(&w, &live)) {
142            Entry::Vacant(entry) => {
143                entry.insert(out.len());
144                None
145            }
146            Entry::Occupied(entry) => {
147                let mut i = *entry.get();
148                loop {
149                    if same(&out[i]) {
150                        break Some(i);
151                    }
152                    if next[i] == NONE {
153                        next[i] = out.len();
154                        break None;
155                    }
156                    i = next[i];
157                }
158            }
159        };
160        match found {
161            Some(i) => out[i].weight += w.weight,
162            None => {
163                next.push(NONE);
164                out.push(w);
165            }
166        }
167    }
168    out
169}
170
171/// Merge returned values only when their caller-visible restrictions also agree.
172pub fn merge_values(pairs: Vec<Returned>) -> Vec<Returned> {
173    let mut out: Vec<Returned> = Vec::with_capacity(pairs.len());
174    let mut index: FxHashMap<(Value, Constraints), usize> = FxHashMap::default();
175    for (v, w, constraints) in pairs {
176        match index.entry((v.clone(), constraints.clone())) {
177            Entry::Occupied(entry) => out[*entry.get()].1 += w,
178            Entry::Vacant(entry) => {
179                entry.insert(out.len());
180                out.push((v, w, constraints));
181            }
182        }
183    }
184    out
185}
186
187#[cfg(test)]
188mod tests {
189    use super::*;
190
191    fn world(slots: Vec<Value>) -> World {
192        World {
193            slots,
194            constraints: Default::default(),
195            inherited: Default::default(),
196            weight: Weight::ONE,
197            run: 0,
198        }
199    }
200
201    #[test]
202    fn worlds_merge_on_their_live_slots() {
203        let mut live = SlotSet::with_capacity(3);
204        live.insert(0);
205        // Slot 1 is dead: worlds that differ only there are the same world.
206        let worlds = vec![
207            world(vec![Value::Int(1.into()), Value::Int(7.into()), Value::Dead]),
208            world(vec![Value::Int(2.into()), Value::Int(7.into()), Value::Dead]),
209            world(vec![
210                Value::Int(1.into()),
211                Value::list(vec![Value::Int(8.into())]),
212                Value::Dead,
213            ]),
214        ];
215        let mut merged = merge(worlds, &live, true);
216        assert_eq!(merged.len(), 2);
217        assert_eq!(merged[0].slots[0], Value::Int(1.into()));
218        assert_eq!(merged[0].weight.to_f64(), 2.0);
219        assert_eq!(merged[1].weight.to_f64(), 1.0);
220        clear_dead(&mut merged, &live);
221        assert!(merged.iter().all(|w| w.slots[1] == Value::Dead));
222        assert_eq!(merged[1].slots[0], Value::Int(2.into()));
223    }
224}