演算法圖鑑
Sorting · 04 / 09

Merge Sort合併排序

切半、各自排、合併,穩定

用在:外部排序大檔案、鏈結串列排序、逆序對

時間複雜度O(n log n)
空間複雜度O(n)
難度進階
前置知識Recursion、Insertion Sort

01為什麼需要它

120 GB 的檔案,機器只有 16 GB 記憶體

一份 120 GB 的點擊紀錄要依使用者 ID 排序,整份讀進記憶體根本放不下,任何需要隨機存取整個陣列的排序都用不了。

為什麼用它每次讀 10 GB 進來排好、寫成一個有序的暫存檔,得到 12 個;再同時打開這 12 個檔,每個只看最前面那一筆,挑最小的輸出。合併只需要循序讀寫,正好是磁碟最擅長的。Unix 的 sort 指令和資料庫的 ORDER BY 在記憶體不夠時都這樣做,PostgreSQL 查詢計畫裡的 external merge 就是它。

後台表格:點一下欄位,同組內的順序不能亂

訂單列表有 3 萬筆,已經依下單時間排好。客服點「物流狀態」欄位排序,希望同一個狀態裡的訂單仍然照時間排。

為什麼用它合併時兩值相等一律先取左段,左段的元素原本就排在前面,所以相等元素的相對順序永遠不變,這叫穩定。有了穩定排序,多欄位排序只要「先排次要欄位、再排主要欄位」。Python 的 sort、Java 對物件的 Arrays.sort 都用以合併為核心的 TimSort,就是為了保證這一點。

公開 API 接受使用者上傳的資料來排序

服務接受最多 100 萬筆數字並回傳排序結果。有人刻意構造資料,讓快速排序的 pivot 每次都選到極端值。

為什麼用它快速排序最壞退化成 O(n²),100 萬筆約 5×10¹¹ 次比較,服務直接卡死。合併排序永遠從正中間切,切法和資料內容無關,最壞也只要約 2×10⁷ 次比較,惡意輸入找不到弱點。

看到這些關鍵字就想到它:記憶體放不下、外部排序、需要穩定排序、最壞也要 n log n、鏈結串列排序、合併兩段有序、逆序對或「右邊有幾個比我小」。

02核心概念

兩段已經排好的資料要合成一段,只要兩個指標各指段首,每次比較、把較小的放到輸出、那一邊前進,n 個元素放 n 次就好,O(n)。合併排序把排序問題轉成合併問題:把陣列從中間切成兩半,遞迴把兩半各自排好,再合併。一路切到只剩 0 或 1 個元素時,它天生有序,這就是 base case。

正確性用歸納法:假設遞迴回傳時兩半都已有序,只要合併正確,整段就有序。合併的不變量是「輸出區已放好的 k 個元素,是兩段裡最小的 k 個,而且有序」。兩段各自有序,所以剩下的元素裡最小的一定是兩個指標所指的其中一個,取較小的放上去,不變量維持;放滿時整段完成。穩定性來自比較寫成 a[i] <= a[j]:相等時取左段,而左段的元素在原陣列本來就在右段之前。寫成 < 會先取右邊,結果仍然有序,但穩定性就沒了。

複雜度:T(n) = 2T(n/2) + O(n)。畫成遞迴樹,第 d 層有 2^d 段、每段長 n/2^d,這一層所有合併加起來剛好處理 n 個元素;切到長度 1 要 ⌈log₂ n⌉ 層,所以總共 O(n log n)。切法不看資料內容,最好、平均、最壞都是 O(n log n),每次合併的比較次數介於較短那段的長度和 len − 1 之間,只影響常數。若加上「a[mid-1] <= a[mid] 就跳過合併」的檢查,已排序的輸入降到 O(n)。空間:合併要暫存陣列,整個排序共用一塊 O(n) 的 buf 即可,遞迴堆疊再加 O(log n)。由下而上的迭代版省掉遞迴,但暫存陣列一樣是 O(n)。

常見錯誤:每次遞迴都切片出新陣列(a[:mid]a[mid:]),複雜度沒變,但配置記憶體的次數多、常數大,應該先配置一塊 buf 共用;區間開閉混用,[lo, mid)[mid, hi)hi - lo <= 1 是一套,改成閉區間就要整套跟著改,否則會漏格或無限遞迴。和鄰近演算法比:快速排序原地、快取友善、平均通常更快,但最壞 O(n²) 且不穩定;堆積排序只要 O(1) 空間,也不穩定。合併排序用 O(n) 空間換到「穩定+最壞 O(n log n)」。短段落用插入排序反而更快,所以 TimSort 會先用插入排序把短段整理好再合併。合併這一步本身在串列上只改指標(見合併串列),在合併時順便計數就是逆序對。

03演算法步驟

  1. 1定義 sort(lo, hi):排好半開區間 [lo, hi)hi - lo <= 1 時直接回傳,這是 base case。
  2. 2mid = (lo + hi) // 2,遞迴 sort(lo, mid)sort(mid, hi),回來時兩段都已有序。
  3. 3合併:i = loj = mid,比較 a[i]a[j],較小的寫進暫存陣列、那一邊的指標前進。相等時取左邊(<=),保持穩定。
  4. 4其中一段用完後,另一段剩下的本來就有序,整段照抄;最後把暫存陣列的 [lo, hi) 寫回原陣列。
  5. 5暫存陣列在最外層配置一次、所有合併共用,不要每層都建新陣列。
  6. 6實務上的改進:a[mid-1] <= a[mid] 時兩段已接得起來,跳過合併;段長很短(例如十幾個以下)時改用插入排序。不想用遞迴就改成由下而上,段長 1、2、4… 逐輪合併。

04互動示範

共用陣列 [5, 2, 9, 1, 7, 3, 8, 4]。四列是遞迴的第 0 到第 3 層:切半時整段往下搬一層,合併時從下一層逐個取回上一層,虛線格代表這個位置的值目前在別層。黃色是合併時左右兩段的指標,藍色是剛放進輸出的位置,綠色是已排好的區段。數一數比較次數:第 2 層四次合併各 1 次、第 1 層兩次各 3 次、第 0 層 7 次,共 17 次,每一層都不超過 n = 8。

開始[5, 2, 9, 1, 7, 3, 8, 4] · 由小到大
遞迴的每一層(上:整段,下:切到只剩一個)
0
52917384
1
········
2
········
3
········
黃色是合併時左右兩段的指標,藍色是剛放進輸出的位置,綠色是已排好的區段
步驟 0/39合併排序分兩個階段:先一路切半直到每段只有一個元素(一個元素天生有序),再把相鄰兩段合併回去。

05程式碼

由上而下的遞迴版和由下而上的迭代版共用同一個合併函式。遞迴版直接對應演算法步驟;迭代版不用遞迴,段長 1、2、4… 一輪一輪合併,也正是外部排序一輪輪合併暫存檔的形狀。兩版都只配置一次暫存陣列、用 <= 保持穩定;最後用訂單資料示範穩定性,C++ 對照標準庫的 std::stable_sort

# 把 src[lo:mid] 與 src[mid:hi] 兩段有序的合併到 dst[lo:hi]
def merge(src, dst, lo, mid, hi, key):
    i, j = lo, mid
    for k in range(lo, hi):
        # 左邊還有,而且右邊用完或「左 <= 右」就取左邊;相等時取左邊才穩定
        if i < mid and (j == hi or key(src[i]) <= key(src[j])):
            dst[k] = src[i]
            i += 1
        else:
            dst[k] = src[j]
            j += 1


# 由上而下:切半、遞迴排好兩半、合併。半開區間 [lo, hi)
def merge_sort(a, key=lambda x: x):
    buf = a[:]                              # 暫存陣列只配置一次,所有合併共用

    def sort(lo, hi):
        if hi - lo <= 1:                    # 0 或 1 個元素天生有序
            return
        mid = (lo + hi) // 2
        sort(lo, mid)
        sort(mid, hi)
        if key(a[mid - 1]) <= key(a[mid]):  # 兩段本來就接得起來,跳過合併
            return
        merge(a, buf, lo, mid, hi, key)
        for k in range(lo, hi):             # 合併結果寫回原陣列
            a[k] = buf[k]

    sort(0, len(a))
    return a


# 由下而上:不用遞迴,段長 1、2、4、8… 一輪一輪兩兩合併
def merge_sort_bottom_up(a, key=lambda x: x):
    n = len(a)
    src, dst = a, [None] * n                # 只多配置一塊暫存陣列
    width = 1
    while width < n:
        for lo in range(0, n, 2 * width):
            merge(src, dst, lo, min(lo + width, n), min(lo + 2 * width, n), key)
        src, dst = dst, src                 # 這一輪的輸出是下一輪的輸入
        width *= 2
    if src is not a:                        # 結果最後落在暫存陣列就搬回來
        a[:] = src
    return a


if __name__ == "__main__":
    print(merge_sort([5, 2, 9, 1, 7, 3, 8, 4]))            # [1, 2, 3, 4, 5, 7, 8, 9]
    print(merge_sort_bottom_up([5, 2, 9, 1, 7, 3, 8, 4]))  # [1, 2, 3, 4, 5, 7, 8, 9]
    # 穩定性:訂單已依時間排好,依狀態排序後,同狀態的訂單仍照時間排
    orders = [("A01", "shipped"), ("A02", "pending"), ("A03", "shipped"), ("A04", "pending")]
    print(merge_sort(orders, key=lambda o: o[1]))
    # [('A02', 'pending'), ('A04', 'pending'), ('A01', 'shipped'), ('A03', 'shipped')]

06練習題

  • LeetCode 2570Merge Two 2D Arrays by Summing Values(單獨練合併這一步)Easy
  • LeetCode 912Sort an Array(由上而下與由下而上各寫一次)Medium
  • LeetCode 148Sort List(串列版,由下而上可做到 O(1) 額外空間)Medium
  • LeetCode 937Reorder Data in Log Files(依賴穩定排序)Medium
  • LeetCode 315Count of Smaller Numbers After Self(合併時順便計數)Hard
  • LeetCode 493Reverse Pairs(合併前先用雙指標計數)Hard