Python生成器嵌套调用问题:如何复用另一生成器输出
解决方案:嵌套生成器串联实现
我看你是想把两个生成器串起来:先用res_gen批量生成预处理好的图像块数据,再把这些数据传给train_datagen生成带噪声的训练batch,最终输出train_datagen的结果对吧?你的思路是对的,但sub_gen的写法有语法问题,而且还有一些可以优化的点,我来帮你修正:
首先看核心问题:嵌套生成器的正确写法
你原来的sub_gen里yield each train_datagen(r)是无效语法,因为train_datagen返回的是一个生成器对象,你需要遍历这个生成器的每一个输出,再逐个yield出去,而不是直接yield生成器本身。
修正后的sub_gen应该是这样:
def sub_gen(batch_size=4): # 遍历res_gen的每一次输出(即预处理好的res1数组) for res1 in res_gen(): # 把res1传给train_datagen,遍历它生成的每个batch for batch in train_datagen(res1, batch_size=batch_size): yield batch
这样每次调用sub_gen(),就会先从res_gen拿到预处理好的图像块,再通过train_datagen生成带噪声的输入和干净标签的batch,逐个输出给训练流程。
优化res_gen的多进程使用
另外你的res_gen里每次循环都创建新的Pool(num_threads),这会频繁创建和销毁进程,非常影响效率。建议把Pool的创建移到循环外面,用上下文管理器来管理:
from multiprocessing import Pool import numpy as np # 假设file_list和sigma是已经定义好的全局变量或传入参数 file_list = [...] # 你的文件路径列表 sigma = 25 # 噪声标准差,根据你的需求调整 # Generator 1:批量生成预处理后的图像块 def res_gen(num_threads=4): # 只创建一次进程池,复用它 with Pool(num_threads) as p: while True: # 分批次处理file_list for i in range(0, len(file_list), num_threads): batch_files = file_list[i:min(i+num_threads, len(file_list))] # 用进程池并行生成图像块 patch = p.map(gen_patches, batch_files) # 合并所有图像块 res = [] for x in patch: res += x # 预处理:reshape、归一化 res1 = np.array(res).reshape((-1, res[0].shape[0], res[0].shape[1], 1)) res1 = res1.astype('float32') / 255.0 yield res1
这里做了几个优化:
- 用
with Pool(...)上下文管理器,自动管理进程池的创建和销毁,避免资源泄漏 - 用
-1自动计算batch维度,代码更简洁 - 明确提取
batch_files,逻辑更清晰
完整的可运行代码示例
把所有部分整合起来,完整代码如下:
from multiprocessing import Pool import numpy as np # 假设这些是你的全局配置/依赖 file_list = ["file1.jpg", "file2.jpg", ...] # 替换成你的实际文件列表 sigma = 25 # 噪声强度 # 假设gen_patches是你已经实现的函数,输入文件路径,返回图像块列表 def gen_patches(file_path): # 示例逻辑:读取文件,生成图像块 # 这里替换成你自己的实现 return [np.random.rand(64,64) for _ in range(10)] # Generator 1:批量生成预处理后的图像块 def res_gen(num_threads=4): with Pool(num_threads) as p: while True: for i in range(0, len(file_list), num_threads): batch_files = file_list[i:min(i+num_threads, len(file_list))] patch = p.map(gen_patches, batch_files) res = [] for x in patch: res += x res1 = np.array(res).reshape((-1, res[0].shape[0], res[0].shape[1], 1)) res1 = res1.astype('float32') / 255.0 yield res1 # Generator 2:生成带噪声的训练batch def train_datagen(res1, batch_size=4): indices = list(range(res1.shape[0])) while True: np.random.shuffle(indices) for i in range(0, len(indices), batch_size): batch_indices = indices[i:i+batch_size] ge_batch_y = res1[batch_indices] # 添加高斯噪声 noise = np.random.normal(0, sigma/255.0, ge_batch_y.shape) ge_batch_x = ge_batch_y + noise yield ge_batch_x, ge_batch_y # 顶层生成器:串联两个生成器 def sub_gen(batch_size=4): for res1 in res_gen(): for batch in train_datagen(res1, batch_size=batch_size): yield batch # 使用示例: if __name__ == "__main__": gen = sub_gen(batch_size=4) # 获取一个batch x, y = next(gen) print(f"Input shape: {x.shape}, Label shape: {y.shape}")
关键说明
- 当你在训练模型时(比如Keras的
model.fit_generator),直接传入sub_gen()即可,它会源源不断地输出训练用的(x, y) batch - 注意
gen_patches函数需要是可被进程池序列化的(不能用lambda或者无法pickle的对象),如果有问题可以考虑用multiprocessing.get_context('spawn')来创建Pool - 如果你的
file_list很大,res_gen会循环往复地处理所有文件,train_datagen会在每个res1内部打乱后生成batch,直到res1的所有数据都被用完,再自动获取下一个res1继续处理
内容的提问来源于stack exchange,提问作者Prachi
相关产品推荐
相关产品推荐

