mod = 998244353 def main(): h, w = list(map(int, input().split())) S = [input() for _ in range(h)] X = [[0, 0, 0, 0] for _ in range(h)] for y in range(h): dp = [0] * 4 dp[0] = 1 for x in range(w): ndp = [0] * 4 if S[y][x] != "B": for bit in range(4): nbit = bit * 2 if 4 <= nbit: nbit ^= 3 nbit = nbit % 4 ndp[nbit] = (ndp[nbit] + dp[bit]) % mod if S[y][x] != "W": for bit in range(4): nbit = (bit^2) * 2 if 4 <= nbit: nbit ^= 3 nbit = nbit % 4 ndp[nbit] = (ndp[nbit] + dp[bit]) % mod dp = ndp X[y] = dp dp = [0] * 4 dp[0] = 1 for y in range(h): ndp = [0] * 4 for i in range(4): for j in range(4): k = i^j if k == 0: v = 0 elif k == 1: v = 2 elif k == 2: v = 3 else: v = 1 ndp[v] = (ndp[v] + dp[i]*X[y][j]) % mod dp = ndp return dp[0] print(main())