分治与主定理 —— a、b、d 与 f(n) 的取法

课程:算法设计与分析 · 分治
日期:2026-09-16


一、分治三步

步骤 做什么
分(Divide) 把原问题拆成若干规模更小、形式相同的子问题
治(Conquer) 递归求解子问题;规模足够小时直接求解(递归出口)
合(Combine) 把子问题的解合并成原问题的解

核心前提:子问题相互独立(不重叠)+ 合并代价可接受。

递推式的建立

T(n) = a · T(n/b) + f(n)
符号 含义 计算
a 本层递归调用了几次 数代码里出现几次递归调用
b 规模缩小的倍数 子问题规模是 n/2 → b = 2;是 n/3 → b = 3
f(n) 本层「分」+「合」的代价 只算本层,不含递归调用

二、f(n) 怎么取:O(1) 还是 O(n)?

d 就是本层开销关于 n 的幂次

本层开销 f(n) 幂次 d 典型场景
O(1) d = 0 取中点、比较一次、访问一个节点、移动一个盘子
O(n) d = 1 把 n 个元素过一遍(merge / partition / 扫跨中点)
O(n²) d = 2 要处理 n 个元素的每一对
O(n log n) ⚠️ 非 n^d 形式 简化版主定理不适用
O(log n) ⚠️ 非 n^d 形式 同上

d = 0 为什么合法?因为 O(1) = O(n^0),于是 b^d = b^0 = 1。二分查找就是这种情况。

操作判据:合并时要不要把 n 个元素都过一遍?

这是唯一需要问的问题。

  • 要扫一遍 → d = 1:归并的 merge、快排的 partition、最大子数组和的跨中点扫描
  • 只做常数次操作 → d = 0:二分取中点比一次、遍历二叉树访问根、汉诺塔移动一个盘
  • 子结果要做 O(n²) 配对 → d = 2

三、主定理简化版:只比 a 与 b^d

适用形式:

T(n) = a · T(n/b) + O(n^d)        (a ≥ 1, b > 1, d ≥ 0)

展开递归树,第 k 层有 a^k 个子问题,每个规模 n / b^k:

第 k 层总代价 = a^k × (n / b^k)^d = n^d × (a / b^d)^k

每一层代价 = 上一层 × a / b^d,所以 a 与 b^d 谁大谁小,决定了每层工作量是涨、平还是落:

  • a:子问题数量增长的倍数
  • b^d:规模缩小 b 倍时,单个子问题代价缩小的倍数
情况 条件 每层代价变化 谁决定总量 复杂度
1 a < b^d 越往下越小 根层 n^d O(n^d)
2 a = b^d 每层一样大 所有层(多乘一个层数) O(n^d · log n)
3 a > b^d 越往下越大 叶子层 n^(log_b a) O(n^(log_b a))

四、例题速查表

算法 递推式 a b d b^d 比较 复杂度
二分查找 T(n)=T(n/2)+O(1) 1 2 0 1 = O(log n)
归并排序 T(n)=2T(n/2)+O(n) 2 2 1 2 = O(n log n)
二叉树遍历 T(n)=2T(n/2)+O(1) 2 2 0 1 > O(n)
快速选择(第K大) T(n)=T(n/2)+O(n) 1 2 1 2 < O(n)
最大子数组和 T(n)=2T(n/2)+O(n) 2 2 1 2 = O(n log n)
最近点对 T(n)=2T(n/2)+O(n) 2 2 1 2 = O(n log n)
Karatsuba 乘法 T(n)=3T(n/2)+O(n) 3 2 1 2 > O(n^1.585)
Strassen 矩阵乘 T(n)=7T(n/2)+O(n²) 7 2 2 4 > O(n^2.807)
普通分块矩阵乘 T(n)=8T(n/2)+O(n²) 8 2 2 4 > O(n^3)

五、代码模板:归并排序

// 归并排序:T(n) = 2T(n/2) + O(n)  →  a=2, b=2, d=1  →  O(n log n)
// 依赖:#include <vector> + using std::vector;
// (<bits/stdc++.h> 是 GCC 专有头文件,macOS 自带 clang 没有,别写)
void mergeArr(vector<int>& a, int l, int mid, int r); // C++ 必须先声明(Java 无此要求);名字避开 std::merge

void mergeSort(vector<int>& a, int l, int r) {
if (l >= r) return; // 治:规模足够小,直接求解
int mid = l + (r - l) / 2;
mergeSort(a, l, mid); // 治:递归左半
mergeSort(a, mid + 1, r); // 治:递归右半
mergeArr(a, l, mid, r); // 合:本层 O(n)
}

void mergeArr(vector<int>& a, int l, int mid, int r) {
vector<int> tmp;
tmp.reserve(r - l + 1);
int i = l, j = mid + 1;
while (i <= mid && j <= r) tmp.push_back(a[i] <= a[j] ? a[i++] : a[j++]);
while (i <= mid) tmp.push_back(a[i++]);
while (j <= r) tmp.push_back(a[j++]);
for (int k = 0; k < (int)tmp.size(); ++k) a[l + k] = tmp[k];
}