Skip to main content

competitive/graph/
grid.rs

1use super::{Graph, Neighbor, VertexMap};
2use std::{iter::Map, marker::PhantomData, ops::Range};
3
4#[derive(Debug, Clone, Copy)]
5pub struct GridGraph<A> {
6    pub height: usize,
7    pub width: usize,
8    _marker: PhantomData<fn() -> A>,
9}
10
11impl GridGraph<Adj4> {
12    pub fn new_adj4(height: usize, width: usize) -> Self {
13        Self::new(height, width)
14    }
15    pub fn adj4(&self, vid: (usize, usize)) -> GridAdjacency<'_, Adj4> {
16        GridAdjacency {
17            g: self,
18            xy: vid,
19            diter: GridDirectionIter::default(),
20            _marker: PhantomData,
21        }
22    }
23}
24impl GridGraph<Adj8> {
25    pub fn new_adj8(height: usize, width: usize) -> Self {
26        Self::new(height, width)
27    }
28    pub fn adj8(&self, vid: (usize, usize)) -> GridAdjacency<'_, Adj8> {
29        GridAdjacency {
30            g: self,
31            xy: vid,
32            diter: GridDirectionIter::default(),
33            _marker: PhantomData,
34        }
35    }
36}
37
38impl<A> GridGraph<A> {
39    pub fn new(height: usize, width: usize) -> Self {
40        Self {
41            height,
42            width,
43            _marker: PhantomData,
44        }
45    }
46    #[inline]
47    pub fn move_by_diff(&self, xy: (usize, usize), dxdy: (isize, isize)) -> Option<(usize, usize)> {
48        let nx = xy.0.wrapping_add(dxdy.0 as usize);
49        let ny = xy.1.wrapping_add(dxdy.1 as usize);
50        if nx < self.height && ny < self.width {
51            Some((nx, ny))
52        } else {
53            None
54        }
55    }
56    #[inline]
57    pub fn flat(&self, xy: (usize, usize)) -> usize {
58        xy.0 * self.width + xy.1
59    }
60    #[inline]
61    pub fn unflat(&self, pos: usize) -> (usize, usize) {
62        (pos / self.width, pos % self.width)
63    }
64}
65
66impl<A> Graph for GridGraph<A>
67where
68    GridDirectionIter<A>: Iterator<Item = GridDirection>,
69{
70    type Vertex = (usize, usize);
71    type Label = GridDirection;
72    type Vertices<'g>
73        = GridVertices
74    where
75        A: 'g;
76    type Neighbors<'g>
77        = Map<
78        GridAdjacency<'g, A>,
79        fn(((usize, usize), GridDirection)) -> Neighbor<(usize, usize), GridDirection>,
80    >
81    where
82        A: 'g;
83
84    #[inline]
85    fn vsize(&self) -> usize {
86        self.height * self.width
87    }
88
89    #[inline]
90    fn vertices(&self) -> Self::Vertices<'_> {
91        GridVertices {
92            xrange: 0..self.height,
93            yrange: 0..self.width,
94        }
95    }
96
97    #[inline]
98    fn neighbors(&self, vertex: Self::Vertex) -> Self::Neighbors<'_> {
99        GridAdjacency {
100            g: self,
101            xy: vertex,
102            diter: GridDirectionIter::default(),
103            _marker: PhantomData,
104        }
105        .map(Into::into)
106    }
107}
108
109#[derive(Debug, Clone)]
110pub struct GridVertices {
111    xrange: Range<usize>,
112    yrange: Range<usize>,
113}
114
115impl Iterator for GridVertices {
116    type Item = (usize, usize);
117    fn next(&mut self) -> Option<Self::Item> {
118        loop {
119            if self.xrange.start >= self.xrange.end {
120                return None;
121            }
122            if let Some(ny) = self.yrange.next() {
123                return Some((self.xrange.start, ny));
124            }
125            self.yrange.start = 0;
126            self.xrange.start += 1;
127        }
128    }
129}
130
131#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
132pub enum GridDirection {
133    U = 0isize,
134    L = 1isize,
135    R = 2isize,
136    D = 3isize,
137    UL = 4isize,
138    UR = 5isize,
139    DL = 6isize,
140    DR = 7isize,
141}
142
143impl GridDirection {
144    pub fn dxdy(self) -> (isize, isize) {
145        match self {
146            GridDirection::U => (-1, 0),
147            GridDirection::L => (0, -1),
148            GridDirection::R => (0, 1),
149            GridDirection::D => (1, 0),
150            GridDirection::UL => (-1, -1),
151            GridDirection::UR => (-1, 1),
152            GridDirection::DL => (1, -1),
153            GridDirection::DR => (1, 1),
154        }
155    }
156    pub fn ndxdy(self, d: usize) -> (isize, isize) {
157        let d = d as isize;
158        match self {
159            GridDirection::U => (-d, 0),
160            GridDirection::L => (0, -d),
161            GridDirection::R => (0, d),
162            GridDirection::D => (d, 0),
163            GridDirection::UL => (-d, -d),
164            GridDirection::UR => (-d, d),
165            GridDirection::DL => (d, -d),
166            GridDirection::DR => (d, d),
167        }
168    }
169}
170
171#[derive(Debug, Clone, Copy)]
172pub enum Adj4 {}
173#[derive(Debug, Clone, Copy)]
174pub enum Adj8 {}
175
176#[derive(Debug, Clone)]
177pub struct GridDirectionIter<A> {
178    dir: Option<GridDirection>,
179    _marker: PhantomData<fn() -> A>,
180}
181impl<A> Default for GridDirectionIter<A> {
182    fn default() -> Self {
183        Self {
184            dir: Some(GridDirection::U),
185            _marker: PhantomData,
186        }
187    }
188}
189
190impl Iterator for GridDirectionIter<Adj4> {
191    type Item = GridDirection;
192    fn next(&mut self) -> Option<Self::Item> {
193        if let Some(dir) = &mut self.dir {
194            let cdir = Some(*dir);
195            self.dir = match dir {
196                GridDirection::U => Some(GridDirection::L),
197                GridDirection::L => Some(GridDirection::R),
198                GridDirection::R => Some(GridDirection::D),
199                _ => None,
200            };
201            cdir
202        } else {
203            None
204        }
205    }
206}
207impl Iterator for GridDirectionIter<Adj8> {
208    type Item = GridDirection;
209    fn next(&mut self) -> Option<Self::Item> {
210        if let Some(dir) = &mut self.dir {
211            let cdir = Some(*dir);
212            self.dir = match dir {
213                GridDirection::U => Some(GridDirection::L),
214                GridDirection::L => Some(GridDirection::R),
215                GridDirection::R => Some(GridDirection::D),
216                GridDirection::D => Some(GridDirection::UL),
217                GridDirection::UL => Some(GridDirection::UR),
218                GridDirection::UR => Some(GridDirection::DL),
219                GridDirection::DL => Some(GridDirection::DR),
220                GridDirection::DR => None,
221            };
222            cdir
223        } else {
224            None
225        }
226    }
227}
228
229#[derive(Debug, Clone)]
230pub struct GridAdjacency<'g, A> {
231    g: &'g GridGraph<A>,
232    xy: (usize, usize),
233    diter: GridDirectionIter<A>,
234    _marker: PhantomData<fn() -> A>,
235}
236
237impl<A> Iterator for GridAdjacency<'_, A>
238where
239    GridDirectionIter<A>: Iterator<Item = GridDirection>,
240{
241    type Item = ((usize, usize), GridDirection);
242    fn next(&mut self) -> Option<Self::Item> {
243        for dir in self.diter.by_ref() {
244            match self.g.move_by_diff(self.xy, dir.dxdy()) {
245                Some(nxy) => return Some((nxy, dir)),
246                None => continue,
247            }
248        }
249        None
250    }
251}
252
253impl<A, T> VertexMap<T> for GridGraph<A>
254where
255    Self: Graph<Vertex = (usize, usize)>,
256{
257    type Vmap = Vec<T>;
258
259    #[inline]
260    fn construct_vmap<F>(&self, f: F) -> Self::Vmap
261    where
262        F: FnMut() -> T,
263    {
264        let mut map = Vec::with_capacity(self.height * self.width);
265        map.resize_with(self.height * self.width, f);
266        map
267    }
268
269    #[inline]
270    fn vmap_get<'a>(&self, map: &'a Self::Vmap, (x, y): Self::Vertex) -> &'a T {
271        assert!(x < self.height, "expected 0..{}, but {}", self.height, x);
272        assert!(y < self.width, "expected 0..{}, but {}", self.width, y);
273        &map[x * self.width + y]
274    }
275
276    #[inline]
277    fn vmap_get_mut<'a>(&self, map: &'a mut Self::Vmap, (x, y): Self::Vertex) -> &'a mut T {
278        assert!(x < self.height, "expected 0..{}, but {}", self.height, x);
279        assert!(y < self.width, "expected 0..{}, but {}", self.width, y);
280        &mut map[x * self.width + y]
281    }
282}
283
284#[cfg(test)]
285mod tests {
286    use super::GridGraph;
287    use crate::{
288        graph::{ShortestPathExt, VertexMap},
289        num::Saturating,
290        tools::Xorshift,
291    };
292
293    #[test]
294    fn grid_graph_apsp() {
295        let mut rng = Xorshift::default();
296        const A: u64 = 1_000_000_000;
297        for _ in 0..30 {
298            let h = rng.rand(8) as usize + 1;
299            let w = rng.rand(8) as usize + 1;
300
301            let weight: Vec<_> = std::iter::repeat_with(|| Saturating(rng.rand(A - 1) + 1))
302                .take(8)
303                .collect();
304
305            let g = GridGraph::new_adj4(h, w);
306            let cost: Vec<Vec<Vec<_>>> = (0..h)
307                .map(|i| {
308                    (0..w)
309                        .map(|j| {
310                            g.standard_sp_additive()
311                                .dijkstra([(i, j)], |dir| weight[dir as usize])
312                        })
313                        .collect()
314                })
315                .collect();
316            let cost2: Vec<Vec<_>> = g
317                .standard_sp_additive()
318                .warshall_floyd_ap(|dir| weight[dir as usize]);
319            for (i, row) in cost.iter().enumerate() {
320                for (j, source_cost) in row.iter().enumerate() {
321                    for ni in 0..h {
322                        for nj in 0..w {
323                            assert_eq!(
324                                g.vmap_get(source_cost, (ni, nj)),
325                                g.vmap_get(g.vmap_get(&cost2, (i, j)), (ni, nj))
326                            );
327                        }
328                    }
329                }
330            }
331
332            let g = GridGraph::new_adj8(h, w);
333            let cost: Vec<Vec<Vec<_>>> = (0..h)
334                .map(|i| {
335                    (0..w)
336                        .map(|j| {
337                            g.standard_sp_additive()
338                                .dijkstra([(i, j)], |dir| weight[dir as usize])
339                        })
340                        .collect()
341                })
342                .collect();
343            let cost2: Vec<Vec<_>> = g
344                .standard_sp_additive()
345                .warshall_floyd_ap(|dir| weight[dir as usize]);
346            for (i, row) in cost.iter().enumerate() {
347                for (j, source_cost) in row.iter().enumerate() {
348                    for ni in 0..h {
349                        for nj in 0..w {
350                            assert_eq!(
351                                g.vmap_get(source_cost, (ni, nj)),
352                                g.vmap_get(g.vmap_get(&cost2, (i, j)), (ni, nj))
353                            );
354                        }
355                    }
356                }
357            }
358        }
359    }
360}