Python技术迷

找出无序数组中第 K 个最小元素

在日常开发中,我们经常会遇到这样的问题:一个数组乱糟糟的,顺序完全随机,现在要你找出第 K 小的元素。说白了就是从这个一堆数里捞出“排名第 K 的那位小兄弟”。

这个问题不复杂,但你要是真碰上面试官问你这个,不仅得给出一个能跑的代码,还得让人家看出你代码写得清晰、有思路、知道 tradeoff,最好还能顺便提点优化策略。今天就来聊聊这个经典的题目,顺便也整理下几种解法。

一、最简单粗暴的方法:排序解决战斗

你可能第一个想到的就是:排序啊,把数组排个序不就好了?然后拿第 K 个元素。

defkth_smallest(nums, k):
    nums.sort()
return nums[k - 1]

这个思路很对,代码也好理解,面试官也不会觉得你写错了。但问题是——排序的时间复杂度是 O(n log n),有点“大材小用”的感觉,因为你其实只想知道前 K 小的元素,没必要把整个数组都排一遍。

二、用 Python 内置的 heapq 优雅一点点

如果你写 Python,可能知道 heapq 这个模块,它帮你操作堆,专门干这种“找最小几个”的活儿。

import heapq

defkth_smallest(nums, k):
return heapq.nsmallest(k, nums)[-1]

heapq.nsmallest(k, nums) 会把前 K 小的元素找出来,内部是用堆结构实现的,时间复杂度是 O(n log k),比全排序要省一点时间,尤其当 K 比 n 小得多时。

这个方法最大的优点是:代码非常简洁,而且读起来也很清晰。但要注意,这里其实还是把前 K 个最小值都保存下来了,在空间上会占用 O(k) 的额外内存。

三、用最大堆(Max Heap)节省点空间

虽然 Python 的 heapq 默认是最小堆,但我们可以“曲线救国”地用最大堆来做这个事:

思路是:我们维护一个大小为 K 的最大堆,把前 K 个数先放进去,然后接下来的每个新元素都和堆顶(也就是当前第 K 小的)比,如果更小,就替换。

不过因为 Python 只有最小堆,所以我们得自己取负值来模拟最大堆:

import heapq

defkth_smallest(nums, k):
    max_heap = [-num for num in nums[:k]]
    heapq.heapify(max_heap)

for num in nums[k:]:
if -num > max_heap[0]:  # 相当于 num < -max_heap[0]
            heapq.heappop(max_heap)
            heapq.heappush(max_heap, -num)

return -max_heap[0]

这个方法的优势在于:你只维护一个固定大小的堆,内存消耗低,而且时间复杂度依旧是 O(n log k)。不过因为有负号转换,代码读起来没那么直观。

四、进阶一点:快排的“切片版”做法

这个方法很有意思,而且在某些场景下速度非常快。思路是:用快速排序的 partition 操作来“猜位置”。它不像传统排序那样一遍排到底,而是通过随机选一个 pivot,把数组切成两半:

  • 一边是小于 pivot 的
  • 一边是大于等于 pivot 的

然后你根据 K 和 pivot 的位置继续递归下去,直到找到第 K 小的那个。

import random

defquickselect(nums, k):
defpartition(left, right, pivot_index):
        pivot = nums[pivot_index]
# 把 pivot 移到末尾
        nums[pivot_index], nums[right] = nums[right], nums[pivot_index]
        store_index = left

# 把小于 pivot 的值放到左边
for i in range(left, right):
if nums[i] < pivot:
                nums[store_index], nums[i] = nums[i], nums[store_index]
                store_index += 1

# 把 pivot 放回正确位置
        nums[right], nums[store_index] = nums[store_index], nums[right]
return store_index

    left, right = 0, len(nums) - 1
whileTrue:
        pivot_index = random.randint(left, right)
        pos = partition(left, right, pivot_index)
if pos == k - 1:
return nums[pos]
elif pos < k - 1:
            left = pos + 1
else:
            right = pos - 1

这个算法的平均时间复杂度是 O(n),最坏情况是 O(n^2)(比如 pivot 每次都选得特别烂),但实际表现通常非常好,是很多面试官喜欢听到的解法。

最后,面试官最喜欢听的回答是?

如果你碰到面试官问你这个问题,最好的策略是这样回答:

“这个问题有几种解法,最简单的是直接排序,复杂度是 O(n log n),但效率不算最高。如果 K 比较小,我们可以用最小堆或最大堆来优化到 O(n log k),Python 里直接用 heapq.nsmallest(k, nums)[-1] 就能实现。而如果追求更快的平均时间复杂度,可以考虑用 Quickselect 算法,它是快排的一种变种,平均复杂度是 O(n),实际运行效率也不错。当然,具体用哪个方法也要看数据规模和是否允许改变原数组。”

这种回答不仅让你显得思路清晰,而且表现出你对各种方法的优缺点都很了解,也说明你不是只会调库,而是真的懂算法的程序员。

对编程、职场感兴趣的同学,大家可以联系我微信:golang404,拉你进入“程序员交流群”。
🔥虎哥私藏精品 热门推荐🔥

虎哥作为一名老码农,整理了全网最全《python高级架构师资料合集》。

资料包含了《IDEA视频教程》、《最全python面试题库》、《最全项目实战源码及视频》及《毕业设计系统源码》,总量高达650GB,全部免费领取。