演算法圖鑑
Sorting · 05 / 09

Quick Sort快速排序

選 pivot 分兩邊,平均最快

用在:大多數語言內建排序的基礎,Quick Select 找第 k 大

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

01為什麼需要它

電商搜尋結果依價格排序

一次查詢撈出 200 萬件商品,要在記憶體裡依價格由低到高排好再分頁。價格是 double,同價的商品誰先誰後無所謂。

為什麼用它快速排序在原陣列上交換,不必像合併排序再開一份 200 萬格的暫存陣列(多 16 MB);分割是從頭到尾循序掃,對 CPU 快取很友善,常數比其他 O(n log n) 排序小。不要求穩定時,語言內建的排序多半以它為主體,例如 Java 對 double[] 的 Arrays.sort 用的就是雙軸快速排序。

監控面板上的 p50 與 p99 延遲

每分鐘收到 120 萬筆 API 回應時間,要算出 p50 和 p99,也就是排序後第 60 萬個與第 118 萬 8 千個位置上的值。全部排好是 O(n log n),但其他 119 萬多個位置的順序根本用不到。

為什麼用它Quick Select 用同樣的分割,但分割完只往答案所在的那一邊走,另一邊整個丟掉。每次切在中間時總共只掃 n + n/2 + n/4 + … ≤ 2n 個元素,隨機選 pivot 下期望仍是 O(n)。C++ 的 std::nth_element 就是這個想法。

三千萬筆訂單依狀態分組

訂單只有「待付款、已付款、出貨中、已送達、已取消」五種狀態,要把三千萬筆原地依狀態排好,不想再多開一份陣列。

為什麼用它重複值極多時,一般的分割把等於 pivot 的元素全擠到同一邊,最壞會退化成 O(n²)。三路分割一次把「等於 pivot」的整段收在中間、不再遞迴,每往下一層至少少一種狀態,最多五層、每層掃 O(n),接近線性,而且全程原地。

看到這些關鍵字就想到它:原地排序、不要求穩定、平均最快、pivot、分割(partition)、第 k 小/中位數/百分位數、大量重複值用三路分割、語言內建 sort。

02核心概念

快速排序是反過來的分治。合併排序直接從中間切,功夫花在合併;快速排序則是把功夫花在前面的分割(partition):選一個 pivot,把 ≤ pivot 的元素搬到左邊、> pivot 的搬到右邊。分割完,pivot 左邊每個都不比它大、右邊每個都比它大,所以它已經站在排序後的最終位置,之後再也不用動。左右兩段各自遞迴就好,沒有合併步驟,全程在原陣列上交換。

為什麼正確?示範用的 Lomuto 分割a[hi] 當 pivot,用 ij 把範圍切成三段並維持不變量:a[lo..i] 都 ≤ pivot、a[i+1..j-1] 都 > pivot、a[j..hi-1] 還沒看。看 a[j] 時,若它 > pivot,j 前進就自然併入第二段;若它 ≤ pivot,i 先前進一格再交換 a[i]a[j],第二段開頭那個大元素被換到第二段尾端,a[j] 則併入第一段,三段的性質都沒被破壞。掃完時第三段是空的,把 pivot 換到 i+1 就得到「左 ≤ pivot < 右」。再對長度做歸納:左右兩段遞迴排好後,左段每個值 ≤ pivot < 右段每個值,整段就有序。

複雜度:同一層遞迴的所有分割加起來掃 O(n),所以時間取決於有幾層。pivot 每次切在中間時 T(n) = 2T(n/2) + O(n),log n 層,O(n log n);每次都是極值時(已排序的資料固定取尾端當 pivot)T(n) = T(n−1) + O(n),n 層,O(n²)。改成隨機選 pivot 後,元素互不相同時期望比較次數約 2n ln n ≈ 1.39 n log₂ n,不管輸入原本怎麼排都一樣,沒有哪種排列能穩定觸發最壞情況。空間方面分割是原地的,額外成本只有遞迴堆疊:平均深度 O(log n),但兩邊都直接遞迴時最壞深度是 O(n)。只遞迴較短的一邊、較長的一邊用迴圈接著處理,每深一層範圍至少減半,深度保證不超過 log₂ n,這才是 O(log n) 空間的由來。

常見的坑有三個。第一是大量重複值:Lomuto 把等於 pivot 的元素全放左邊,所有值都相同時每次只少一個,隨機 pivot 也救不了;要改用三路分割,切成 < pivot== pivot> pivot 三段,中間整段一次定位。第二是不穩定:遠距離交換會打亂同值元素的先後,例如 [2₁, 2₂, 1] 以 1 為 pivot 分割後變成 [1, 2₂, 2₁]。第三是遞迴範圍沒有排除 pivot,寫成 [lo, p] 時範圍可能完全不縮小而無限遞迴。和鄰近課程比:合併排序最壞也是 O(n log n) 而且穩定,但要 O(n) 暫存陣列;堆積排序最壞 O(n log n) 又只要 O(1) 空間,但存取跳來跳去,實務上較慢。所以不穩定的內建排序多以快速排序為主體再加保險(例如 introsort 在遞迴太深時改用堆積排序、小段落改用插入排序),需要穩定時(Python 的 sorted、Java 的物件排序)則用合併型的 Timsort。

03演算法步驟

  1. 1範圍 [lo, hi] 只剩 0 或 1 個元素(lo >= hi)就直接返回。
  2. 2[lo, hi] 隨機選一個索引,和 a[hi] 交換,讓它當 pivot。
  3. 3Lomuto 分割:i = lo − 1jlo 掃到 hi − 1,遇到 a[j] <= pivoti += 1 並交換 a[i]a[j]
  4. 4掃完交換 a[i+1]a[hi]p = i + 1 就是 pivot 的最終位置。
  5. 5處理 [lo, p−1][p+1, hi]:較短的一邊遞迴,較長的一邊更新 lohi 後回到步驟 1 用迴圈繼續,堆疊深度就不會超過 log n。
  6. 6資料有大量重複值時改用三路分割(ltigt 三個指標),等於 pivot 的整段不再遞迴;只要第 k 小就用 Quick Select,每次只往 k 所在的那一段繼續。

04互動示範

共用陣列 [5, 2, 9, 1, 7, 3, 8, 4],固定取範圍最後一個元素當 pivot 做 Lomuto 分割(示範不隨機,每次播放才會一樣)。黃色是 pivot,藍色是剛交換的兩格,綠色是已定位的元素,灰色在目前範圍之外;分割進行中,下方會列出 ≤ pivot 區、> pivot 區與還沒看的元素。第一刀以 4 為 pivot 切成 3 個和 4 個,還算平均;之後以 3、5、7 為 pivot 的三次分割,pivot 都剛好是當時範圍裡的極值,範圍每次只縮小 1,這就是退化成 O(n²) 的樣子。

開始[5, 2, 9, 1, 7, 3, 8, 4] · Lomuto 分割
陣列
52917384
黃色是 pivot,藍色是剛交換的兩格,綠色是已定位的元素,灰色在目前範圍之外。i 是 ≤ pivot 區的右界,j 是掃描指標,p 標出 pivot。
步驟 0/29Lomuto 分割:取範圍最後一個元素當 pivot,用指標 j 從左掃到右,把 ≤ pivot 的元素往前集中到 i 之前,最後把 pivot 放到 i+1,它就永遠定位了。

05程式碼

Lomuto 分割(和示範相同)、隨機 pivot 並只遞迴較短一邊的快速排序、處理大量重複值的三路分割版本,以及用三路分割實作的 Quick Select。兩種分割都列出來,是因為 Lomuto 最好懂,但遇到大量重複值只有三路分割撐得住,Quick Select 也因此選用三路分割。

import random


def partition(a, lo, hi):
    """Lomuto 分割:以 a[hi] 為 pivot,回傳 pivot 最後的位置"""
    pivot = a[hi]
    i = lo - 1                              # a[lo..i] 都 <= pivot
    for j in range(lo, hi):
        if a[j] <= pivot:
            i += 1
            a[i], a[j] = a[j], a[i]
    a[i + 1], a[hi] = a[hi], a[i + 1]       # pivot 放到兩區中間,從此不再移動
    return i + 1


def quick_sort(a, lo=0, hi=None):
    """隨機 pivot + 只遞迴較短的一邊:平均 O(n log n),堆疊深度 O(log n)"""
    if hi is None:
        hi = len(a) - 1
    while lo < hi:
        r = random.randint(lo, hi)          # 隨機選 pivot,換到尾端再分割
        a[r], a[hi] = a[hi], a[r]
        p = partition(a, lo, hi)
        if p - lo < hi - p:                 # 左邊較短:遞迴左邊,右邊留給迴圈
            quick_sort(a, lo, p - 1)
            lo = p + 1
        else:
            quick_sort(a, p + 1, hi)
            hi = p - 1
    return a


def partition3(a, lo, hi):
    """三路分割:a[lo..lt-1] < pivot、a[lt..gt] == pivot、a[gt+1..hi] > pivot"""
    pivot = a[random.randint(lo, hi)]
    lt, i, gt = lo, lo, hi
    while i <= gt:
        if a[i] < pivot:
            a[lt], a[i] = a[i], a[lt]
            lt += 1
            i += 1
        elif a[i] > pivot:
            a[i], a[gt] = a[gt], a[i]
            gt -= 1                         # 換過來的元素還沒看過,i 不前進
        else:
            i += 1
    return lt, gt


def quick_sort_3way(a, lo=0, hi=None):
    """大量重複值也不退化:等於 pivot 的整段一次定位"""
    if hi is None:
        hi = len(a) - 1
    if lo >= hi:
        return a
    lt, gt = partition3(a, lo, hi)
    quick_sort_3way(a, lo, lt - 1)
    quick_sort_3way(a, gt + 1, hi)
    return a


def quick_select(a, k):
    """第 k 小的值(k 從 0 起算),平均 O(n);會改變 a 的順序"""
    lo, hi = 0, len(a) - 1
    while True:
        lt, gt = partition3(a, lo, hi)
        if k < lt:
            hi = lt - 1                     # 答案在左段,右邊整個丟掉
        elif k > gt:
            lo = gt + 1
        else:
            return a[k]                     # k 落在「等於 pivot」那段


if __name__ == "__main__":
    b = [5, 2, 9, 1, 7, 3, 8, 4]
    print(partition(b, 0, len(b) - 1), b)   # 3 [2, 1, 3, 4, 7, 9, 8, 5](和示範第一次分割相同)
    print(quick_sort([5, 2, 9, 1, 7, 3, 8, 4]))           # [1, 2, 3, 4, 5, 7, 8, 9]
    print(quick_sort_3way([3, 1, 3, 3, 2, 1, 3, 2]))      # [1, 1, 2, 2, 3, 3, 3, 3]
    nums = [5, 2, 9, 1, 7, 3, 8, 4]
    print(quick_select(nums[:], 3))                       # 4(第 4 小)
    print(quick_select(nums[:], len(nums) - 2))           # 8(第 2 大)

06練習題

  • LeetCode 905Sort Array By Parity(一次分割:偶數放左、奇數放右)Easy
  • LeetCode 75Sort Colors(三路分割,也叫荷蘭國旗問題)Medium
  • LeetCode 2161Partition Array According to Given Pivot(要保持原本的相對順序,交換式分割會打亂它)Medium
  • LeetCode 912Sort an Array(固定取尾端當 pivot 容易超時,加上隨機與三路分割)Medium
  • LeetCode 215Kth Largest Element in an Array(Quick Select,注意大量重複值)Medium
  • LeetCode 324Wiggle Sort II(Quick Select 找中位數再三路分割)Medium