Use this to learn the idea, then write your own version.
12 3 4from typing import List5 6 7class UnionFind:8 """Represents a data structure that tracks a set of elements partitioned9 into a number of disjoint (non-overlapping) subsets.10 11 Landau notation: O(α(n)), where α(n) is the inverse Ackermann function.12 13 See:14 https:www.youtube.com/watch?v=zV3Ul2pA2Fw15 https:en.wikipedia.org/wiki/Disjoint-set_data_structure16 https:atcoder.jp/contests/abc120/submissions/444494217 https:atcoder.jp/contests/abc292/submissions/3941007518 https:github.com/not522/ac-library-python/blob/master/atcoder/dsu.py19 """20 21 def __init__(self, number_count: int) -> None:22 """23 Args:24 number_count: The size of elements (greater than 2).25 """26 self.number_count = number_count27 self.parent_numbers = [-1 for _ in range(number_count)]28 self.edge_count = [0 for _ in range(number_count)]29 self.group_count = number_count30 31 def find_root(self, number: int) -> int:32 """Follows the chain of parent pointers from number up the tree until33 it reaches a root element, whose parent is itself.34 Args:35 number: The trees id (0-index).36 37 Returns:38 The index of a root element.39 """40 if self.parent_numbers[number] < 0:41 return number42 43 self.parent_numbers[number] = self.find_root(self.parent_numbers[number])44 return self.parent_numbers[number]45 46 def get_group_size(self, number: int) -> int:47 """48 Args:49 number: The trees id (0-index).50 51 Returns:52 The size of group.53 """54 return -self.parent_numbers[self.find_root(number)]55 56 def is_same_group(self, number_x: int, number_y: int) -> bool:57 """Represents the roots of tree number_x and number_y are in the same58 group.59 Args:60 number_x: The trees x (0-index).61 number_y: The trees y (0-index).62 """63 return self.find_root(number_x) == self.find_root(number_y)64 65 def merge_if_needs(self, number_x: int, number_y: int) -> bool:66 """Uses find_root to determine the roots of the tree number_x and67 number_y belong to. If the roots are distinct, the trees are combined68 by attaching the roots of one to the root of the other.69 Args:70 number_x: The trees x (0-index).71 number_y: The trees y (0-index).72 """73 x = self.find_root(number_x)74 y = self.find_root(number_y)75 76 self.edge_count[x] += 177 78 if x == y:79 return False80 81 self.group_count -= 182 83 if self.parent_numbers[x] > self.parent_numbers[y]:84 x, y = y, x85 86 self.parent_numbers[x] += self.parent_numbers[y]87 self.parent_numbers[y] = x88 self.edge_count[x] += self.edge_count[y]89 return True90 91 def get_roots(self) -> List[int]:92 return [i for i, x in enumerate(self.parent_numbers) if x < 0]93 94 def get_groups(self) -> List[List[int]]:95 roots: List[int] = [self.find_root(i) for i in range(self.number_count)]96 groups: List[List[int]] = [[] for _ in range(self.number_count)]97 98 for i in range(self.number_count):99 groups[roots[i]].append(i)100 101 return list(filter(lambda g: g, groups))102 103 def get_edge_count(self, number: int) -> int:104 return self.edge_count[number]105 106 def get_group_count(self) -> int:107 return self.group_count108 109 110class UnionFind2D:111 """Extends UnionFind to two dimensions.112 113 See:114 https:atcoder.jp/contests/past202010-open/submissions/21472171115 """116 117 def __init__(self, height: int, width: int) -> None:118 self.height: int = height119 self.width: int = width120 self.size: int = height * width121 self.uf: UnionFind = UnionFind(self.size)122 123 def find_root(self, x: int, y: int) -> int:124 assert 0 <= x < self.width125 assert 0 <= y < self.height126 127 return self.uf.find_root(self._to_number(x, y))128 129 def get_group_size(self, x: int, y: int) -> int:130 assert 0 <= x < self.width131 assert 0 <= y < self.height132 133 return self.uf.get_group_size(self._to_number(x, y))134 135 def is_same_group(self, x1: int, y1: int, x2: int, y2: int) -> bool:136 assert 0 <= x1 < self.width137 assert 0 <= y1 < self.height138 assert 0 <= x2 < self.width139 assert 0 <= y2 < self.height140 141 return self.find_root(x1, y1) == self.find_root(x2, y2)142 143 def merge_if_needs(self, x1: int, y1: int, x2: int, y2: int) -> bool:144 assert 0 <= x1 < self.width145 assert 0 <= y1 < self.height146 assert 0 <= x2 < self.width147 assert 0 <= y2 < self.height148 149 return self.uf.merge_if_needs(self._to_number(x1, y1), self._to_number(x2, y2))150 151 def get_roots(self) -> List[int]:152 return self.uf.get_roots()153 154 def get_groups(self) -> List[List[int]]:155 """156 Returns:157 List of trees id (0-index).158 """159 return self.uf.get_groups()160 161 def get_edge_count(self, x: int, y: int) -> int:162 assert 0 <= x < self.width163 assert 0 <= y < self.height164 165 return self.uf.get_edge_count(self._to_number(x, y))166 167 def get_group_count(self) -> int:168 return self.uf.get_group_count()169 170 def _to_number(self, x: int, y: int) -> int:171 """172 Args:173 x, y: Coordinates in grid (0-index).174 175 Returns:176 The trees id (0-index).177 """178 return x + self.width * y179 180 def _to_yx(self, number: int) -> tuple[int, int]:181 """182 Args:183 The trees id (0-index).184 185 Returns:186 y, x: Coordinates in grid (0-index).187 """188 return divmod(number, self.width)189 190 191def main():192 import sys193 194 input = sys.stdin.readline195 196 h, w = map(int, input().split())197 s = [list(input().rstrip()) for _ in range(h)]198 199 uf = UnionFind2D(height=h, width=w)200 red_count = 0201 202 for y in range(h):203 for x in range(w):204 if s[y][x] == ".":205 red_count += 1206 continue207 208 if (y + 1 < h) and s[y + 1][x] == "#":209 uf.merge_if_needs(x, y, x, y + 1)210 if (x + 1 < w) and s[y][x + 1] == "#":211 uf.merge_if_needs(x, y, x + 1, y)212 213 dxy = [(-1, 0), (1, 0), (0, -1), (0, 1), (-1, -1), (1, -1), (-1, 1), (1, 1)]214 dxy = dxy[:4]215 mod = 998244353216 inv = pow(red_count, -1, mod) 217 group_count = uf.get_group_count() - red_count218 ans = 0219 220 for y in range(h):221 for x in range(w):222 if s[y][x] == "#":223 continue224 225 roots = set()226 227 for dx, dy in dxy:228 nx = x + dx229 ny = y + dy230 231 if not (0 <= nx < w):232 continue233 if not (0 <= ny < h):234 continue235 if s[ny][nx] == ".":236 continue237 238 root = uf.find_root(nx, ny)239 roots.add(root)240 241 ans += group_count - len(roots) + 1242 ans %= mod243 244 print(ans * inv % mod)245 246 247if __name__ == "__main__":248 main()249