Quick Sort快速排序
選 pivot 分兩邊,平均最快。
用在:大多數語言內建排序的基礎,Quick Select 找第 k 大
01為什麼需要它
一次查詢撈出 200 萬件商品,要在記憶體裡依價格由低到高排好再分頁。價格是 double,同價的商品誰先誰後無所謂。
為什麼用它快速排序在原陣列上交換,不必像合併排序再開一份 200 萬格的暫存陣列(多 16 MB);分割是從頭到尾循序掃,對 CPU 快取很友善,常數比其他 O(n log n) 排序小。不要求穩定時,語言內建的排序多半以它為主體,例如 Java 對 double[] 的 Arrays.sort 用的就是雙軸快速排序。
每分鐘收到 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,用 i、j 把範圍切成三段並維持不變量: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範圍
[lo, hi]只剩 0 或 1 個元素(lo >= hi)就直接返回。 - 2在
[lo, hi]隨機選一個索引,和a[hi]交換,讓它當 pivot。 - 3Lomuto 分割:
i = lo − 1;j從lo掃到hi − 1,遇到a[j] <= pivot就i += 1並交換a[i]、a[j]。 - 4掃完交換
a[i+1]與a[hi],p = i + 1就是 pivot 的最終位置。 - 5處理
[lo, p−1]與[p+1, hi]:較短的一邊遞迴,較長的一邊更新lo或hi後回到步驟 1 用迴圈繼續,堆疊深度就不會超過 log n。 - 6資料有大量重複值時改用三路分割(
lt、i、gt三個指標),等於 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²) 的樣子。
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