F - Random Vault Heist 解説
by
harurun4635
このような 条件を満たさなくなるまで、 ランダムに残りから選ぶ という設定の問題では、 ランダムに一列にならべて、前から条件を満たさなくなるまで選ぶ と考えるとわかりやすいかもしれません。
以下金庫の集合を \(T\) として \(A_{\text{all}} = \sum_{i \in T} A_i \) とします。
強盗が盗む条件を
- 今盗んだ金庫の集合を \(S\) としたとき、 \(\sum_{i \in S} A_i \lt X\) であれば次の金庫を開ける
と言い換えてみましょう。すると、求める期待値は以下のように言い換えることができます。
すべての \(\sum_{i \in S} A_i \lt X\) である \(S\) について、
- 「集合 \(S\) を盗んだ」という状態に到達する確率 \(\times\) その次に盗む金額の期待値
の総和
前者は「\(N\) 個の金庫を一列に並べた時、特定の \(|S|\) 個が prefix である確率」ですから \(\displaystyle \frac{1}{\binom{N}{|S|}}\) です。
後者は、残った \(S\) 以外の金庫から一様ランダムに選ばれますから \(\displaystyle \frac{A_{\text{all}} - \sum_{i \in S} A_i}{N - |S|}\) です。
\(\sum_{i \in S} A_i \lt X\) であるすべての \(S\) について上の値を求めたいですが、\(S\) は最大 \(2^N\) 通り存在し、実行制限時間に間に合いません。
しかし、\(|S| = k\) が同じであれば、分母の値がすべて同じです。よって、すべての \(k\) について
- 条件を満たすサイズ \(k\) の集合が何個あるか
- それらの \(\sum_{i \in S} A_i\) の総和はいくつか
が分かればよいです。これを直接用いるのは難しいですが、この工夫によって高速に求めることができます。
半分全列挙 を用いましょう。\(T\) を \(T_l\) と \(T_r\) の \(2\) つに分けます。
「\(T_l\) から集合 \(S_l\) を選び、 \(T_r\) から \(k\) 要素を選ぶような \(S\)」についてまとめて計算することを考えると
- \(|S_r| = k\)
- \(\sum_{i \in S_r} A_i \lt X - \sum_{i \in S_l} A_i\)
という条件を満たす \(S_r \subset T_r\) について
- 条件を満たす \(S_r\) が何個あるか
- それらの \(\sum_{i \in S_r} A_i\) の総和はいくつか
が求まればよいです。
これは、事前に \(|S_r| = k\) ごとに
- \(\sum_{i \in S_r} A_i\) の昇順に並べた配列
- その prefix sum
を用意しておけば二分探索などで求めることができます。そして、すべての \((S_l, k)\) について期待値を求め、その総和を取ればよいです。
計算量を考えましょう。以下 \(l = |T_l|, r = |T_r|\) とします。
前計算では \(2^r\) 要素の列挙と、要素数 \(k\) ごとの sort が必要なので \(O(r \cdot 2 ^ {r})\) かかります。そして \(2^l\) 個ある集合 \(S_l\) と \(k ~ (0 \le k \le r)\) について、二分探索が必要なので \(O(2^l \cdot r^2)\) かかります。
\(l, r\) をともに \(N/2\) 程度にとることで、全体で \(O(N^2 2^{N/2})\) で解くことができます。
\(S_l\) をあらかじめ sort しておき、二分探索を尺取法にすることで \(O(N 2^{N/2})\) で解くことができます。
また TL が厳しいと感じた場合、\(O(r2^r + r^22^l)\) からわかるように \(r\) を大きめにとると良いでしょう。
実装例(PyPy)
from math import comb
from bisect import bisect_left
n, x = map(int, input().split())
a = list(map(int, input().split()))
mod = 998244353
def gen(a):
n = len(a)
r = [[] for _ in range(n + 1)]
r[0].append(0)
for x in a:
for k in reversed(range(n)):
if r[k]: r[k+1].extend(s + x for s in r[k])
return r
h = n // 2
tls, trs = gen(a[:h]), gen(a[h:])
a_all = sum(a) % mod
q = [pow((n - k) * comb(n, k), -1, mod) for k in range(n)] + [0] # 分母の値
ans = 0
for k, tr in enumerate(trs):
tr.sort()
pre = [0] * (len(tr) + 1) # prefix sum
for i in range(len(tr)):
pre[i+1] = (pre[i] + tr[i]) % mod
for i, tl in enumerate(tls):
for sl in tl:
c = bisect_left(tr, x - sl)
p = (c * (a_all - sl % mod) - pre[c]) % mod
ans += q[i+k] * p % mod
ans %= mod
print(ans)
投稿日時:
最終更新:
