D - お土産の組み合わせ / Souvenir Combinations 解説 by admin
Claude 4.6 Opus (Thinking)概要
\(N\) 種類のお菓子と \(M\) 種類の飲み物から1つずつ選んだ組み合わせの満足度(美味しさの積)のうち、大きい方から \(K\) 個の合計を求める問題です。
考察
素朴なアプローチとその問題点
最も単純な方法は、\(N \times M\) 通りの積をすべて計算してソートし、上位 \(K\) 個を取ることです。しかし \(N, M\) がそれぞれ最大 \(2 \times 10^5\) なので、組み合わせは最大 \(4 \times 10^{10}\) 通りとなり、到底間に合いません。
重要な気づき
\(A\) と \(B\) をそれぞれ降順にソートすると、次の性質が成り立ちます。
- \(A[0] \times B[0]\) が全組み合わせの中で最大
- \(A[i] \times B[j]\) に対して、次に大きい候補は \(A[i+1] \times B[j]\) や \(A[i] \times B[j+1]\) のどちらか
例えば \(A = [5, 3, 1]\), \(B = [4, 2]\) の場合:
| \(B[0]=4\) | \(B[1]=2\) | |
|---|---|---|
| \(A[0]=5\) | 20 | 10 |
| \(A[1]=3\) | 12 | 6 |
| \(A[2]=1\) | 4 | 2 |
最大値 \(20\)(\(i=0, j=0\))を取った後の候補は \(12\)(\(i=1, j=0\))と \(10\)(\(i=0, j=1\))です。このように、現在取り出した要素の「右」と「下」だけを候補に加えていけば、全探索をせずに上位 \(K\) 個を効率的に取り出せます。
解決方法
これは「ソート済み行列から上位 \(K\) 個を取り出す」という典型問題で、優先度付きキュー(ヒープ)を使って解くことができます。
アルゴリズム
- \(A\) と \(B\) をそれぞれ降順にソートする。
- 最大ヒープ(Pythonでは値を負にしたmin-heap)に初期状態として \((A[0] \times B[0],\ 0,\ 0)\) を入れる。
- 以下を \(K\) 回繰り返す:
- ヒープから最大の値を持つ \((i, j)\) を取り出し、\(A[i] \times B[j]\) を合計に加算する。
- \((i+1, j)\) が未訪問かつ範囲内なら、ヒープに追加して訪問済みに記録する。
- \((i, j+1)\) が未訪問かつ範囲内なら、ヒープに追加して訪問済みに記録する。
- 合計値を出力する。
重複防止が重要です。例えば \((1, 1)\) は \((0, 1)\) からの「下」と \((1, 0)\) からの「右」の両方から到達できるため、visited 集合を用いて同じ \((i, j)\) を二度追加しないようにします。
計算量
- 時間計算量: \(O(N \log N + M \log M + K \log K)\)
- ソートに \(O(N \log N + M \log M)\)
- ヒープ操作を \(K\) 回行い、各回で最大2つの要素を追加するため、ヒープサイズは最大 \(O(K)\)。各操作は \(O(\log K)\)。
- 空間計算量: \(O(N + M + K)\)
- 配列 \(A, B\) に \(O(N + M)\)、ヒープと訪問済み集合に \(O(K)\)。
実装のポイント
Pythonの
heapqはmin-heapなので、値を負にして格納することで max-heap として使用しています。visited集合により同じ \((i, j)\) の重複追加を防ぎます。これがないと同じ組み合わせが複数回カウントされてしまいます。\(K \leq 2 \times 10^5\) という制約があるため、ヒープのサイズが爆発せず効率的に動作します。
sys.stdin.buffer.read()を用いた高速入力により、大きな入力でもTLEを回避しています。ソースコード
import heapq
import sys
def main():
input_data = sys.stdin.buffer.read().split()
idx = 0
N = int(input_data[idx]); idx += 1
M = int(input_data[idx]); idx += 1
K = int(input_data[idx]); idx += 1
A = [int(input_data[idx + i]) for i in range(N)]; idx += N
B = [int(input_data[idx + i]) for i in range(M)]; idx += M
# Sort both in descending order
A.sort(reverse=True)
B.sort(reverse=True)
# Use a max-heap (negate values for min-heap)
# State: (-A[i]*B[j], i, j)
# Start with (A[0]*B[0]) which is the maximum
heap = [(-A[0] * B[0], 0, 0)]
visited = set()
visited.add((0, 0))
total = 0
for _ in range(K):
neg_val, i, j = heapq.heappop(heap)
total += -neg_val
# Push (i+1, j) and (i, j+1)
if i + 1 < N and (i + 1, j) not in visited:
visited.add((i + 1, j))
heapq.heappush(heap, (-A[i + 1] * B[j], i + 1, j))
if j + 1 < M and (i, j + 1) not in visited:
visited.add((i, j + 1))
heapq.heappush(heap, (-A[i] * B[j + 1], i, j + 1))
print(total)
main()
この解説は claude4.6opus-thinking によって生成されました。
投稿日時:
最終更新: