m = 998244353 n, m_val = map(int, input().split()) a, b = 1, 1 for i in range(n): a = a * (n + m_val - i) % m b = b * (i + 1) % m cur_a = [a, 0] r = [cur_a] for i in range(1, 5001): nxt_a = [] for j in range(i + 1): t1 = cur_a[j - 1] * (i - j) * (n + m_val + j) % m if j > 0 else 0 t2 = cur_a[j] * (j + 1) * (m_val + j - i + 1) % m if j < len(cur_a) else 0 nxt_a.append((t1 + t2) % m) nxt_a.append(0) cur_a = nxt_a sm = sum(cur_a[:-1]) % m b = b * (n + i) % m r.append((sm * pow(b, -1, m)) % m) print(sum(r[int(i)] for i in input().split()) % m)