算法题目
有 m 个排序好的数组(从小到大排序),组成的一个二维数组,请你找到最小的第 k 个数。Python 实现,例如数组 \[1, 3, 5, 7, 8, 9, 2, 4, 6],k=5。
题目分析
题意:①全局无序,②子数组有序
给定 m 个各自升序排列的一维数组构成二维数组,全局所有元素中找出第 k 小的数。
示例:[[1,3,5],[7,8,9],[2,4,6]],所有元素排序:[1,2,3,4,5,6,7,8,9],k=5 → 结果为 5。
介绍三种解法:(从暴力到最优)
方法 1:暴力合并排序(最简单,适合数据量小)
- 把所有数组摊平成一维,整体排序取下标 k-1,一行就能写完。
- 缺点:时间复杂度
(O(Nlog N)),N为总元素数,大数据低效。
python
def kth_smallest_brute(matrix, k):
arr = []
for row in matrix:
arr.extend(row)
arr.sort()
return arr[k-1]
# 测试
if __name__ == "__main__":
data = [[1, 3, 5], [7, 8, 9], [2, 4, 6]]
print(kth_smallest_brute(data, 5)) # 输出 5
方法 2:最小堆(优先队列,经典多路归并)
核心思路:多路有序数组归并,用小顶堆每次弹出最小值,弹出第 k 次即为答案。
- 先把每个数组第一个元素入堆,记录:(值, 数组索引, 元素在数组内下标)
- 循环弹出堆顶最小元素,计数 + 1;若计数 ==k 直接返回
- 如果弹出元素所在数组还有下一个元素,继续入堆
- 复杂度:
(O(klog m)),m 为数组个数,k 不大时效率极高。 - 说明:需要用到Python自带的包
heapq,有取巧嫌疑。
python
import heapq
def kth_smallest_heap(matrix, k):
heap = []
m = len(matrix)
for i in range(m): # 初始化堆:每个数组第一个元素入堆
val = matrix[i][0]
heapq.heappush(heap, (val, i, 0))
cnt = 0
while heap:
val, arr_idx, elem_idx = heapq.heappop(heap)
cnt += 1
if cnt == k:
return val
if elem_idx + 1 < len(matrix[arr_idx]): # 当前数组还有下一个元素则入堆
next_val = matrix[arr_idx][elem_idx + 1]
heapq.heappush(heap, (next_val, arr_idx, elem_idx + 1))
if __name__ == "__main__":
test_arr = [[1, 3, 5], [7, 8, 9], [2, 4, 6]]
print(kth_smallest_heap(test_arr, 5)) # 5
方法 3:二分查找最优解法(推荐大数据量)
思路:利用值域二分(自己想到的方法),
- 最小值
left = 所有数组首元素最小值,最大值 right = 所有数组尾元素最大值 mid = (left+right)//2,统计二维数组中≤mid 的元素总数countcount < k:说明答案在右半区间left=mid+1count ≥k:答案在左半区间right=mid- 最终 left=right 就是第 k 小数
- 复杂度 :
(O(m log S)),S 为数值值域范围,性能上限最高。
python
import bisect
def kth_smallest_binary(matrix, k):
# 确定二分上下界
left = min(row[0] for row in matrix)
right = max(row[-1] for row in matrix)
def count_less_or_equal(x):
"""统计所有数组中 <= x 的元素个数(每行有序,二分加速)"""
total = 0
for row in matrix:
total += bisect.bisect_right(row, x) # bisect_right 返回插入点,即小于等于x的数量
return total
while left < right:
mid = (left + right) // 2
cnt = count_less_or_equal(mid)
if cnt < k:
left = mid + 1
else:
right = mid
return left
if __name__ == "__main__":
arr = [[1, 3, 5], [7, 8, 9], [2, 4, 6]]
print(kth_smallest_binary(arr, 5)) # 5
总结
- 数据很小 → 暴力一行版;
- 面试常规、多路归并考点 → 堆解法;
- 海量数据、追求极致效率 → 值域二分法。