Skip to main content

probl_engine/
chain.rs

1//! Solving loops as absorbing Markov chains (docs/semantics.md, section 10).
2//!
3//! A loop's states are its worlds at the loop's start, told apart by the
4//! variables read later. Running the body once from each state gives the
5//! chance of going on to each state, and of leaving the loop. The expected
6//! number of times each state is visited, from the weights the loop starts
7//! with, then says how much weight leaves by each way out.
8//!
9//! States are solved one strongly connected group at a time, in the order
10//! weight flows through them. Within a group, states are eliminated one by
11//! one (Grassmann, Taksar and Heyman, 1985). Every step adds nonnegative
12//! numbers: the chance of leaving a state is the sum of its ways out, never
13//! 1 minus the chance of staying, so a chain whose exits are rare loses no
14//! precision.
15
16use rustc_hash::{FxHashMap, FxHashSet};
17use std::cmp::Reverse;
18use std::collections::BinaryHeap;
19
20/// One step of a chain from each of its states.
21#[derive(Clone, Debug, Default)]
22pub struct Chain {
23    /// For each state: the chance of going on to each state, itself
24    /// included, each state at most once.
25    pub next: Vec<Vec<(usize, f64)>>,
26    /// For each state: the chance of leaving the chain in one step, by any
27    /// way out.
28    pub leave: Vec<f64>,
29}
30
31#[derive(Clone, Debug, PartialEq)]
32pub enum Solution {
33    /// The expected number of visits to each state.
34    Visits(Vec<f64>),
35    /// A state from which the chain can never be left.
36    Stuck(usize),
37    /// Solving would take more steps than allowed.
38    TooBig,
39}
40
41impl Chain {
42    /// The expected number of visits to each state, starting from the
43    /// chances in `start`, in at most `budget` elimination steps, which are
44    /// taken from it. Every state must be reachable from `start`.
45    pub fn visits(&self, start: &[f64], budget: &mut u64) -> Solution {
46        let n = self.next.len();
47        let mut inflow = start.to_vec();
48        let mut visits = vec![0.0; n];
49        // Tarjan's algorithm finds the groups in reverse order of flow.
50        let mut groups = strongly_connected(&self.next);
51        groups.reverse();
52        let mut group_of = vec![0; n];
53        for (g, states) in groups.iter().enumerate() {
54            for &s in states {
55                group_of[s] = g;
56            }
57        }
58        for (g, states) in groups.iter().enumerate() {
59            let solved = if let [s] = states[..] {
60                // One state: it leaves to later groups, or comes back to itself.
61                let away: f64 = self.next[s].iter().filter(|&&(j, _)| j != s).map(|(_, p)| p).sum();
62                let d = self.leave[s] + away;
63                if d <= 0.0 {
64                    return Solution::Stuck(s);
65                }
66                visits[s] = inflow[s] / d;
67                true
68            } else {
69                match self.eliminate(states, g, &group_of, &inflow, &mut visits, budget) {
70                    Ok(()) => true,
71                    Err(solution) => return solution,
72                }
73            };
74            debug_assert!(solved);
75            // What leaves the group flows into later ones.
76            for &s in states {
77                for &(j, p) in &self.next[s] {
78                    if group_of[j] != g {
79                        inflow[j] += visits[s] * p;
80                    }
81                }
82            }
83        }
84        if visits.iter().any(|v| !v.is_finite()) {
85            return Solution::TooBig;
86        }
87        Solution::Visits(visits)
88    }
89
90    /// Solve one strongly connected group by eliminating its states, fewest
91    /// connections first.
92    fn eliminate(
93        &self,
94        states: &[usize],
95        group: usize,
96        group_of: &[usize],
97        inflow: &[f64],
98        visits: &mut [f64],
99        budget: &mut u64,
100    ) -> Result<(), Solution> {
101        let m = states.len();
102        let local: FxHashMap<usize, usize> = states.iter().enumerate().map(|(i, &s)| (s, i)).collect();
103        // Within the group: steps to its states, and the chance of leaving
104        // it, including to later groups.
105        let mut out: Vec<FxHashMap<usize, f64>> = vec![FxHashMap::default(); m];
106        let mut into: Vec<FxHashSet<usize>> = vec![FxHashSet::default(); m];
107        let mut leave = vec![0.0; m];
108        let mut mass = vec![0.0; m];
109        for (i, &s) in states.iter().enumerate() {
110            leave[i] = self.leave[s];
111            mass[i] = inflow[s];
112            for &(j, p) in &self.next[s] {
113                if group_of[j] == group {
114                    let j = local[&j];
115                    *out[i].entry(j).or_insert(0.0) += p;
116                    if j != i {
117                        into[j].insert(i);
118                    }
119                } else {
120                    leave[i] += p;
121                }
122            }
123        }
124        let degree = |i: usize, out: &[FxHashMap<usize, f64>], into: &[FxHashSet<usize>]| {
125            (into[i].len() as u64) * (out[i].len() as u64)
126        };
127        let mut queue: BinaryHeap<Reverse<(u64, usize)>> =
128            (0..m).map(|i| Reverse((degree(i, &out, &into), i))).collect();
129        let mut alive = vec![true; m];
130        // Each state eliminated, in order, to work out the visits to them
131        // afterwards in reverse order.
132        let mut steps: Vec<Eliminated> = Vec::with_capacity(m);
133        while let Some(Reverse((d, k))) = queue.pop() {
134            if !alive[k] || d != degree(k, &out, &into) {
135                continue;
136            }
137            alive[k] = false;
138            // Stays at k don't count: they only make each visit longer.
139            out[k].remove(&k);
140            let away: f64 = out[k].values().sum();
141            let d_k = leave[k] + away;
142            if d_k <= 0.0 {
143                return Err(Solution::Stuck(states[k]));
144            }
145            let succs: Vec<(usize, f64)> = out[k].iter().map(|(&j, &p)| (j, p)).collect();
146            let preds: Vec<(usize, f64)> = into[k]
147                .iter()
148                .map(|&i| (i, out[i].remove(&k).expect("an edge into k")))
149                .collect();
150            let cost = (preds.len() * succs.len() + 1) as u64;
151            if cost > *budget {
152                return Err(Solution::TooBig);
153            }
154            *budget -= cost;
155            // Paths through k: i → k, then k's stays, then k → j or out.
156            for &(i, p_ik) in &preds {
157                let f = p_ik / d_k;
158                for &(j, p_kj) in &succs {
159                    *out[i].entry(j).or_insert(0.0) += f * p_kj;
160                    if j != i {
161                        into[j].insert(i);
162                    }
163                }
164                leave[i] += f * leave[k];
165            }
166            // Weight starting at k goes on to its successors.
167            for &(j, p_kj) in &succs {
168                mass[j] += mass[k] * p_kj / d_k;
169                into[j].remove(&k);
170            }
171            for &(i, _) in &preds {
172                queue.push(Reverse((degree(i, &out, &into), i)));
173            }
174            for &(j, _) in &succs {
175                queue.push(Reverse((degree(j, &out, &into), j)));
176            }
177            steps.push(Eliminated {
178                state: k,
179                mass: mass[k],
180                leaving: d_k,
181                from: preds,
182            });
183        }
184        // Visits to k: the weight that starts there or arrives from a state
185        // eliminated after it, times the expected stays, 1 / d_k.
186        let mut local_visits = vec![0.0; m];
187        for step in steps.into_iter().rev() {
188            let arriving: f64 = step.from.iter().map(|&(i, p)| local_visits[i] * p).sum();
189            local_visits[step.state] = (step.mass + arriving) / step.leaving;
190        }
191        for (i, &s) in states.iter().enumerate() {
192            visits[s] = local_visits[i];
193        }
194        Ok(())
195    }
196}
197
198/// A state as it was eliminated: what the visits to it are made of.
199struct Eliminated {
200    state: usize,
201    /// The weight starting there, including what reached it through states
202    /// eliminated before it.
203    mass: f64,
204    /// The chance of leaving it for another state, or out, each visit.
205    leaving: f64,
206    /// The states still there that could reach it, and with what chance.
207    from: Vec<(usize, f64)>,
208}
209
210/// The strongly connected groups of states, each group's states in no
211/// particular order, and the groups in reverse topological order (Tarjan's
212/// algorithm, without recursion).
213fn strongly_connected(next: &[Vec<(usize, f64)>]) -> Vec<Vec<usize>> {
214    const UNSEEN: usize = usize::MAX;
215    let n = next.len();
216    let mut index = vec![UNSEEN; n];
217    let mut low = vec![0; n];
218    let mut on_stack = vec![false; n];
219    let mut stack = Vec::new();
220    let mut groups = Vec::new();
221    let mut counter = 0;
222    // (state, position in its list of successors)
223    let mut calls: Vec<(usize, usize)> = Vec::new();
224    for root in 0..n {
225        if index[root] != UNSEEN {
226            continue;
227        }
228        calls.push((root, 0));
229        while let Some(&mut (v, ref mut at)) = calls.last_mut() {
230            if *at == 0 {
231                index[v] = counter;
232                low[v] = counter;
233                counter += 1;
234                stack.push(v);
235                on_stack[v] = true;
236            }
237            if let Some(&(w, _)) = next[v].get(*at) {
238                *at += 1;
239                if index[w] == UNSEEN {
240                    calls.push((w, 0));
241                } else if on_stack[w] {
242                    low[v] = low[v].min(index[w]);
243                }
244                continue;
245            }
246            calls.pop();
247            if let Some(&(parent, _)) = calls.last() {
248                low[parent] = low[parent].min(low[v]);
249            }
250            if low[v] == index[v] {
251                let mut group = Vec::new();
252                loop {
253                    let w = stack.pop().expect("v is on the stack");
254                    on_stack[w] = false;
255                    group.push(w);
256                    if w == v {
257                        break;
258                    }
259                }
260                groups.push(group);
261            }
262        }
263    }
264    groups
265}
266
267#[cfg(test)]
268mod tests {
269    use super::*;
270
271    fn close(a: f64, b: f64, tol: f64) {
272        assert!((a - b).abs() <= tol * b.abs().max(1.0), "{a} vs {b}");
273    }
274
275    fn visits(chain: &Chain, start: &[f64]) -> Vec<f64> {
276        match chain.visits(start, &mut u64::MAX.clone()) {
277            Solution::Visits(v) => v,
278            other => panic!("{other:?}"),
279        }
280    }
281
282    /// The expected visits by summing the chain's steps until they vanish.
283    fn by_iterating(chain: &Chain, start: &[f64]) -> Vec<f64> {
284        let n = chain.next.len();
285        let (mut total, mut now) = (start.to_vec(), start.to_vec());
286        for _ in 0..200_000 {
287            let mut after = vec![0.0; n];
288            for (i, row) in chain.next.iter().enumerate() {
289                for &(j, p) in row {
290                    after[j] += now[i] * p;
291                }
292            }
293            if after.iter().sum::<f64>() < 1e-18 {
294                break;
295            }
296            for (t, a) in total.iter_mut().zip(&after) {
297                *t += a;
298            }
299            now = after;
300        }
301        total
302    }
303
304    #[test]
305    fn one_state_that_stays() {
306        let chain = Chain {
307            next: vec![vec![(0, 0.75)]],
308            leave: vec![0.25],
309        };
310        assert_eq!(visits(&chain, &[1.0]), [4.0]);
311    }
312
313    #[test]
314    fn gamblers_ruin_is_exact() {
315        // A fair walk on 1..=9, leaving at 0 or 10: from k, the chance of
316        // reaching 10 is k / 10, and the expected visits to j from k are
317        // 2 min(j, k) (10 − max(j, k)) / 10.
318        let n = 9;
319        let next: Vec<Vec<(usize, f64)>> = (0..n)
320            .map(|i| {
321                let mut row = Vec::new();
322                if i > 0 {
323                    row.push((i - 1, 0.5));
324                }
325                if i + 1 < n {
326                    row.push((i + 1, 0.5));
327                }
328                row
329            })
330            .collect();
331        let leave = (0..n).map(|i| if i == 0 || i == n - 1 { 0.5 } else { 0.0 }).collect();
332        let chain = Chain { next, leave };
333        for k in 1..=9 {
334            let mut start = vec![0.0; n];
335            start[k - 1] = 1.0;
336            let v = visits(&chain, &start);
337            for j in 1..=9 {
338                let expected = 2.0 * (j.min(k) * (10 - j.max(k))) as f64 / 10.0;
339                close(v[j - 1], expected, 1e-13);
340            }
341            // Reaching 10 is leaving from 9.
342            close(v[n - 1] * 0.5, k as f64 / 10.0, 1e-13);
343        }
344    }
345
346    #[test]
347    fn rare_exits_lose_no_precision() {
348        // Two states that pass the weight back and forth, leaving with a
349        // chance of 1e-15 each step: 1e15 visits, which 1 − (chance of
350        // staying) couldn't compute.
351        let chain = Chain {
352            next: vec![vec![(1, 1.0 - 1e-15)], vec![(0, 1.0)]],
353            leave: vec![1e-15, 0.0],
354        };
355        let v = visits(&chain, &[1.0, 0.0]);
356        close(v[0], 1e15, 1e-12);
357        close(v[0] * 1e-15, 1.0, 1e-12);
358    }
359
360    #[test]
361    fn chains_that_cannot_be_left_are_found() {
362        // 0 leads to the cycle 1 ⇄ 2, which never leaves.
363        let chain = Chain {
364            next: vec![vec![(1, 0.5)], vec![(2, 1.0)], vec![(1, 1.0)]],
365            leave: vec![0.5, 0.0, 0.0],
366        };
367        assert!(matches!(
368            chain.visits(&[1.0, 0.0, 0.0], &mut u64::MAX.clone()),
369            Solution::Stuck(1 | 2)
370        ));
371        let chain = Chain {
372            next: vec![vec![(0, 1.0)]],
373            leave: vec![0.0],
374        };
375        assert_eq!(chain.visits(&[1.0], &mut u64::MAX.clone()), Solution::Stuck(0));
376    }
377
378    #[test]
379    fn a_budget_stops_large_eliminations() {
380        let n = 50;
381        let next: Vec<Vec<(usize, f64)>> = (0..n)
382            .map(|i| (0..n).map(|j| (j, 0.9 / n as f64)).filter(|&(j, _)| j != i).collect())
383            .collect();
384        let chain = Chain {
385            next,
386            leave: vec![0.1 + 0.9 / 50.0; n],
387        };
388        let mut start = vec![0.0; n];
389        start[0] = 1.0;
390        assert_eq!(chain.visits(&start, &mut 1000), Solution::TooBig);
391        assert!(matches!(
392            chain.visits(&start, &mut u64::MAX.clone()),
393            Solution::Visits(_)
394        ));
395    }
396
397    /// Random chains with several groups, cycles and self-loops, against
398    /// summing their steps.
399    #[test]
400    fn random_chains_agree_with_iterating() {
401        let mut rng = crate::continuous::Rng::new(5);
402        for _ in 0..200 {
403            let n = 1 + (rng.uniform() * 12.0) as usize;
404            let mut next = Vec::new();
405            let mut leave = Vec::new();
406            for _ in 0..n {
407                let mut row: Vec<(usize, f64)> = Vec::new();
408                let mut weights = Vec::new();
409                for j in 0..n {
410                    if rng.uniform() < 0.3 {
411                        row.push((j, 0.0));
412                        weights.push(rng.uniform());
413                    }
414                }
415                let out = rng.uniform() * 0.5 + 0.05;
416                weights.push(out);
417                let total: f64 = weights.iter().sum();
418                for ((_, p), w) in row.iter_mut().zip(&weights) {
419                    *p = w / total;
420                }
421                next.push(row);
422                leave.push(out / total);
423            }
424            let chain = Chain { next, leave };
425            let start: Vec<f64> = (0..n).map(|_| rng.uniform()).collect();
426            let solved = visits(&chain, &start);
427            let iterated = by_iterating(&chain, &start);
428            for (a, b) in solved.iter().zip(&iterated) {
429                close(*a, *b, 1e-9);
430            }
431        }
432    }
433}