Official

D - お土産の組み合わせ / Souvenir Combinations Editorial 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\) 個を取り出す」という典型問題で、優先度付きキュー(ヒープ)を使って解くことができます。

アルゴリズム

  1. \(A\)\(B\) をそれぞれ降順にソートする。
  2. 最大ヒープ(Pythonでは値を負にしたmin-heap)に初期状態として \((A[0] \times B[0],\ 0,\ 0)\) を入れる。
  3. 以下を \(K\) 回繰り返す:
    • ヒープから最大の値を持つ \((i, j)\) を取り出し、\(A[i] \times B[j]\) を合計に加算する。
    • \((i+1, j)\) が未訪問かつ範囲内なら、ヒープに追加して訪問済みに記録する。
    • \((i, j+1)\) が未訪問かつ範囲内なら、ヒープに追加して訪問済みに記録する。
  4. 合計値を出力する。

重複防止が重要です。例えば \((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 によって生成されました。

posted:
last update: