NumPy中将迭代器输出转为数组拼接的最优简洁实现方案问询
你第二种写法的性能已经比逐次concatenate的方案好一个量级,但还可以写得更简洁,同时兼容batch维度大小不一致的场景。
首先说下第一种写法慢的根本原因:每次调用np.concatenate都会重新申请整块内存,把已有的全量数据和新批次数据拷贝到新内存里,总拷贝量随批次数量呈平方级增长,数据量越大速度越慢,生产环境不要这么写。
第二种用列表收集所有批次再转数组的思路是对的:列表只存储每个批次数组的引用,不会额外拷贝数据,最后一次性做内存分配和全量拷贝,时间复杂度是线性的,性能已经摸到了纯numpy实现的上限。只是写法上多了一步没必要的reshape,而且对不等长批次的兼容性差。
推荐写法
直接把生成器转成列表后传入np.concatenate,一步完成拼接,不需要手动调整维度:
data_set = np.concatenate(list(generator), axis=0)
这个写法的优势:
- 性能和你当前的列表解包方案完全一致,没有额外开销
- 不需要计算目标shape做reshape,代码更短、逻辑更直观
- 自动适配不同批次样本数不一致的情况,只要所有批次除第0维外的其他维度(图片高、宽、通道数)统一,就能正常运行
注意事项
- 这个方法仅适用于生成器产出的全量数据可以完整放进内存的场景,如果生成器是无限输出的(比如开启了repeat的数据集迭代器),会直接占满内存。
- 不要为了所谓的“省内存”绕回逐次concatenate的写法,
list(generator)存储的只是数组引用,内存开销可以忽略不计。 - 不需要尝试用
np.fromiter实现类似逻辑,这个接口处理多维数组需要提前固定dtype和输出shape,写法冗余,性能也没有优势。
内容的提问来源于stack exchange,提问作者Johnzy
相关产品推荐
相关产品推荐

