ABC123 D - Cake 123
https://atcoder.jp/contests/abc123/tasks/abc123_d
ABC440-E を解くためにはまずこの問題を理解する必要がありました。
概要
数列 A, B, C は全て降順ソートしておくこととします。美味しい順に K 個のケーキの作り方を挙げていくわけですが、最初はケーキ $ A_1 + B_1 + C_1 $ から出発し、i, j, k のどれかを +1 しながら調べていく方法が考えられます。しかしこれをやると $ X \times Y \times Z $ 通り調べることになるので計算量オーバーです。そこで、工夫しながら探索していくのがこの解法です。
ケーキを作る上で重要な法則があります。あるケーキ $ A_i + B_j + C_k $ を作ることを考えます。このとき、このケーキよりも美味しさが一段階小さなケーキは
- $ A_{i+1} + B_{j} + C_{k} $
- $ A_{i} + B_{j+1} + C_{k} $
- $ A_{i} + B_{j} + C_{k+1} $
のいずれかです。言い換えると、あるケーキとこれらのケーキの間に挟まるようなケーキ($ A_i + B_j + C_k > d > A_{i+1} + B_{j} + C_{k} $ みたいになるような美味しさ $ d $)は存在しません。
探索方法
A□
B□
C□
これがスタート地点。$(i, j, k) = (1, 1, 1) $ です。以後、隣接する1マス、より正確には右方向に隣接する1マスに進んで探索範囲を拡げていきます。
A□□□□□
B□□□□□□□
C□□
ここで、今 $(i, j, k) = (5, 7, 2) $ を上位 $ e $番目の美味しさとしてチェックしたとします ($ e <= K $)。じゃあ次に見るべきはこれよりも一段階美味しさが小さなケーキ
- (6, 7, 2)
- (5, 8, 2)
- (5, 7, 3)
この3つです。
公式解説によれば、キューを用意して次の操作を行います。
- キューに残された中から美味しさが最大のケーキを取り出す。
- 取り出した最大のケーキに対し、それよりも一段階美味しくない3つのケーキをキューに追加する。
これを K 回繰り返すと上位 K 個のケーキを列挙することができます。heapq を使えば heappop が K 回、heappush が $ 3 \times K $ 回です。$ K <= 3000 $ なので計算は余裕で間に合いますね。
また、同じケーキをキューに重複して追加してしまわないようにします。キューに追加したことのあるケーキを (i, j, k) のタプルで表現し、set 型変数に入れておきます。
実装
- heapq は最小値を取り出す仕組みなので全ての美味しさに -1 をかけて正負を反転させておく。
- seen で既読を管理し、既読なら push しない。
- それぞれの探索範囲 X, Y, Z を超えるならもう push しない。
import heapq
X, Y, Z, K = map(int, input().split())
A = list(map(int, input().split()))
B = list(map(int, input().split()))
C = list(map(int, input().split()))
A.sort(reverse=True)
B.sort(reverse=True)
C.sort(reverse=True)
que = [[-A[0]-B[0]-C[0], 0, 0, 0]]
seen = {(0, 0, 0)}
heapq.heapify(que)
ans = []
for _ in range(K):
top = heapq.heappop(que)
ans.append(top[0])
[a, b, c] = top[1:]
if a + 1 < X and (a+1, b, c) not in seen:
heapq.heappush(que, [-A[a+1]-B[b]-C[c], a+1, b, c])
seen.add((a+1, b, c))
if b + 1 < Y and (a, b+1, c) not in seen:
heapq.heappush(que, [-A[a]-B[b+1]-C[c], a, b+1, c])
seen.add((a, b+1, c))
if c + 1 < Z and (a, b, c+1) not in seen:
heapq.heappush(que, [-A[a]-B[b]-C[c+1], a, b, c+1])
seen.add((a, b, c+1))
for an in ans:
print(-an)