Stream数据中位数计算优化:解决输入循环导致的TLE问题
解决大输入流中位数计算超时问题
哥们,我太懂你这种被百万级输入卡超时的痛苦了!500万条数据用普通的for _ in range(n)循环读输入,频繁的IO操作绝对是拖慢速度的元凶。咱们分两步来解决这个问题,先搞定输入读取,再优化中位数计算:
一、先把输入读取速度拉满!
普通的input()函数每次调用都会触发一次系统IO,对于百万级别的数据来说,这种频繁的IO开销会直接把时间耗光。最快的方式是一次性读取所有输入,用sys.stdin.read()直接把整个输入流读进来,再分割成数字列表——亲测这个方法比循环读快至少一个数量级。
代码示例:
import sys def main(): # 一次性读取所有输入,自动分割所有空白符(包括换行、空格) all_nums = list(map(int, sys.stdin.read().split())) n = len(all_nums) # 接下来处理中位数...
不管你的输入是每行一个数字,还是用空格分隔的一行数字,这个方法都能完美处理,完全避开了循环读输入的低效问题。
二、优化中位数的计算逻辑
解决了输入问题后,咱们再看中位数计算。因为输入长度是奇数,中位数就是第(n+1)//2大的数字,这里有两种高效方案:
方案1:一次性处理所有数据——快速选择算法
如果不需要实时计算(可以先读完所有数据再算),快速选择算法的平均时间复杂度是O(n),比排序的O(n log n)快很多,适合处理超大数据集。
你可以用Python标准库的heapq.nlargest来实现(底层是堆优化的快速选择),代码简洁又高效:
import heapq import sys def main(): all_nums = list(map(int, sys.stdin.read().split())) n = len(all_nums) # 找第(n+1)//2大的数,取最后一个元素就是中位数 median = heapq.nlargest((n+1)//2, all_nums)[-1] print(median)
如果想自己实现更纯粹的快速选择(避免递归栈溢出,建议用迭代版),也可以这么写:
import random import sys def find_kth_largest(nums, k): left, right = 0, len(nums)-1 while left <= right: pivot_idx = random.randint(left, right) pivot = nums[pivot_idx] # 把pivot移到最右边 nums[pivot_idx], nums[right] = nums[right], nums[pivot_idx] # 分区:大于pivot的放左边,小于的放右边 store_idx = left for i in range(left, right): if nums[i] > pivot: nums[store_idx], nums[i] = nums[i], nums[store_idx] store_idx += 1 # 把pivot移到正确的位置 nums[store_idx], nums[right] = nums[right], nums[store_idx] if store_idx + 1 == k: return nums[store_idx] elif store_idx + 1 < k: left = store_idx + 1 else: right = store_idx -1 def main(): all_nums = list(map(int, sys.stdin.read().split())) n = len(all_nums) median = find_kth_largest(all_nums, (n+1)//2) print(median)
方案2:实时流计算——双堆法
如果必须边读边计算中位数(比如数据是真正的流,不能一次性存下),可以用大顶堆+小顶堆的组合:
- 大顶堆存输入中较小的一半元素(用负数模拟,因为Python的
heapq默认是小顶堆) - 小顶堆存输入中较大的一半元素
- 始终保持大顶堆的大小比小顶堆大1(因为输入长度是奇数),这样大顶堆的堆顶就是中位数
代码示例:
import heapq import sys def main(): max_heap = [] # 存较小的一半,用负数模拟大顶堆 min_heap = [] # 存较大的一半 all_nums = list(map(int, sys.stdin.read().split())) for num in all_nums: # 先把当前数加入大顶堆 heapq.heappush(max_heap, -num) # 平衡两个堆:把大顶堆的最大元素移到小顶堆 heapq.heappush(min_heap, -heapq.heappop(max_heap)) # 保证大顶堆的大小始终比小顶堆大1 if len(max_heap) < len(min_heap): heapq.heappush(max_heap, -heapq.heappop(min_heap)) # 中位数就是大顶堆的堆顶(取负数还原) print(-max_heap[0])
最后总结
- 优先优化输入读取:用
sys.stdin.read()一次性读入所有数据,这是解决你当前TLE问题最直接的手段。 - 选择合适的中位数计算方法:一次性处理选快速选择,实时流处理选双堆法,两种方法的时间复杂度都能满足5秒的时间限制。
内容的提问来源于stack exchange,提问作者Gareth Ma
相关产品推荐
相关产品推荐

