Skip to main content

competitive/graph/
closure.rs

1use super::{DirectedGraph, Graph, Neighbor, VertexMap};
2use std::{collections::HashMap, hash::Hash, iter::Map, marker::PhantomData, ops::Range};
3
4pub struct UsizeGraph<Fa> {
5    vsize: usize,
6    adj: Fa,
7}
8impl<Fa> UsizeGraph<Fa> {
9    pub fn new(vsize: usize, adj: Fa) -> Self {
10        Self { vsize, adj }
11    }
12}
13
14impl<Fa, I, T> Graph for UsizeGraph<Fa>
15where
16    I: Iterator<Item = (usize, T)>,
17    Fa: Fn(usize) -> I,
18{
19    type Vertex = usize;
20    type Label = T;
21    type Vertices<'g>
22        = Range<usize>
23    where
24        Fa: 'g;
25    type Neighbors<'g>
26        = Map<I, fn((usize, T)) -> Neighbor<usize, T>>
27    where
28        Fa: 'g;
29
30    #[inline]
31    fn vsize(&self) -> usize {
32        self.vsize
33    }
34
35    #[inline]
36    fn vertices(&self) -> Self::Vertices<'_> {
37        0..self.vsize
38    }
39
40    #[inline]
41    fn neighbors(&self, vertex: Self::Vertex) -> Self::Neighbors<'_> {
42        (self.adj)(vertex).map(Into::into)
43    }
44}
45
46impl<Fa> DirectedGraph for UsizeGraph<Fa> where UsizeGraph<Fa>: Graph {}
47
48impl<Fa, T> VertexMap<T> for UsizeGraph<Fa>
49where
50    Self: Graph<Vertex = usize>,
51{
52    type Vmap = Vec<T>;
53    #[inline]
54    fn construct_vmap<F>(&self, f: F) -> Self::Vmap
55    where
56        F: FnMut() -> T,
57    {
58        let mut v = Vec::with_capacity(self.vsize);
59        v.resize_with(self.vsize, f);
60        v
61    }
62    #[inline]
63    fn vmap_get<'a>(&self, map: &'a Self::Vmap, vid: Self::Vertex) -> &'a T {
64        &map[vid]
65    }
66    #[inline]
67    fn vmap_get_mut<'a>(&self, map: &'a mut Self::Vmap, vid: Self::Vertex) -> &'a mut T {
68        &mut map[vid]
69    }
70}
71
72pub struct ClosureGraph<V, Fv, Fa> {
73    vs: Fv,
74    adj: Fa,
75    _marker: PhantomData<fn() -> V>,
76}
77
78impl<V, Fv, Fa> ClosureGraph<V, Fv, Fa> {
79    pub fn new(vs: Fv, adj: Fa) -> Self {
80        Self {
81            vs,
82            adj,
83            _marker: PhantomData,
84        }
85    }
86}
87
88impl<V, Fv, Fa, Iv, Ia, T> Graph for ClosureGraph<V, Fv, Fa>
89where
90    V: Eq + Copy,
91    Iv: Iterator<Item = V>,
92    Fv: Fn() -> Iv,
93    Ia: Iterator<Item = (V, T)>,
94    Fa: Fn(V) -> Ia,
95{
96    type Vertex = V;
97    type Label = T;
98    type Vertices<'g>
99        = Iv
100    where
101        V: 'g,
102        Fv: 'g,
103        Fa: 'g;
104    type Neighbors<'g>
105        = Map<Ia, fn((V, T)) -> Neighbor<V, T>>
106    where
107        V: 'g,
108        Fv: 'g,
109        Fa: 'g;
110
111    #[inline]
112    fn vsize(&self) -> usize {
113        (self.vs)().count()
114    }
115
116    #[inline]
117    fn vertices(&self) -> Self::Vertices<'_> {
118        (self.vs)()
119    }
120
121    #[inline]
122    fn neighbors(&self, vertex: Self::Vertex) -> Self::Neighbors<'_> {
123        (self.adj)(vertex).map(Into::into)
124    }
125}
126
127impl<V, Fv, Fa> DirectedGraph for ClosureGraph<V, Fv, Fa> where Self: Graph {}
128
129impl<V, Fv, Fa, T> VertexMap<T> for ClosureGraph<V, Fv, Fa>
130where
131    V: Eq + Copy + Hash,
132    T: Clone,
133    Self: Graph<Vertex = V>,
134{
135    type Vmap = (HashMap<V, T>, T);
136    #[inline]
137    fn construct_vmap<F>(&self, mut f: F) -> Self::Vmap
138    where
139        F: FnMut() -> T,
140    {
141        (HashMap::new(), f())
142    }
143    #[inline]
144    fn vmap_get<'a>(&self, (map, val): &'a Self::Vmap, vid: Self::Vertex) -> &'a T {
145        map.get(&vid).unwrap_or(val)
146    }
147    #[inline]
148    fn vmap_get_mut<'a>(&self, (map, val): &'a mut Self::Vmap, vid: Self::Vertex) -> &'a mut T {
149        map.entry(vid).or_insert_with(|| val.clone())
150    }
151}
152
153#[cfg(test)]
154mod tests {
155    use crate::{
156        graph::{ClosureGraph, GridGraph, ShortestPathExt, UsizeGraph, VertexMap},
157        num::Saturating,
158        tools::Xorshift,
159    };
160    use std::iter::repeat_with;
161
162    #[test]
163    fn closure_graph_sssp() {
164        let mut rng = Xorshift::default();
165        const A: u64 = 1_000_000_000;
166        for _ in 0..30 {
167            let h = rng.rand(8) as usize + 1;
168            let w = rng.rand(8) as usize + 1;
169
170            let weight: Vec<_> = repeat_with(|| Saturating(rng.rand(A - 1) + 1))
171                .take(8)
172                .collect();
173            let visitable: Vec<Vec<bool>> =
174                repeat_with(|| repeat_with(|| rng.gen_bool(0.8)).take(w).collect())
175                    .take(h)
176                    .collect();
177
178            let g = GridGraph::new_adj8(h, w);
179            let g1 = UsizeGraph::new(h * w, |u| {
180                g.adj8(g.unflat(u)).filter_map(|a| {
181                    if visitable[a.0.0][a.0.1] {
182                        Some((g.flat(a.0), a.1))
183                    } else {
184                        None
185                    }
186                })
187            });
188            let g2 = ClosureGraph::new(
189                || {
190                    (0..h)
191                        .flat_map(|i| (0..w).map(move |j| (i, j)))
192                        .filter(|&(i, j)| visitable[i][j])
193                },
194                |u| g.adj8(u).filter(|&((i, j), _)| visitable[i][j]),
195            );
196            for (i, visitable) in visitable.iter().enumerate() {
197                for (j, visitable) in visitable.iter().enumerate() {
198                    assert_eq!((i, j), g.unflat(g.flat((i, j))));
199                    if !visitable {
200                        continue;
201                    }
202                    let cost1 = g1
203                        .standard_sp_additive()
204                        .dijkstra([g.flat((i, j))], |dir| weight[dir as usize]);
205                    let cost2 = g2
206                        .standard_sp_additive()
207                        .dijkstra([(i, j)], |dir| weight[dir as usize]);
208                    for ni in 0..h {
209                        for nj in 0..w {
210                            assert_eq!(
211                                g1.vmap_get(&cost1, g.flat((ni, nj))),
212                                g2.vmap_get(&cost2, (ni, nj))
213                            );
214                        }
215                    }
216                }
217            }
218        }
219    }
220
221    #[test]
222    fn closure_graph_apsp() {
223        let mut rng = Xorshift::default();
224        const A: u64 = 1_000_000_000;
225        for _ in 0..30 {
226            let h = rng.rand(8) as usize + 1;
227            let w = rng.rand(8) as usize + 1;
228
229            let weight: Vec<_> = repeat_with(|| Saturating(rng.rand(A - 1) + 1))
230                .take(8)
231                .collect();
232
233            let g = GridGraph::new_adj4(h, w);
234            let cost: Vec<Vec<Vec<_>>> = (0..h)
235                .map(|i| {
236                    (0..w)
237                        .map(|j| {
238                            g.standard_sp_additive()
239                                .dijkstra([(i, j)], |dir| weight[dir as usize])
240                        })
241                        .collect()
242                })
243                .collect();
244            let g2 = ClosureGraph::new(
245                || (0..h).flat_map(|i| (0..w).map(move |j| (i, j))),
246                |u| g.adj4(u),
247            );
248            let cost2 = g2
249                .standard_sp_additive()
250                .warshall_floyd_ap(|dir| weight[dir as usize]);
251            for (i, row) in cost.iter().enumerate() {
252                for (j, source_cost) in row.iter().enumerate() {
253                    for ni in 0..h {
254                        for nj in 0..w {
255                            assert_eq!(
256                                g.vmap_get(source_cost, (ni, nj)),
257                                g2.vmap_get(g2.vmap_get(&cost2, (i, j)), (ni, nj))
258                            );
259                        }
260                    }
261                }
262            }
263        }
264    }
265}