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}