結果

問題 No.3753 Certainly a Cretan
コンテスト
ユーザー marc2825
提出日時 2026-08-19 01:50:08
言語 PyPy3
(7.3.23 + ACL)
コンパイル:
pypy3 -mpy_compile _filename_
実行:
pypy3 _filename_
結果
AC  
実行時間 439 ms / 2,500 ms
+ 360µs
コード長 3,475 bytes
記録
記録タグの例:
初AC ショートコード 純ショートコード 純主流ショートコード 最速実行時間
コンパイル時間 66 ms
コンパイル使用メモリ 82,964 KB
実行使用メモリ 160,492 KB
最終ジャッジ日時 2026-10-02 20:58:20
合計ジャッジ時間 8,665 ms
ジャッジサーバーID
(参考情報)
judge1_0 / judge2_0
このコードへのチャレンジ
(要ログイン)
ファイルパターン 結果
sample AC * 2
other AC * 46
権限があれば一括ダウンロードができます

ソースコード

diff #
raw source code

MOD = 998244353

N, Q = map(int, input().split())
S = list(input())

fact = [1] * (N + 1)
invfact = [1] * (N + 1)

for i in range(1, N + 1):
    fact[i] = fact[i - 1] * i % MOD

invfact[N] = pow(fact[N], MOD - 2, MOD)

for i in range(N, 0, -1):
    invfact[i - 1] = invfact[i] * i % MOD


def C(n, k):
    if k < 0 or k > n:
        return 0
    return fact[n] * invfact[k] % MOD * invfact[n - k] % MOD


def cat(r):
    return fact[2 * r] * invfact[r] % MOD * invfact[r + 1] % MOD


def invcat(r):
    return invfact[2 * r] * fact[r] % MOD * fact[r + 1] % MOD


# 位置 x が
# ・x = 0
# ・偶数位置の Y/N の切り替わり
# のどちらかなら 1 を持つ
bit = [0] * (N + 1)

bit[1] = 1
cnt = 1

bad = 0
P = 1
last = 0

for i in range(1, N):
    if S[i - 1] == S[i]:
        continue

    if i % 2 == 1:
        bad += 1
    else:
        bit[i + 1] = 1
        cnt += 1

        P = P * cat((i - last) // 2) % MOD
        last = i


# BIT を O(N) で構築
for i in range(1, N + 1):
    j = i + (i & -i)
    if j <= N:
        bit[j] += bit[i]


def add(pos, x):
    i = pos + 1

    while i <= N:
        bit[i] += x
        i += i & -i


def prefix_sum(pos):
    if pos < 0:
        return 0

    i = pos + 1
    res = 0

    while i > 0:
        res += bit[i]
        i -= i & -i

    return res


# 集合に入っている位置のうち、
# 左から k 個目の位置を返す
# k は 1-indexed
def kth(k):
    idx = 0
    d = 1 << (N.bit_length() - 1)

    while d > 0:
        nxt = idx + d

        if nxt <= N and bit[nxt] < k:
            idx = nxt
            k -= bit[nxt]

        d >>= 1

    return idx


def prev_boundary(x):
    k = prefix_sum(x - 1)
    return kth(k)


def next_boundary(x):
    k = prefix_sum(x)

    if k == cnt:
        return -1

    return kth(k + 1)


def insert_boundary(x):
    global P, cnt

    l = prev_boundary(x)
    r = next_boundary(x)

    if r == -1:
        P = P * cat((x - l) // 2) % MOD
    else:
        P = P * invcat((r - l) // 2) % MOD
        P = P * cat((x - l) // 2) % MOD
        P = P * cat((r - x) // 2) % MOD

    add(x, 1)
    cnt += 1


def erase_boundary(x):
    global P, cnt

    l = prev_boundary(x)
    r = next_boundary(x)

    if r == -1:
        P = P * invcat((x - l) // 2) % MOD
    else:
        P = P * invcat((x - l) // 2) % MOD
        P = P * invcat((r - x) // 2) % MOD
        P = P * cat((r - l) // 2) % MOD

    add(x, -1)
    cnt -= 1


def change_boundary(x):
    global bad

    if x <= 0 or x >= N:
        return

    diff = S[x - 1] != S[x]

    if x % 2 == 1:
        if diff:
            bad -= 1
        else:
            bad += 1

    else:
        if diff:
            erase_boundary(x)
        else:
            insert_boundary(x)


for _ in range(Q):
    t, x = map(int, input().split())

    if t == 1:
        change_boundary(x - 1)
        change_boundary(x)

        if S[x - 1] == 'Y':
            S[x - 1] = 'N'
        else:
            S[x - 1] = 'Y'

    else:
        K = x

        if bad > 0:
            print(0)
            continue

        B = kth(cnt)
        L = N - B
        q = K - B // 2

        if S[-1] == 'Y':
            if not (0 <= q <= L // 2):
                print(0)
                continue

            f = (C(L, q) - C(L, q - 1)) % MOD

        else:
            if not ((L + 1) // 2 <= q <= L):
                print(0)
                continue

            f = (C(L, q) - C(L, q + 1)) % MOD

        print(P * f % MOD)
0