数组中的第 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 的做法是:

  1. 5 个元素分成一组。
  2. 每组内部排序,取这一组的中位数。
  3. 把所有组的中位数放到数组前面。
  4. 递归选择这些中位数的中位数作为 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。