演算法圖鑑
Sorting · 06 / 09

Heap Sort堆積排序

先 heapify 再逐個取出,原地

用在:記憶體受限又要保證 n log n 的場合

時間複雜度O(n log n)
空間複雜度O(1)
難度進階
前置知識Binary Heap、Selection Sort

01為什麼需要它

作業系統核心裡的 sort()

Linux 核心開機時要把上千筆例外處理表(exception table)排好,之後還有各種表格要排序。核心堆疊只有 8 KB 到 16 KB,不能放心遞迴;有些場合不方便配置記憶體;要排的內容也不一定受核心控制。

為什麼用它堆積排序只用陣列本身,sift down 是迴圈不是遞迴,額外空間 O(1),而且不管輸入長什麼樣子都是 O(n log n)。快速排序平均較快,但最壞 O(n²) 可以被刻意觸發;合併排序要 O(n) 的暫存。Linux 的 lib/sort.c 選的就是堆積排序。

內建排序的保險絲

一個 API 接受使用者上傳 100 萬筆數字再排序。有人摸清了你的快速排序怎麼選 pivot,特地送來讓每次分割都極度不平均的資料,比較次數從約 2,000 萬暴增到約 5,000 億。

為什麼用它混合排序平常跑快速排序,一旦發現遞迴層數超過約 2 log n,就把這一段交給堆積排序:它最壞也是 O(n log n),而且同樣原地,不會丟掉快速排序不用暫存陣列的優勢。.NET 的 Array.Sort 與 Rust 的 sort_unstable 都拿堆積排序當最壞情況的退路。

放榜查詢:大多數人只看前幾頁

30 萬名考生的成績要依分數由高到低分頁顯示,每頁 50 名。絕大多數人只看第一、二頁,但沒人知道會不會有人一路翻到最後。

為什麼用它先花 O(n) 把成績 heapify 成最大堆積(約 60 萬次比較以內),之後每要一名就取出一次堆頂,約 36 次比較。第一頁總共不到 61 萬次,全部排好則要約 1,000 萬次。真有人翻到最後一頁,也不過是做完一次完整的堆積排序。

看到這些關鍵字就想到它:原地排序、O(1) 額外空間、最壞也要 O(n log n)、不能遞迴、怕惡意輸入卡出最壞情況、快速排序的退路、邊排邊取出最大的幾個。

02核心概念

堆積排序可以看成換了工具的選擇排序:選擇排序每輪要掃一遍才找得到極值,O(n);如果未排序區本身就是一個最大堆積,最大值就在 a[0],拿走後修復只要 O(log n)。整個演算法分兩個階段,全部在原陣列上完成。heapify:先把整個陣列整理成最大堆積。取出:反覆把堆頂 a[0] 和堆積的最後一格交換,堆積大小減一,再對新的根做 sift down。陣列因此分成兩段,前面 [0, size) 是堆積,後面 [size, n) 是已排好的尾端,尾端一輪長一格。

正確性靠一個不變量:每輪開始時,a[size..n−1] 是全體最大的 n − size 個元素且已由小到大排好,a[0..size−1] 是其餘元素組成的最大堆積。堆頂是堆積裡最大的,又不大於尾端的任何一個,所以把它換到 size − 1 之後尾端依然有序、多了一格。換到根的那個元素只破壞了根這一個位置,左右子樹仍是堆積,一次 sift down 就修好,不變量延續到下一輪;size 降到 1 時整個陣列有序。建堆也是同樣的道理:索引 n // 2 以後都是葉節點,本身就是堆積,從 n // 2 − 1 往前處理,輪到節點 i 時它的兩棵子樹已經是堆積,一次 sift down 就讓以 i 為根的子樹也成為堆積。

複雜度:heapify 看起來是 n/2 次 sift down、每次 O(log n),但高度為 h 的節點最多 ⌈n / 2^(h+1)⌉ 個,每個最多下沉 h 層,總和 Σ h · n / 2^(h+1) = O(n),因為大部分節點都在底層、根本沉不了幾層。取出階段做 n − 1 次 sift down,每次最多走樹高 ⌊log₂ n⌋ 層,O(n log n)。它沒有「運氣不好」的輸入:已排序、反序、隨機,最好、平均、最壞都是 O(n log n)(像全部相等這種特例反而會變成 O(n),因為每次 sift down 第一步就停)。sift down 寫成迴圈時額外空間是 O(1);寫成遞迴則要 O(log n) 的呼叫堆疊,就不算真正原地了。

常見的錯有三個:sift down 判斷子節點存不存在時要比 size(目前的堆積大小)而不是 n,否則已排好的尾端會被拉回堆積;建堆的迴圈要從 n // 2 − 1 往前跑,往後跑時子樹還不是堆積;由小到大要用最大堆積,原地用最小堆積會排成由大到小。它也不穩定[2a, 2b, 1] 排完是 [1, 2b, 2a],兩個 2 的相對順序反了。和鄰居比較:合併排序穩定,但要 O(n) 暫存;快速排序平均更快,因為分割是循序掃過相鄰的記憶體,而 sift down 從 i 跳到 2i + 1,大陣列上幾乎每一步都是快取未命中。所以堆積排序很少當主力,而是在「不能多用記憶體、又不能接受最壞 O(n²)」時出場。Binary Heap 講的是堆積本身,Top-K 用的是大小為 K 的堆積,這裡則是把整個陣列就地變成堆積。

03演算法步驟

  1. 1sift_down(a, i, size):在 i2i + 12i + 2 中找最大的,子節點索引要 < size 才算存在。最大的是 i 就停,否則交換並把 i 移到那個子節點,重複。
  2. 2建堆in // 2 − 1 往下到 0,對每個 i 呼叫 sift_down(a, i, n)。完成後 a[0] 是最大值。
  3. 3取出endn − 1 往下到 1,交換 a[0]a[end],這一輪的最大值落在 end,之後不再移動。
  4. 4對新的根呼叫 sift_down(a, 0, end)。此時堆積大小是 end,傳 n 會把已排好的尾端捲回去。
  5. 5迴圈結束,陣列由小到大排好。要由大到小就把比較反過來(最小堆積);只要最大的前 k 個,取出 k 次就停,成本 O(n + k log n)。

04互動示範

共用陣列 [5, 2, 9, 1, 7, 3, 8, 4]。上方的樹只畫目前還在堆積裡的部分,節點下方的 [i] 是它在陣列裡的索引;下方是同一份陣列,綠色是已排好的尾端。黃色是正在比較的父子節點,藍色是剛交換的兩格。步驟 1 到 8 是 heapify,建完是 [9, 7, 8, 4, 2, 3, 5, 1];之後每次取出都是「堆頂換到尾端」一次交換,再讓新的根往下沉,下沉的層數不會超過樹高。

開始[5, 2, 9, 1, 7, 3, 8, 4] · 最大堆積 · 由小到大
4[7]1[3]2[1]7[4]5[0]3[5]9[2]8[6]
陣列(前 8 格是堆積,綠色尾端已排好)
52917384
黃色是正在比較的父子節點,藍色是剛交換的兩格。樹只畫堆積的部分,節點下方是它在陣列裡的索引。
步驟 0/31堆積排序分兩個階段:先把陣列原地整理成最大堆積(heapify),再反覆把堆頂(最大值)換到尾端、縮小堆積、修復堆頂。

05程式碼

核心是手寫的 sift_down 與兩階段的 heap_sort。Python 另外用 heapq 寫了一個邊排邊輸出的產生器,對應放榜翻頁的情境:它不是原地的,但示範了只取前 k 個時的 O(n + k log n)。C++ 用標準庫的 std::make_heapstd::pop_heap 重寫同一個演算法,傳入 std::greater 就變成由大到小。

import heapq
from itertools import islice


def sift_down(a, i, size):
    """把 a[i] 往下沉,直到它不小於兩個子節點。只看前 size 格。"""
    while True:
        l, r, largest = 2 * i + 1, 2 * i + 2, i
        if l < size and a[l] > a[largest]:     # 子節點存在要比 size,不是 len(a)
            largest = l
        if r < size and a[r] > a[largest]:
            largest = r
        if largest == i:                       # 最大堆積性質成立
            return
        a[i], a[largest] = a[largest], a[i]
        i = largest


# 堆積排序:原地、O(1) 額外空間、最壞 O(n log n)、不穩定
def heap_sort(a):
    n = len(a)
    # 階段一:heapify。索引 n // 2 之後都是葉節點,從最後一個非葉節點往前沉
    for i in range(n // 2 - 1, -1, -1):
        sift_down(a, i, n)
    # 階段二:堆頂是最大值,換到尾端固定,堆積縮小一格後修復堆頂
    for end in range(n - 1, 0, -1):
        a[0], a[end] = a[end], a[0]
        sift_down(a, 0, end)                   # 只修前 end 格,尾端已排好
    return a


# 變形:邊排邊輸出(不是原地)。O(n) 建堆,之後每取一個 O(log n)
# 只取前 k 個時總成本 O(n + k log n),全部取完就是一次完整的堆積排序
def iter_largest(nums):
    h = [-x for x in nums]                     # heapq 是最小堆積,取負當最大堆積
    heapq.heapify(h)
    while h:
        yield -heapq.heappop(h)


if __name__ == "__main__":
    print(heap_sort([5, 2, 9, 1, 7, 3, 8, 4]))     # [1, 2, 3, 4, 5, 7, 8, 9]
    scores = [62, 95, 71, 88, 95, 40, 79]
    print(list(islice(iter_largest(scores), 3)))   # [95, 95, 88](第一頁只要前 3 名)

06練習題

  • LeetCode 506Relative Ranks(從最大堆積依序取出,第幾個出來就是第幾名)Easy
  • LeetCode 1636Sort Array by Increasing Frequency(改寫 sift_down 的比較:先比次數,次數相同時值大的在前)Easy
  • LeetCode 912Sort an Array(手寫堆積排序,O(1) 額外空間且最壞 O(n log n))Medium
  • LeetCode 215Kth Largest Element in an Array(heapify 後只取出 k 次,是提早停下的堆積排序)Medium
  • LeetCode 1962Remove Stones to Minimize the Total(原地 heapify,反覆修改堆頂再往下沉)Medium