結果
| 問題 | No.3762 Glowing Utility Pole |
| コンテスト | |
| ユーザー |
ei1333333
|
| 提出日時 | 2026-10-04 19:17:17 |
| 言語 | PyPy3 (7.3.23 + ACL) |
| 結果 |
AC
不安定
|
| 実行時間 | 488 ms / 2,000 ms |
| + 541µs | |
| コード長 | 2,911 bytes |
| 記録 | |
| コンパイル時間 | 67 ms |
| コンパイル使用メモリ | 83,628 KB |
| 実行使用メモリ | 132,140 KB |
| 最終ジャッジ日時 | 2026-10-09 20:52:59 |
| 合計ジャッジ時間 | 6,193 ms |
|
ジャッジサーバーID (参考情報) |
judge3_0 / judge1_0 |
(要ログイン)
| ファイルパターン | 結果 |
|---|---|
| sample | AC * 3 |
| other | AC * 47 |
ソースコード
import sys
from math import comb
MOD = 998244353
def solve(N, M, C):
T = C.count(0)
# 以下の配列の添字 j は、元コードの包除原理の j + 1 に対応。
# j = 0 の項は、各右端について r を加えることで別処理する。
# 最終出現位置を新しい順に並べたときの係数
coef = [
[
comb(M - rank - 1, j - 1) if j <= M - rank else 0
for j in range(1, M)
]
for rank in range(M)
]
top = [comb(M, j) for j in range(1, M)]
inv_m = pow(M, MOD - 2, MOD)
right_ratio = [
(M - j) * inv_m % MOD
for j in range(1, M)
]
left_ratio = [
M * pow(M - j, MOD - 2, MOD) % MOD
for j in range(1, M)
]
prefix = [0] * (M - 1)
left_weight = [1] * (M - 1)
right_weight = [1] * (M - 1)
# 各色の最終出現位置を新しい順に管理する。
# 位置そのものではなく、その位置での累積和を保存する。
recent_colors = []
recent_prefix = []
correction = [0] * (M - 1)
ans = 0
for r, c in enumerate(C, 1):
# 現在位置までの prefix_weight を更新
for j in range(M - 1):
s = prefix[j] + left_weight[j]
if s >= MOD:
s -= MOD
prefix[j] = s
if c == 0:
for j in range(M - 1):
left_weight[j] = (
left_weight[j] * left_ratio[j] % MOD
)
right_weight[j] = (
right_weight[j] * right_ratio[j] % MOD
)
else:
if c in recent_colors:
i = recent_colors.index(c)
del recent_colors[i]
del recent_prefix[i]
recent_colors.insert(0, c)
recent_prefix.insert(0, prefix[:])
# 固定色の最終出現位置が変わった場合だけ再計算
correction = [0] * (M - 1)
for rank, past in enumerate(recent_prefix):
cc = coef[rank]
for j in range(min(M - 1, M - rank)):
correction[j] += cc[j] * past[j]
for j in range(M - 1):
correction[j] %= MOD
# 元コードの j = 0 の項
ans += r
for j in range(M - 1):
term = (top[j] * prefix[j] - correction[j]) % MOD
term = term * right_weight[j] % MOD
# 元コードでは j + 1 番目の項なので符号に注意
if j & 1:
ans += term
else:
ans -= term
ans %= MOD
return ans * pow(M, T, MOD) % MOD
def main():
data = list(map(int, sys.stdin.buffer.read().split()))
if not data:
return
N, M = data[:2]
C = data[2:]
del data
print(solve(N, M, C))
if __name__ == "__main__":
main()
ei1333333