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