リリース・改善中
分割統治法 ガイド · 2/6
この章は現在、英語でのみ提供しています。
This chapter traces four algorithms by hand on small inputs: merge sort, counting inversions during the merge, fast exponentiation and quickselect. Working through the traces once makes the code in the next chapter easy to follow.
Take [5, 2, 4, 7, 1, 3, 2, 6]. The array is cut in half until every piece holds one element, then the pieces are merged back level by level.
| Level | Pieces after merging |
|---|---|
| 0 (single elements) | [5] [2] [4] [7] [1] [3] [2] [6] |
| 1 | [2, 5] [4, 7] [1, 3] [2, 6] |
| 2 | [2, 4, 5, 7] [1, 2, 3, 6] |
| 3 | [1, 2, 2, 3, 4, 5, 6, 7] |
Every level touches each element once, and there are log2 8 = 3 levels: that is where O(n log n) comes from. The merge compares the front elements of two sorted lists and repeatedly takes the smaller one.
def merge(left, right):
out, i, j = [], 0, 0
while i < len(left) and j < len(right):
if left[i] <= right[j]: # <= keeps equal elements in order (stable)
out.append(left[i])
i += 1
else:
out.append(right[j])
j += 1
out.extend(left[i:])
out.extend(right[j:])
return out
print(merge([2, 4, 5, 7], [1, 2, 3, 6])) # [1, 2, 2, 3, 4, 5, 6, 7]An inversion is a pair of positions i < j with a[i] > a[j]. Checking every pair costs O(n^2). Divide and conquer sorts inversions into three kinds: both elements in the left half, both in the right half, or one on each side (crossing). The first two kinds are counted recursively. Crossing pairs come almost for free during the merge: when right[j] is taken before left[i], it is smaller than left[i] and every element after it in the left list, so it forms len(left) - i crossing inversions.
The final merge of the trace above, [2, 4, 5, 7] with [1, 2, 3, 6]:
| Taken | From | Left elements still waiting | Added |
|---|---|---|---|
| 1 | right | 2, 4, 5, 7 | 4 |
| 2 | left | - | 0 |
| 2 | right |
| 4, 5, 7 |
| 3 |
| 3 | right | 4, 5, 7 | 3 |
| 4 | left | - | 0 |
| 5 | left | - | 0 |
| 6 | right | 7 | 1 |
| 7 | left | - | 0 |
The final merge adds 11. The lower levels add 1 for [5] with [2], 1 for [2, 5] with [4, 7] (4 is less than 5) and 1 for [1, 3] with [2, 6] (2 is less than 3). The array therefore has 14 inversions. Because the tie 2 <= 2 takes the left element first, equal values are never counted.
Multiplying b by itself e - 1 times is hopeless when e is large. Halve the exponent instead: b^e = (b^(e/2))^2 when e is even, and b^e = (b^((e-1)/2))^2 * b when e is odd. Each call halves e, so at most about 2 log2 e multiplications are needed.
Trace for 3^13 (13 is 1101 in binary):
| Call | Half result | Computation | Value |
|---|---|---|---|
power(3, 0) | - | base case | 1 |
power(3, 1) | 1 | 1 * 1 * 3 | 3 |
power(3, 3) | 3 | 3 * 3 * 3 | 27 |
power(3, 6) | 27 | 27 * 27 | 729 |
power(3, 13) | 729 | 729 * 729 * 3 | 1594323 |
def power(b, e):
if e == 0:
return 1
half = power(b, e // 2)
return half * half * b if e % 2 else half * half
print(power(3, 13)) # 1594323Five calls replace twelve multiplications; for e = 10^18 it is about 60 calls instead of 10^18 steps.
Quickselect finds the k-th smallest element (0-based) without sorting everything. It partitions around a pivot as quicksort does, then continues only in the side that contains position k.
Trace with the last element as pivot, a = [7, 2, 9, 4, 1, 8, 3], k = 3:
| Range | Pivot | Range after partition | Pivot position | Next |
|---|---|---|---|---|
| 0..6 | 3 | [2, 1, 3, 4, 7, 8, 9] | 2 | k is to the right: 3..6 |
| 3..6 | 9 | [4, 7, 8, 9] | 6 | k is to the left: 3..5 |
| 3..5 | 8 | [4, 7, 8] | 5 | k is to the left: 3..4 |
| 3..4 | 7 | [4, 7] | 4 | k is to the left: 3..3 |
| 3..3 | - | [4] | - | answer 4 |
After the first step each partition peels off a single element. That is exactly how the O(n^2) worst case arises with a poor pivot. A random pivot makes long runs like this unlikely and gives O(n) expected time.
import random
def quickselect(a, k):
"""k-th smallest (0-based) element of a; reorders a in place."""
lo, hi = 0, len(a) - 1
while lo < hi:
p = random.randint(lo, hi) # random pivot
a[p], a[hi] = a[hi], a[p]
pivot, store = a[hi], lo
for i in range(lo, hi): # Lomuto partition
if a[i] < pivot:
a[i], a[store] = a[store], a[i]
store += 1
a[store], a[hi] = a[hi], a[store]
if k == store:
return a[k]
if k < store:
hi = store - 1
else:
lo = store + 1
return a[k]
print(quickselect([7, 2, 9, 4, 1, 8, 3], 3)) # 4log n levels of O(n) work each.len(left) - i whenever the right side wins.O(log e) multiplications.O(n) expected time.
コメント 0件
ログイン · ログインするとコメントできます。
最初のコメントを書いてみましょう。