Go 數據結構和算法篇:歸併排序

一、實現原理


今天介紹比前面三種排序算法性能更好的排序算法 —— 歸併排序。

所謂歸併排序,指的是如果要排序一個數據序列,我們可以先把該序列從中間分成前後兩部分,然後對這兩部分分別做排序操作,再將排好序的兩部分合並在一起,這樣整個數據序列就都有序了。

歸併排序使用了分治思想,分治,顧名思義,就是分而治之,將一個大問題分解成小的子問題來解決。說到這裏,可能你會聯想起我們之前講到的一個編程技巧 —— 遞歸,沒錯,歸併排序就是通過遞歸來實現的。這個遞歸的公式是每次都將傳入的待排序數據序列一分爲二,直到變成不能繼續分割的最小區間單元,然後將最小區間單元數據排序後合併起來,最終返回的就是排序好的數據序列了。圖示如下:

歸併排序圖示

由於涉及到遞歸,所以歸併排序從理解上要比前面三個排序要困難一些,還是建議通過這個動態圖幫助理解:https://visualgo.net/zh/sorting(在界面頂部選擇歸併排序,然後在左下角點擊執行即可)。

二、示例代碼

通過上面的分析,我們知道歸併 = 遞歸 + 合併,對應的 Go 實現代碼如下:

 1package main
 2
 3import (
 4    "fmt"
 5)
 6
 7// 歸併排序
 8func mergeSort(nums []int) []int {
 9    if len(nums) <= 1 {
10        return nums
11    }
12
13    // 獲取分區位置
14    p := len(nums) / 2
15    // 通過遞歸分區
16    left := mergeSort(nums[0:p])
17    right := mergeSort(nums[p:])
18    // 排序後合併
19    return merge(left, right)
20}
21
22// 排序合併
23func merge(left []int, right []int) []int {
24    i, j := 0, 0
25    m, n := len(left), len(right)
26    // 用於存放結果集
27    var result []int
28    for {
29        // 任何一個區間遍歷完,則退出
30        if i >= m || j >= n {
31            break
32        }
33        // 對所有區間數據進行排序
34        if left[i] <= right[j] {
35            result = append(result, left[i])
36            i++
37        } else {
38            result = append(result, right[j])
39            j++
40        }
41    }
42
43    // 如果左側區間還沒有遍歷完,將剩餘數據放到結果集
44    if i != m {
45        for ; i < m; i++ {
46            result = append(result, left[i])
47        }
48    }
49
50    // 如果右側區間還沒有遍歷完,將剩餘數據放到結果集
51    if j != n {
52        for ; j < n; j++ {
53            result = append(result, right[j])
54        }
55    }
56
57    // 返回排序後的結果集
58    return result
59}
60
61func main() {
62    nums := []int{4, 5, 6, 7, 8, 3, 2, 1}
63    sortedNums := mergeSort(nums)
64    fmt.Println(sortedNums)
65}
66
67

運行上述代碼,打印結果如下:

三、性能分析

最後我們來看下歸併排序的性能:

歸併排序的時間複雜度推導過程

歸併的思路是將一個複雜的問題 a 遞歸拆解爲子問題 b 和 c,再將子問題計算結果合併,最終得到問題的答案,這裏我們將歸併排序總的時間複雜度設爲 T(n),則 T(n) = 2*T(n/2) + n,其中 T(n/2) 是遞歸拆解的第一步對應子問題的時間複雜度,n 則是排序合併函數的時間複雜度(一個循環遍歷),依次類推,我們可以推導 T(n) 的計算邏輯如下:

1T(n) = 2*T(n/2) + n
2        = 2*(2*T(n/4) + n/2) + n = 4*T(n/4) + 2*n
3        = 4(2*T(n/8) + n/4) + 2*n = 8*T(n/8) + 3*n
4        = ...
5        = 2^k*T(n/2^k) + k*n
6
7

遞歸到最後,T(n/2k)≈T(1),也就是 n/2k = 1,計算歸併排序的時間複雜度,就演變成了計算 k 的值,2k = n,所以 k=log2n,我們把 k 的值帶入上述 T(n) 的推導公式,得到:

1T(n) = n*T(1) + n*log2n = n(C + log2n)
2
3

注:上述公式中 2 是下標,即 log2n。

把常量和低階忽略,所以 T(n) = nlogn

本文由 Readfog 進行 AMP 轉碼,版權歸原作者所有。
來源https://mp.weixin.qq.com/s/fGAnnZxjH_A2ULJDAxxTgA