数组中的第 K 个最大元素
给定整数数组 nums 和整数 k,请返回数组中第 k 个最大的元素。
注意,需要找的是数组排序后的第 k 个最大的元素,而不是第 k 个不同的元素。
你必须设计并实现时间复杂度为 O(n) 的算法解决此问题。
示例 1:
输入:nums = [3,2,1,5,6,4], k = 2
输出:5
示例 2:
输入:nums = [3,2,3,1,2,4,5,5,6], k = 4
输出:4
提示:
1 <= k <= nums.length <= 10^5-10^4 <= nums[i] <= 10^4
解法一(随机化快速选择):
class Solution {
public:
int findKthLargest(vector<int>& nums, int k) {
int target = nums.size() - k;
int left = 0;
int right = nums.size() - 1;
while (left <= right)
{
auto range = randomPartition(nums, left, right);
if (target < range.first)
{
right = range.first - 1;
}
else if (target > range.second)
{
left = range.second + 1;
}
else
{
return nums[target];
}
}
return -1;
}
private:
pair<int, int> randomPartition(vector<int>& nums, int left, int right) {
static mt19937 gen(random_device{}());
uniform_int_distribution<int> dist(left, right);
int pivotIndex = dist(gen);
int pivot = nums[pivotIndex];
return partition(nums, left, right, pivot);
}
pair<int, int> partition(vector<int>& nums, int left, int right, int pivot) {
int less = left;
int cur = left;
int greater = right;
while (cur <= greater)
{
if (nums[cur] < pivot)
{
swap(nums[less], nums[cur]);
less++;
cur++;
}
else if (nums[cur] > pivot)
{
swap(nums[cur], nums[greater]);
greater--;
}
else
{
cur++;
}
}
return {less, greater};
}
};
解法二(BFPRT,最坏线性时间):
class Solution {
public:
int findKthLargest(vector<int>& nums, int k) {
int target = nums.size() - k;
return select(nums, 0, nums.size() - 1, target);
}
private:
int select(vector<int>& nums, int left, int right, int target) {
while (true)
{
if (left == right)
{
return nums[left];
}
int pivot = medianOfMedians(nums, left, right);
auto range = partition(nums, left, right, pivot);
if (target < range.first)
{
right = range.first - 1;
}
else if (target > range.second)
{
left = range.second + 1;
}
else
{
return nums[target];
}
}
}
int medianOfMedians(vector<int>& nums, int left, int right) {
int size = right - left + 1;
if (size <= 5)
{
sort(nums.begin() + left, nums.begin() + right + 1);
return nums[left + size / 2];
}
int medianIndex = left;
for (int i = left; i <= right; i += 5)
{
int groupRight = min(i + 4, right);
sort(nums.begin() + i, nums.begin() + groupRight + 1);
int median = i + (groupRight - i) / 2;
swap(nums[medianIndex], nums[median]);
medianIndex++;
}
int mid = left + (medianIndex - left) / 2;
return select(nums, left, medianIndex - 1, mid);
}
pair<int, int> partition(vector<int>& nums, int left, int right, int pivot) {
int less = left;
int cur = left;
int greater = right;
while (cur <= greater)
{
if (nums[cur] < pivot)
{
swap(nums[less], nums[cur]);
less++;
cur++;
}
else if (nums[cur] > pivot)
{
swap(nums[cur], nums[greater]);
greater--;
}
else
{
cur++;
}
}
return {less, greater};
}
};
核心思想
如果直接排序,排序后取第 k 个最大元素很简单。
但排序需要 O(n log n),题目要求设计 O(n) 的算法,所以不能把完整排序作为最终解法。
这题最关键的观察是:
找第
k大只需要知道它最终排在哪个位置,不需要把整个数组都排好序。
如果把数组按升序排列,第 k 个最大元素的下标是:
target = nums.size() - k
例如数组长度是 6,第 2 大在升序排序后的下标就是:
6 - 2 = 4
也就是说,问题可以转化成:
找到升序排列后下标为 target 的元素
快速选择就是围绕这个目标下标不断缩小搜索范围。
解法一:随机化快速选择
快速选择和快速排序很像,都会选一个基准值 pivot,然后把数组分成几部分。
这里使用三路划分:
[left, less - 1] 小于 pivot
[less, greater] 等于 pivot
[greater + 1, right] 大于 pivot
划分完成后,只需要看目标下标 target 落在哪个区域:
- 如果
target < less,继续在左边找 - 如果
target > greater,继续在右边找 - 如果
less <= target <= greater,说明答案就是pivot
因为每次只递归或循环处理一边,所以平均时间复杂度是 O(n)。
为什么使用三路划分
数组中可能有重复元素。
题目要求的是排序后的第 k 个最大元素,不是第 k 个不同元素。
例如:
nums = [3,2,3,1,2,4,5,5,6]
两个 5 都要参与排序位置计算。
如果只做普通二路划分,在大量重复元素时可能会反复处理相等元素,效率不稳定。
三路划分把所有等于 pivot 的元素一次性放到中间。
如果 target 落在中间区域,就可以直接返回,不需要继续搜索。
解法二:BFPRT
随机化快速选择的平均时间复杂度是 O(n),但最坏情况下仍然可能退化到 O(n^2)。
如果严格要求最坏时间复杂度也是 O(n),可以使用 BFPRT,也叫“中位数的中位数”。
它和快速选择的整体流程一样,区别在于 pivot 的选择方式。
BFPRT 的做法是:
- 每
5个元素分成一组。 - 每组内部排序,取这一组的中位数。
- 把所有组的中位数放到数组前面。
- 递归选择这些中位数的中位数作为
pivot。
这样选出的 pivot 不会太偏。
每次划分后,可以保证丢掉足够多的元素,因此最坏时间复杂度也能控制在 O(n)。
为什么 target = n - k
如果数组升序排序,下标从 0 开始:
第 1 大 -> 下标 n - 1
第 2 大 -> 下标 n - 2
第 k 大 -> 下标 n - k
所以代码中先计算:
int target = nums.size() - k;
然后按照“找第 target 小”的思路处理。
这样可以统一使用升序划分逻辑。
边界情况
如果 k = 1,要找的是数组最大值。
对应升序下标:
n - 1
如果 k = nums.length,要找的是数组最小值。
对应升序下标:
0
如果数组中有重复元素,重复元素都按出现次数参与排序位置计算。
例如:
nums = [5,5,4], k = 2
第 2 大仍然是 5。
三路划分可以自然处理这种情况。
正确性证明
我们证明:快速选择算法能返回数组中第 k 个最大的元素。
结论 1:第 k 大元素等价于升序下标为 n - k 的元素
数组长度为 n。
升序排序后,最大元素在下标 n - 1,第二大元素在下标 n - 2。
依次类推,第 k 大元素在下标 n - k。
所以只要找到升序下标 target = n - k 的元素,就找到了第 k 大元素。
结论 2:三路划分后,可以确定答案所在区域
一次划分后,数组被分成三段:
小于 pivot 的区域
等于 pivot 的区域
大于 pivot 的区域
如果 target 在小于区域,答案一定在左边。
如果 target 在大于区域,答案一定在右边。
如果 target 在等于区域,说明排序后该位置上的元素就是 pivot。
因此每次划分后,都能正确缩小搜索范围或直接得到答案。
结论 3:算法不会丢掉答案
算法只会丢弃 target 不可能所在的区域。
当 target < less 时,中间和右边所有元素的排序位置都不可能是 target。
当 target > greater 时,左边和中间所有元素的排序位置都不可能是 target。
当 target 在 [less, greater] 中时,中间区域所有值都等于 pivot,直接返回正确。
所以答案不会被丢掉。
结论 4:BFPRT 与快速选择的选择逻辑相同
BFPRT 只是用更稳定的方式选择 pivot。
选出 pivot 后,仍然使用三路划分,并根据 target 所在区域继续搜索。
因此它和快速选择有相同的正确性逻辑。
得出结论
由结论 1 可知,问题可以转化成查找升序下标 n - k 的元素。
由结论 2 和结论 3 可知,快速选择每次划分都能保留答案所在区域。
由结论 4 可知,BFPRT 也满足相同的查找逻辑。
因此两个解法都能正确返回数组中第 k 个最大的元素。
举例理解
以:
nums = [3,2,1,5,6,4], k = 2
为例。
数组长度 n = 6。
第 2 大元素对应的升序下标是:
target = 6 - 2 = 4
如果完整排序,数组是:
[1,2,3,4,5,6]
下标 4 的元素是:
5
快速选择不需要真的排完整个数组。
它只需要通过划分不断确认下标 4 所在的区域,最后找到这个位置上的值。
再看:
nums = [3,2,3,1,2,4,5,5,6], k = 4
升序排序后是:
[1,2,2,3,3,4,5,5,6]
第 4 大对应升序下标:
9 - 4 = 5
下标 5 的元素是:
4
所以答案是 4。
复杂度分析
设数组长度为 n。
解法一
随机化快速选择每次期望能丢掉一部分元素。
- 平均时间复杂度:
O(n) - 最坏时间复杂度:
O(n^2) - 空间复杂度:
O(1)
三路划分可以更好地处理大量重复元素。
解法二
BFPRT 通过“中位数的中位数”选择基准值,保证每次划分都不会过于偏斜。
- 最坏时间复杂度:
O(n) - 空间复杂度:
O(1)
实际刷题中,随机化快速选择更常用,代码更短;如果严格要求最坏线性时间,可以使用 BFPRT。