LEETCODE 215Medium

数组中的第K个最大元素

快速排序每趟 partition 都能确定一个元素的最终位次,只要这个位次恰好是 n-k,剩下的排序都可以不做。

问题拆解

找出数组中第 k 个最大的元素,题目明确要求 O(n) 时间——直接排序取下标是 O(n log n),不达标。

排序之所以“超额完成”,是因为它把所有元素的位次都排清楚了,而我们只关心一个位次:升序下标 n - k。快速排序的 partition 恰好有个性质:一趟结束后,pivot 落在它排序后的最终位置上。如果这个位置正好是 n - k,答案直接就是它;偏大就去左半边找,偏小就去右半边找——每次只递归一侧,这就是快速选择。

排序确定所有元素的位次,快速选择只确定一个:每趟 partition 后丢掉一半,n + n/2 + n/4 + … = 2n,平均 O(n)

有一个坑必须提:力扣的测试数据里有全等元素、已排序等针对性用例,固定取第一个或最后一个元素当 pivot 会退化成 O(n²) 直接超时,pivot 必须随机选。

快速选择

private final Random random = new Random();

public int findKthLargest(int[] nums, int k) {
    int target = nums.length - k; // 第 k 大 = 升序下标 n-k
    int left = 0, right = nums.length - 1;
    while (true) {
        int p = partition(nums, left, right);
        if (p == target) return nums[p];
        if (p < target) left = p + 1;
        else right = p - 1;
    }
}

private int partition(int[] nums, int left, int right) {
    // 随机 pivot,避免被构造数据卡成 O(n^2)
    int r = left + random.nextInt(right - left + 1);
    swap(nums, r, right);
    int pivot = nums[right];
    int i = left; // [left, i) 都小于 pivot
    for (int j = left; j < right; j++) {
        if (nums[j] < pivot) {
            swap(nums, i++, j);
        }
    }
    swap(nums, i, right);
    return i;
}

private void swap(int[] nums, int i, int j) {
    int t = nums[i];
    nums[i] = nums[j];
    nums[j] = t;
}
import random

def findKthLargest(nums: list[int], k: int) -> int:
    def partition(left: int, right: int) -> int:
        # 随机 pivot,避免被构造数据卡成 O(n^2)
        r = random.randint(left, right)
        nums[r], nums[right] = nums[right], nums[r]
        pivot = nums[right]
        i = left  # [left, i) 都小于 pivot
        for j in range(left, right):
            if nums[j] < pivot:
                nums[i], nums[j] = nums[j], nums[i]
                i += 1
        nums[i], nums[right] = nums[right], nums[i]
        return i

    target = len(nums) - k  # 第 k 大 = 升序下标 n-k
    left, right = 0, len(nums) - 1
    while True:
        p = partition(left, right)
        if p == target:
            return nums[p]
        if p < target:
            left = p + 1
        else:
            right = p - 1
func findKthLargest(nums []int, k int) int {
    target := len(nums) - k // 第 k 大 = 升序下标 n-k
    left, right := 0, len(nums)-1
    for {
        p := partition(nums, left, right)
        if p == target {
            return nums[p]
        }
        if p < target {
            left = p + 1
        } else {
            right = p - 1
        }
    }
}

func partition(nums []int, left, right int) int {
    // 随机 pivot,避免被构造数据卡成 O(n^2)
    r := left + rand.Intn(right-left+1)
    nums[r], nums[right] = nums[right], nums[r]
    pivot := nums[right]
    i := left // [left, i) 都小于 pivot
    for j := left; j < right; j++ {
        if nums[j] < pivot {
            nums[i], nums[j] = nums[j], nums[i]
            i++
        }
    }
    nums[i], nums[right] = nums[right], nums[i]
    return i
}
pub fn find_kth_largest(nums: Vec<i32>, k: i32) -> i32 {
    let mut nums = nums;
    let target = nums.len() - k as usize; // 第 k 大 = 升序下标 n-k
    let (mut left, mut right) = (0usize, nums.len() - 1);
    // 标准库没有直接的随机数,用系统时钟造一个简易随机源
    let mut seed = std::time::SystemTime::now()
        .duration_since(std::time::UNIX_EPOCH)
        .unwrap()
        .subsec_nanos() as usize;
    loop {
        // 随机 pivot,避免被构造数据卡成 O(n^2)
        seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1);
        let r = left + seed % (right - left + 1);
        nums.swap(r, right);
        let pivot = nums[right];
        let mut i = left; // [left, i) 都小于 pivot
        for j in left..right {
            if nums[j] < pivot {
                nums.swap(i, j);
                i += 1;
            }
        }
        nums.swap(i, right);
        if i == target {
            return nums[i];
        }
        if i < target {
            left = i + 1;
        } else {
            right = i - 1;
        }
    }
}

partition 用的是 Lomuto 写法:i 维护“小于 pivot 的区域”的右边界,j 扫描时遇到更小的就换进去,最后把 pivot 换到 i 处——此时它左边全比它小、右边全不小于它,位次确定。主循环不用递归,改成收缩 [left, right] 区间的迭代,省掉栈开销。Rust 标准库没有随机数,力扣环境也不能引 rand crate,这里用时钟种子加线性同余凑一个够用的随机源。

大小为 k 的小顶堆

如果数据是流式的、或者 k 远小于 n,堆是更稳的选择:维护一个只装 k 个元素的小顶堆,堆顶是“当前第 k 大”。新元素比堆顶大,就替换堆顶;否则它连前 k 名都进不了,直接丢弃。

public int findKthLargest(int[] nums, int k) {
    PriorityQueue<Integer> heap = new PriorityQueue<>(); // 小顶堆
    for (int x : nums) {
        if (heap.size() < k) {
            heap.offer(x);
        } else if (x > heap.peek()) {
            heap.poll(); // 挤掉当前第 k 大
            heap.offer(x);
        }
    }
    return heap.peek();
}
import heapq

def findKthLargest(nums: list[int], k: int) -> int:
    heap = nums[:k]
    heapq.heapify(heap)  # 小顶堆
    for x in nums[k:]:
        if x > heap[0]:
            heapq.heapreplace(heap, x)  # 挤掉当前第 k 大
    return heap[0]
type minHeap []int

func (h minHeap) Len() int            { return len(h) }
func (h minHeap) Less(i, j int) bool  { return h[i] < h[j] }
func (h minHeap) Swap(i, j int)       { h[i], h[j] = h[j], h[i] }
func (h *minHeap) Push(x interface{}) { *h = append(*h, x.(int)) }
func (h *minHeap) Pop() interface{} {
    old := *h
    n := len(old)
    x := old[n-1]
    *h = old[:n-1]
    return x
}

func findKthLargest(nums []int, k int) int {
    h := &minHeap{}
    heap.Init(h)
    for _, x := range nums {
        if h.Len() < k {
            heap.Push(h, x)
        } else if x > (*h)[0] {
            (*h)[0] = x       // 直接替换堆顶再下沉
            heap.Fix(h, 0)
        }
    }
    return (*h)[0]
}
use std::cmp::Reverse;
use std::collections::BinaryHeap;

pub fn find_kth_largest(nums: Vec<i32>, k: i32) -> i32 {
    let k = k as usize;
    // BinaryHeap 是大顶堆,套 Reverse 变成小顶堆
    let mut heap: BinaryHeap<Reverse<i32>> = BinaryHeap::with_capacity(k);
    for x in nums {
        if heap.len() < k {
            heap.push(Reverse(x));
        } else if x > heap.peek().unwrap().0 {
            heap.pop(); // 挤掉当前第 k 大
            heap.push(Reverse(x));
        }
    }
    heap.peek().unwrap().0
}

方向别搞反:求第 k 顶堆——堆顶是 k 个候选里最小的,正是随时可以被淘汰的那个。时间 O(n log k),空间 O(k),虽然渐近上不如快速选择,但不依赖随机性、不修改原数组,工程上常常更受欢迎。

复杂度

指标 复杂度 原因
时间 平均 O(n) / 堆 O(n log k) 快选每趟丢一半,几何级数收敛到 2n;堆每个元素至多一次 log k 的调整
空间 O(1) / O(k) 快选原地划分;堆只存 k 个元素

可以迁移的模式

  • “第 k 个 / 前 k 个”问题先想快速选择和堆,不要惯性全排序;
  • partition 的返回值是 pivot 的最终位次,这个性质还能用来找中位数(如 295 的离线版);
  • pivot 随机化不是可选优化:面对对抗性数据,它是快选保住平均复杂度的前提。

快速选择赢在渐近复杂度,堆赢在稳定和流式友好——两个都值得放进工具箱,按数据形态取用。