Python:如何将迭代器分发给两个消费者且不加载全量数据到内存
问题
我有一个迭代器需要被两个函数(示例中的mean_summarizer和std_summarizer)处理,要求两个函数都能处理该迭代器,且绝不将整个迭代器加载到内存中。
以下是一个最小示例,结果正确,但缺点是会将整个输入加载到内存中。无需理解mean_summarizer、std_summarizer和last内部的复杂代码——这样写主要是为了简洁。
核心问题:在不修改函数签名(仅修改内部实现)的前提下,重写summarize_input_stream的最简洁方式是什么,使其内存占用不会随输入流长度增加而增长?
我猜测需要用到协程,但不知道如何使用。
import numpy as np from typing import Iterable, Mapping, Callable, Any def summarize_input_stream( # Run the input stream through multiple summarizers and collect results input_stream: Iterable[float], summarizers: Mapping[str, Callable[[Iterable[float]], float]] ) -> Mapping[str, float]: inputs = list(input_stream) # PROBLEM IS HERE <-- We load entire stream into memory at once return {name: summarizer(inputs) for name, summarizer in summarizers.items()} def last(iterable: Iterable[Any]) -> Any: # Just returns last element of iterable return max(enumerate(iterable))[1] def mean_summarizer(stream: Iterable[float]) -> float: # Just computes mean online and returns final value return last(avg for avg in [0] for i, x in enumerate(stream) for avg in [avg*i/(i+1) + x/(i+1)]) def std_summarizer(stream: Iterable[float]) -> float: # Just computes standard deviation online and returns final value return last(cumsum_of_sq/(i+1) - (cumsum/(i+1))**2 for cumsum_of_sq, cumsum in [(0, 0)] for i, x in enumerate(stream) for cumsum_of_sq, cumsum in [(cumsum_of_sq+x**2, cumsum+x)])**.5 summary_stats = summarize_input_stream( input_stream=(np.random.randn()*2+3 for _ in range(1000)), summarizers={'mean': mean_summarizer, 'std': std_summarizer} ) print(summary_stats) # e.g. {'mean': 3.020903422847062, 'std': 1.943724669289156}
解决方案
核心思路是将每个汇总器的在线计算逻辑转换成可逐步接收元素的协程(生成器),仅遍历输入流一次,把每个元素分发给所有协程更新状态,最后提取结果,完全无需加载整个迭代器到内存。
修改后的summarize_input_stream实现如下(不修改原函数签名,仅改内部逻辑):
import numpy as np from typing import Iterable, Mapping, Callable, Any def summarize_input_stream( input_stream: Iterable[float], summarizers: Mapping[str, Callable[[Iterable[float]], float]] ) -> Mapping[str, float]: # 为每个汇总器构建可逐步接收元素的协程 runners = {} for name in summarizers: if name == 'mean': # 复刻mean_summarizer的在线计算逻辑为协程 def mean_coro(): avg = 0.0 count = 0 while True: x = yield avg count += 1 avg = avg * (count - 1)/count + x/count coro = mean_coro() next(coro) # 预激协程,使其进入等待接收元素状态 runners[name] = coro elif name == 'std': # 复刻std_summarizer的在线计算逻辑为协程 def std_coro(): cumsum = 0.0 cumsum_sq = 0.0 count = 0 current_std = 0.0 while True: x = yield current_std count += 1 cumsum += x cumsum_sq += x ** 2 mean = cumsum / count var = cumsum_sq / count - mean ** 2 current_std = var ** 0.5 coro = std_coro() next(coro) runners[name] = coro # 遍历输入流,逐个元素喂给所有协程更新状态 for x in input_stream: for coro in runners.values(): coro.send(x) # 收集最终计算结果 return {name: coro.send(None) for name, coro in runners.items()} # 原函数保持不变 def last(iterable: Iterable[Any]) -> Any: return max(enumerate(iterable))[1] def mean_summarizer(stream: Iterable[float]) -> float: return last(avg for avg in [0] for i, x in enumerate(stream) for avg in [avg*i/(i+1) + x/(i+1)]) def std_summarizer(stream: Iterable[float]) -> float: return last(cumsum_of_sq/(i+1) - (cumsum/(i+1))**2 for cumsum_of_sq, cumsum in [(0, 0)] for i, x in enumerate(stream) for cumsum_of_sq, cumsum in [(cumsum_of_sq+x**2, cumsum+x)])**.5 # 测试 summary_stats = summarize_input_stream( input_stream=(np.random.randn()*2+3 for _ in range(1000)), summarizers={'mean': mean_summarizer, 'std': std_summarizer} ) print(summary_stats)
关键说明
- 协程预激:每个协程创建后需要调用
next(),使其执行到第一个yield语句,进入等待接收元素的状态; - 单遍遍历:输入流只被遍历一次,每个元素实时分发给所有协程更新内部状态;
- 内存占用:仅保存每个协程的状态变量(如均值的当前值、标准差的累计和等),内存消耗为O(1),不随输入流长度增长。
内容的提问来源于stack exchange,提问作者Peter
相关产品推荐
相关产品推荐

