You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.14 08:37:59