Official

F - Random Vault Heist Editorial 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)

posted:
last update: