Skip to content

题目链接 · 灵神原题解(署名来源)

核心思路

k 大元素在升序数组中的下标是 nk

  1. nums 中随机选择一个基准元素 pivot。关于为什么要随机,见文末答疑。
  2. 划分 nums。通过交换,把 <pivot 的元素放在 pivot 的左侧,把 pivot 的元素放在 pivot 的右侧。如此划分可以让我们粗略地排序 nums。划分后,pivot 此刻的位置就等于 pivot 在升序数组中的位置。
  3. pivotnums 中的下标为 i
    • 如果 i=nk,那么答案就是 pivot
    • 如果 i>nk,说明答案在 pivot 左侧,我们在其中寻找,回到第一步。
    • 如果 i<nk,说明答案在 pivot 右侧,我们在其中寻找,回到第一步。
    • 这类似 二分查找,只要我们每次能把问题的规模缩小一半,就可以用 O(n) 时间解决(见复杂度分析)。
    • 问题规模缩小后,相当于在 nums 的一个子数组中,继续划分子数组,寻找答案。

然而,如果按照 <pivotpivot 划分数组,这个做法会在数组包含大量重复元素时,划分后的 i 往往是子数组第一个元素的下标,算法会退化至 O(n2)

解决办法:修改第二步,把 < 改成 ,也就是把 pivot 的元素放在 pivot 的左侧,把 pivot 的元素放在 pivot 的右侧。特别地,如果子数组所有元素都相同,这样做可以完美地返回子数组的中心下标(见代码),避免复杂度退化。

具体要如何交换元素?实现细节见代码注释。

python
class Solution:
    def partition(self, nums: List[int], left: int, right: int) -> int:
        """
        在子数组 [left, right] 中随机选择一个基准元素 pivot
        根据 pivot 重新排列子数组 [left, right]
        重新排列后,<= pivot 的元素都在 pivot 的左侧,>= pivot 的元素都在 pivot 的右侧
        返回 pivot 在重新排列后的 nums 中的下标
        特别地,如果子数组的所有元素都等于 pivot,我们会返回子数组的中心下标,避免退化
        """

        # 1. 在子数组 [left, right] 中随机选择一个基准元素 pivot
        i = randint(left, right)
        pivot = nums[i]
        # 把 pivot 与子数组第一个元素交换,避免 pivot 干扰后续划分,从而简化实现逻辑
        nums[i], nums[left] = nums[left], nums[i]

        # 2. 相向双指针遍历子数组 [left + 1, right]
        # 循环不变量:在循环过程中,子数组的数据分布始终如下图
        # [ pivot | <=pivot | 尚未遍历 | >=pivot ]
        #   ^                 ^     ^         ^
        #   left              i     j         right

        i, j = left + 1, right
        while True:
            while i <= j and nums[i] < pivot:
                i += 1
            # 此时 nums[i] >= pivot

            while i <= j and nums[j] > pivot:
                j -= 1
            # 此时 nums[j] <= pivot

            if i >= j:
                break

            # 维持循环不变量
            nums[i], nums[j] = nums[j], nums[i]
            i += 1
            j -= 1

        # 循环结束后
        # [ pivot | <=pivot | >=pivot ]
        #   ^             ^   ^     ^
        #   left          j   i     right

        # 3. 把 pivot 与 nums[j] 交换,完成划分(partition)
        # 为什么与 j 交换?
        # 如果与 i 交换,可能会出现 i = right + 1 的情况,已经下标越界了,无法交换
        # 另一个原因是如果 nums[i] > pivot,交换会导致一个大于 pivot 的数出现在子数组最左边,不是有效划分
        # 与 j 交换,即使 j = left,交换也不会出错
        nums[left], nums[j] = nums[j], nums[left]

        # 交换后
        # [ <=pivot | pivot | >=pivot ]
        #               ^
        #               j

        # 返回 pivot 的下标
        return j

    def findKthLargest(self, nums: list[int], k: int) -> int:
        n = len(nums)
        target_index = n - k  # 第 k 大元素在升序数组中的下标是 n - k
        left, right = 0, n - 1  # 闭区间
        while True:
            i = self.partition(nums, left, right)
            if i == target_index:
                # 找到第 k 大元素
                return nums[i]
            if i > target_index:
                # 第 k 大元素在 [left, i - 1] 中
                right = i - 1
            else:
                # 第 k 大元素在 [i + 1, right] 中
                left = i + 1
cpp
// C++ 版待补充
cpp
class Solution {
    // 在子数组 [left, right] 中随机选择一个基准元素 pivot
    // 根据 pivot 重新排列子数组 [left, right]
    // 重新排列后,<= pivot 的元素都在 pivot 的左侧,>= pivot 的元素都在 pivot 的右侧
    // 返回 pivot 在重新排列后的 nums 中的下标
    // 特别地,如果子数组的所有元素都等于 pivot,我们会返回子数组的中心下标,避免退化
    int partition(vector<int>& nums, int left, int right) {
        // 1. 在子数组 [left, right] 中随机选择一个基准元素 pivot
        int i = left + rand() % (right - left + 1);
        int pivot = nums[i];
        // 把 pivot 与子数组第一个元素交换,避免 pivot 干扰后续划分,从而简化实现逻辑
        swap(nums[i], nums[left]);

        // 2. 相向双指针遍历子数组 [left + 1, right]
        // 循环不变量:在循环过程中,子数组的数据分布始终如下图
        // [ pivot | <=pivot | 尚未遍历 | >=pivot ]
        //   ^                 ^     ^         ^
        //   left              i     j         right

        i = left + 1;
        int j = right;
        while (true) {
            while (i <= j && nums[i] < pivot) {
                i++;
            }
            // 此时 nums[i] >= pivot

            while (i <= j && nums[j] > pivot) {
                j--;
            }
            // 此时 nums[j] <= pivot

            if (i >= j) {
                break;
            }

            // 维持循环不变量
            swap(nums[i], nums[j]);
            i++;
            j--;
        }

        // 循环结束后
        // [ pivot | <=pivot | >=pivot ]
        //   ^             ^   ^     ^
        //   left          j   i     right

        // 3. 把 pivot 与 nums[j] 交换,完成划分(partition)
        // 为什么与 j 交换?
        // 如果与 i 交换,可能会出现 i = right + 1 的情况,已经下标越界了,无法交换
        // 另一个原因是如果 nums[i] > pivot,交换会导致一个大于 pivot 的数出现在子数组最左边,不是有效划分
        // 与 j 交换,即使 j = left,交换也不会出错
        swap(nums[left], nums[j]);

        // 交换后
        // [ <=pivot | pivot | >=pivot ]
        //               ^
        //               j

        // 返回 pivot 的下标
        return j;
    }

public:
    int findKthLargest(vector<int>& nums, int k) {
        srand(time(NULL));
        int n = nums.size();
        int target_index = n - k; // 第 k 大元素在升序数组中的下标是 n - k
        int left = 0, right = n - 1; // 闭区间
        while (true) {
            int i = partition(nums, left, right);
            if (i == target_index) {
                // 找到第 k 大元素
                return nums[i];
            }
            if (i > target_index) {
                // 第 k 大元素在 [left, i - 1] 中
                right = i - 1;
            } else {
                // 第 k 大元素在 [i + 1, right] 中
                left = i + 1;
            }
        }
    }
};

复杂度分析

  • 时间复杂度:期望 O(n),其中 nnums 的长度。在平均情况下,第一次划分(partition)需要处理 n 个元素,第二次平均 n2,第三次平均 n4,依此类推。所以期望时间复杂度为 O(n+n2+n4+)=O(n)
  • 空间复杂度:O(1)

答疑

:如果不随机选择基准元素 pivot,会发生什么?

:比如子数组是有序的,且我们每次都选子数组的第一个(或者最后一个)元素作为 pivot,那么按照算法,j 不变或者移动到最左边,划分是最不均匀的,算法会退化至 O(n2)。随机选 pivot 能使划分在期望意义上是均匀的(j 移动到子数组的中间),保证算法的期望时间复杂度为 O(n)

:代码中的 nums[i] < pivotnums[i] > pivot 能否改成 nums[i] <= pivotnums[i] >= pivot

:这个做法会在子数组所有元素相同时,划分后的 j 是子数组最后一个元素的下标,是最不均匀划分,算法会退化至 O(n2)

:代码中的 i <= j 能否改成 i < j

:这会算错。来看一个例子 nums=[2,1,3]pivot=2。左指针 i=1 移动到 i=2,右指针 j=2 因为不满足 i < j 的条件,无法移动。此时我们交换 2nums[j]=3,得到 [3,1,2],返回 j=2。然而 j=2 左侧有大于 pivot=2 的元素,划分失败。

如果写成 i <= j,那么最终 i=2j=1。此时我们交换 2nums[j]=1,得到 [1,2,3],返回 j=1。这样的划分就是正确的。

附:库函数写法

cpp
class Solution {
public:
    int findKthLargest(vector<int>& nums, int k) {
        ranges::nth_element(nums, nums.end() - k);
        return nums[nums.size() - k];
    }
};

关联题目

如果你理解了划分的过程,那么快速排序算法最难的内容也就理解了。读者可以趁热打铁,完成如下题目:

分类题单

如何科学刷题?

  1. 滑动窗口与双指针(定长/不定长/单序列/双序列/三指针/分组循环)
  2. 二分算法(二分答案/最小化最大值/最大化最小值/第K小)
  3. 单调栈(基础/矩形面积/贡献法/最小字典序)
  4. 网格图(DFS/BFS/综合应用)
  5. 位运算(基础/性质/拆位/试填/恒等式/思维)
  6. 图论算法(DFS/BFS/拓扑排序/基环树/最短路/最小生成树/网络流)
  7. 动态规划(入门/背包/划分/状态机/区间/状压/数位/数据结构优化/树形/博弈/概率期望)
  8. 常用数据结构(前缀和/差分/栈/队列/堆/字典树/并查集/树状数组/线段树)
  9. 数学算法(数论/组合/概率期望/博弈/计算几何/随机算法)
  10. 贪心与思维(基本贪心策略/反悔/区间/字典序/数学/思维/脑筋急转弯/构造)
  11. 链表、二叉树与回溯(前后指针/快慢指针/DFS/BFS/直径/LCA/一般树)
  12. 字符串(KMP/Z函数/Manacher/字符串哈希/AC自动机/后缀数组/子序列自动机)

我的题解精选(已分类)

欢迎关注 B站@灵茶山艾府

本文整理自灵茶山艾府(endlesscheng)的公开内容,仅供个人学习使用

本站仅供个人学习使用,请勿外传