1use rustc_hash::{FxHashMap, FxHashSet};
17use std::cmp::Reverse;
18use std::collections::BinaryHeap;
19
20#[derive(Clone, Debug, Default)]
22pub struct Chain {
23 pub next: Vec<Vec<(usize, f64)>>,
26 pub leave: Vec<f64>,
29}
30
31#[derive(Clone, Debug, PartialEq)]
32pub enum Solution {
33 Visits(Vec<f64>),
35 Stuck(usize),
37 TooBig,
39}
40
41impl Chain {
42 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 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 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 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 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 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 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 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 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 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 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
198struct Eliminated {
200 state: usize,
201 mass: f64,
204 leaving: f64,
206 from: Vec<(usize, f64)>,
208}
209
210fn 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 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 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 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 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 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 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 #[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}