Merge k Sorted Lists
listsの行として、整数のリストをk個受け取ります。各行は非減少順にソートされており、行ごとに長さは異なる場合があり、空の行はありません。
すべての行のすべての値を含む1つのリストにマージし、非減少順にソートして返してください。ある値が1つの行または複数の行に何度か現れる場合、結果にもその回数だけ現れます。
関数
- listsinteger-2d-array
- 行ごとに並べ替えられたリスト(長さは異なる場合があります)
- 戻り値integer-array
- すべての行のすべての値を、1つのソート済みリストに
制約
1 ≤ lists.length ≤ 1041 ≤ lists[i].length、すべての行を合わせて最大104個の値を保持します-104 ≤ lists[i][j] ≤ 104- すべての行は非減少順に並べ替えられています。
例
- 入力
- lists = [[2, 6, 9], [1, 4, 10], [3, 5]]
- 出力
- [1, 2, 3, 4, 5, 6, 9, 10]
- 説明
- 全体で最小の値は1で、2行目の最初の値です。その後、各行は2、4、3で始まるため、次は2で、その後も同様です。3行目は5の後で値がなくなるため、最後に6、9、10が残ります。
- 入力
- lists = [[5], [-2, 5, 7], [0, 5]]
- 出力
- [-2, 0, 5, 5, 5, 7]
- 説明
- 3つの5はそれぞれ異なる行から来ており、3つとも残ります。負の値である
-2は0より前に並びます。
- 入力
- lists = [[4, 8]]
- 出力
- [4, 8]
- 説明
- 1 行だけの場合、マージするものはありません。行はすでにソートされているため、それが答えです。
提出時に隠しテスト+14件
発展問題
すべての行から少なくとも1つの値を含む最小の範囲 [a, b] を見つけます。行の先頭要素を格納した同じヒープと、これまでの最大の先頭要素を使って、O(N log k) で見つけられるでしょうか?
ヒント
1つずつ開いてください。開くたびに少しずつ答えに近づきます。
各行は並べ替えられています。すべての値の中で最小となる可能性があるのはどの値でしょうか?
答えの次の値は、常に各行の最初の未使用の値のうち最小のものです。それを取り出した後、変化するのはそれらの値のうち1つだけです。
各行の最初の未使用値を、行番号を付けて最小ヒープに格納します。最小値を取り出して追加し、その行に次の値があればヒープに追加します。
解説
各行はソートされているため、まだ使用されていない最小の値は、常にいずれかの行の未使用部分の先頭の値です。問題全体は、k 個の行の先頭値の中から最小のものを、値の個数である N 回見つけることです。すべての先頭値を走査すると、値ごとに k ステップかかります。最小ヒープを使えば先頭値を順序どおりに保ち、最小の値を O(log k) で取り出せるため、全体の計算量は O(N·k) から O(N log k) になります。典型的な形式では各リストは連結リストですが、ここでは各行が配列であり、行ごとのインデックスがノードへのポインターの役割を果たします。
すべての値について、k 個のヘッドを比較する
正しいが、最大のテストでは終わらない
考え方
各行につき1つのインデックスpos[r]を保持し、まだ使っていない行rの最初の値、つまり行の先頭を指すようにします。まだ使われていない値のうち最小のものは、必ずこれらの先頭のいずれかです。行rでは、まだ使われていない値はすべてpos[r]以降にあり、行はソートされているので、先頭より小さい値はありません。
そこで、値が残っているすべての行を調べて最小の先頭を見つけ、それを追加して、その行のインデックスを1つ進めます。N個の値をすべて取り出すまで繰り返します。これはマージソートのマージ手順を、2つのリストからk個のリストに拡張したものです。
最初の例では、先頭は2、1、3なので、最初に1が取り出され、2行目の先頭は4になります。次に2(先頭は2、4、3)、続いて3(先頭は6、4、3)、その次に4、そして5が取り出され、3行目は空になります。最後の3ラウンドでは、6と10だけを比較し、次に9と10を比較し、最後は10だけを扱います。
コストは、N個の値それぞれについてk回の比較を行うことです。値が1つずつ入った10^4行では、比較回数は10^8回になります。C、Java、JavaScriptなら1秒未満で処理できますが、Pythonでは10秒以上かかり、Nとkの両方を2倍にすると、どの言語でも4倍遅くなります。トレースを見ると無駄がわかります。各値を選んだ後に変わる先頭は1つだけなのに、次のラウンドでは再びk個すべてを読み取っています。
アルゴリズム
- すべての行について
pos[r] = 0に設定し、値の数をNとして数えます。 N回繰り返します。pos[r]がまだ行の範囲内にあるすべての行を確認し、先頭の値が最小の行を記録します。- その先頭の値を結果に追加し、その行の
posに1を加えます。 - 結果を返します。
def mergeKLists(lists):
pos = [0] * len(lists) # pos[r]: index of the first unused value in row r
total = sum(len(row) for row in lists)
merged = []
for _ in range(total):
best = -1 # the row whose head is the smallest so far
for r in range(len(lists)):
if pos[r] < len(lists[r]) and (best == -1 or lists[r][pos[r]] < lists[best][pos[best]]):
best = r
merged.append(lists[best][pos[best]])
pos[best] += 1
return mergedk 個の先頭要素からなる最小ヒープ
考え方
この走査では、最後の反復から変わったのは1つの先頭要素だけなのに、最小値を見つけるためにk個の先頭要素を読み直します。最小ヒープはまさにこの用途に使えます。数値の集合を保持し、最小値が先頭に来るようにします。先頭の値の取り出しと数値の追加は、どちらもO(log size)のコストがかかります。
各行の最初の値を、行番号をタグとして付けてヒープに入れます。次に、最小のペア(value, row)を取り出し、valueを追加します。その行に次の値があれば、同じタグを付けてヒープに入れます。ヒープには、まだ値が残っている各行の先頭要素が常にちょうど1つずつ入っているため、先頭の値が全体で最も小さい未使用の値になります。これは走査のルールを、より速く実行したものです。
行番号を0から始めて、最初の例を追ってみましょう。ヒープには最初、2(行0)、1(行1)、3(行2)が入っています。1を取り出し、行1の次の値である4を入れます。2を取り出し、行0の6を入れます。3を取り出し、行2の5を入れます。4を取り出し、10を入れます。5を取り出します。行2の値はこれでなくなったので、何も入れず、ヒープには6と10が残ります。6を取り出し、9を入れます。9を取り出し、その後10を取り出します。結果は[1, 2, 3, 4, 5, 6, 9, 10]です。
各値は一度ヒープに入り、一度取り出されます。また、ヒープに保持される要素は最大でもk個なので、2N回の各操作にかかるコストはO(log k)です。N = k = 10^4の場合、これは約2 × 10^4 × 14、つまり3 × 10^5ステップ未満であり、走査の10^8ステップに比べて少なくなります。ヒープが使うメモリはO(k)で、O(N)にはなりません。各行の先頭要素だけを保持し、その後ろにある値は保持しないためです。
いくつかの実装では、各行の先頭要素の順に並べた行番号の配列を使って、ヒープを手作業で構築します。スロットiの子は2i+1と2i+2です(1から数えるLuaとRでは2iと2i+1)。これにより、処理も節約できます。先頭にある行の値を取り出した後、その行の次の値がそれより小さくなることはないため、その行は先頭にとどまり、1回下方向にふるい落とすだけで済みます。取り出しの後に追加する必要はありません。
アルゴリズム
- すべての行
rについて、(lists[r][0], r)を値で順序付けられた最小ヒープに追加します。 - ヒープが空でない間、最小のペア
(value, r)を取り出し、valueを結果に追加します。 - 行
rに次の値がある場合、それとrを一緒に追加します。 - ヒープが空になったら、結果を返します。
import heapq
def mergeKLists(lists):
# The heap holds one (value, row) pair per row that still has values: that row's head.
heap = [(row[0], r) for r, row in enumerate(lists)]
heapq.heapify(heap)
nxt = [1] * len(lists) # nxt[r]: index of row r's next value, not yet in the heap
merged = []
while heap:
value, r = heapq.heappop(heap) # the smallest head of all rows
merged.append(value)
if nxt[r] < len(lists[r]):
heapq.heappush(heap, (lists[r][nxt[r]], r)) # row r's new head takes its place
nxt[r] += 1
return merged
落とし穴と境界ケース
ヒープのロジックは短いものです。バグのほとんどは、ヒープに何を入れるか、そしてどのような順序で並べるかに起因します。
- 値の出どころを記録し忘れる。ヒープに値だけを入れると、pop の後にどの行を進めればよいかわかりません。値と一緒に行番号も格納しましょう。
- 誤って最大ヒープを使う。C++ の
priority_queueと Rust のBinaryHeapは最大の値を先頭に置きます。greater<>またはReverseを使いましょう。Java のPriorityQueueと Python のheapqは、最小の値を先頭に置きます。 - Python の
heapqで値が同じ場合。2つの値が等しいと、タプルの比較は2番目の要素に進みます。行番号なら問題なく比較できますが、連結リストのノードは比較できないため、従来の方法では値が等しいとクラッシュします。2番目の要素には行番号かカウンターを入れましょう。 - 最初にすべての値をプッシュする。正しくソートはできますが、ヒープの要素数が
Nまで増え、処理量がO(N log N)になります。各行の先頭要素を1つずつ保持しましょう。 - 短い行の末尾を越えて読み取る。行ごとに長さが異なるため、次の値をプッシュする前に、その行に次の値があることを確認しましょう。
- 重複を取り除く。異なる行にある等しい値は別々の値であり、すべて結果に含める必要があります。
よくある質問4
k 個のソート済みリストをマージする場合の時間計算量は?
最小ヒープを使うと、計算量は O(N log k) です。ここで、N は値の総数、k はリストの数です。各値は1回ずつヒープに追加され、取り出されます。また、ヒープには最大で k 個の要素が格納されるため、各操作のコストは O(log k) です。出力とは別に必要な追加メモリは O(k) です。
すべての値をまとめて並べ替えてみてはどうでしょうか?
それは正しく、実行時間は O(N log N) です。小さな入力であれば問題ありません。リストがすでにソートされているという事実を無視しているため、ヒープでは log k で済むところを、値ごとに log N のコストがかかります。また、すべての値を一度にメモリ上に保持する必要があります。ヒープなら、ソートではできない、ストリームとして届くリストのマージもできます。
ヒープを使わずに、k 個のソート済みリストをマージできますか?
はい、分割統治法を使います。2つのリストをマージする処理でリストをペアごとにマージし、次にその結果をペアごとにマージする、という手順を繰り返します。log k 回のラウンドがあり、各ラウンドですべての値に一度ずつ触れるため、これも O(N log k) です。リストを1つずつ、拡大していく結果にマージする方法はより遅くなります。前の方の値はマージのたびに再びコピーされるため、合計で O(N·k) になります。
なぜヒープには各リストの先頭だけがあればよいのでしょうか?
各リストはソートされているため、最初の未使用の値は、そのリストに残っている値の中で最小の値です。したがって、すべてのリストの中で最小の値は各リストの先頭の値のうち最小のものであり、リストのより奥にある値がそれを上回ることはありません。先頭の値が取り出されると、同じリストの次の値がそのリストの先頭になり、ヒープ内でその値と入れ替わります。
似た問題
同じ考え方を使う問題です。2〜3問解くとパターンが身につきます。
Python
def mergeKLists(lists):
# ここにコードを書いてくださいケース1
ケース2
ケース3
入力
lists = [[2, 6, 9], [1, 4, 10], [3, 5]]
期待値
[1, 2, 3, 4, 5, 6, 9, 10]