训练数据管道同时承担读取、解码、增强、打乱和组批。多个工作线程调用一个有状态生成器时,索引、随机数和文件句柄可能互相踩踏;简单互斥锁虽然避免同时执行 next(),却把核心路径串行化,也没有保证返回数组不被后续修改。
先按后端和数据形态选择接口
| 接口 | 适合 | 并发所有权 |
|---|---|---|
| tf.data.Dataset | TensorFlow大规模流水线 | 框架调度map、interleave与prefetch |
| keras.utils.PyDataset | 按批索引的Python自定义加载 | Keras管理线程/进程与队列 |
| torch DataLoader | Keras配PyTorch后端 | DataLoader工作进程 |
| 普通数组 | 数据可直接装入内存 | 无需自建并发生成器 |

TensorFlow路径优先组合算子
官方tf.data性能指南建议用并行map/interleave处理读取与转换,再用prefetch让准备下一批与模型计算重叠。并行度可以交给 AUTOTUNE,但cache、shuffle和prefetch都会占用缓冲内存,顺序需按数据大小设计。
PyDataset把工作单元定义成批次
Keras内置训练指南说明,PyDataset通过 __len__ 和 __getitem__ 提供批次,并可配置workers、use_multiprocessing和max_queue_size。索引式接口比共享无限迭代器更容易保证每个批次只出现一次。
class AudioBatches(keras.utils.PyDataset):
def __init__(self, files, labels, batch_size, **kwargs):
super().__init__(**kwargs)
self.files = files
self.labels = labels
self.batch_size = batch_size
self.order = np.arange(len(files))
def __len__(self):
return math.ceil(len(self.files) / self.batch_size)
def __getitem__(self, index):
ids = self.order[index * self.batch_size:(index + 1) * self.batch_size]
return load_batch(self.files[ids]), self.labels[ids]
并发参数交给父类保存
workers、use_multiprocessing和max_queue_size通过构造函数传入 super().__init__,不要在子类里另起一套线程池。批次索引应是纯函数式入口,同一个index在关闭随机增强时返回相同样本集合。
多进程要求对象可以安全复制
启用multiprocessing时,数据集和回调函数需要可序列化;打开的连接、锁、GPU上下文和不可pickle对象不应保存在实例中。不同操作系统的进程启动方式不同,入口脚本使用 if __name__ == '__main__' 保护,并在目标平台实测。
随机增强需要从全局状态中解耦
为epoch、样本索引和工作进程派生随机种子,避免多个副本生成完全相同增强,也避免共享一个无锁随机数生成器。验证与测试关闭随机增强并固定顺序。若在 on_epoch_end 打乱索引,只修改顺序表,不原地移动原始数据。
线程安全还包括对象生命周期
- 每次返回独立批次数组,不复用正在训练的缓冲区。
- 异常时关闭文件、网络连接和子进程。
- 队列有上限,慢消费者不会无限吃内存。
- 最后一个不满批次的策略与分布式副本一致。
- 增强函数不修改共享元数据。
先验收数据覆盖,再比较速度
在样本中临时加入唯一ID,跑多个epoch并统计遗漏、重复和跨集合污染;固定种子重复运行,比较批次ID和张量摘要。并发数变化后标签仍需与输入一一对应。正确性测试通过,再测每秒样本、GPU空闲、CPU与I/O。
锁什么时候仍然合适
只读第三方库内部有不可重入调用、又无法改造时,可以在最小临界区加锁保护该调用。锁不应包住整个解码与增强流程,也不应被包装成性能优化。若所有工作最终排队经过同一把锁,增加线程只增加调度开销。
数据管道的首要指标是每个样本与标签正确且可复现;吞吐是在这条约束内优化的第二目标。
并发加载首先要定义失败如何传播
某个 worker 读取损坏文件、增强函数抛出异常或超时后,主训练循环必须收到明确错误并停止或按策略跳过。若后台线程静默退出,队列可能永久等待;若反复重试同一坏样本,则会制造假死。为每个批次携带索引和尝试次数,限制重试并把失败样本写入独立报告。
覆盖性比吞吐更早验收。一个 epoch 结束后核对样本编号:应出现的是否全部出现,是否重复,最后不足批量如何处理。多进程下不要依赖进程间共享的普通计数器;由确定性索引生成批次,或在主进程汇总返回标识。随机增强的种子可由全局种子、epoch 和样本 ID 派生,使不同并发度仍能复现实验。
队列水位体现生产者与消费者关系
持续为空说明准备速度不足,长期满载说明预取可能占用过多内存或设备处理更慢。分别测量文件读取、解码、增强、组批和设备等待,找到真正瓶颈。小文件随机读可通过合并格式或缓存改善,大图片解码可能适合并行,纯 Python 计算则要考虑进程开销。
训练被取消或发生异常时,关闭队列、通知 worker、等待有期限退出,再释放文件和共享内存。把生成器包上一把锁只能保护同一进程中的临界区,无法解决跨进程序列化、资源继承和异常回收。
最终同时报告样本每秒、首批等待、CPU 与内存、设备空闲、结果覆盖和重复运行一致性。数据管道的价值不是让进度条更快,而是稳定、完整地把预期样本交给模型。
并发度改变批次到达顺序时,含状态的预处理、在线统计或依赖顺序的增强可能得到不同结果。尽量让样本转换成为纯函数,把需要累积的状态放在明确的单线程阶段;否则要把顺序也作为实验配置记录并测试。
本文《Keras数据加载并发:何时用tf.data、PyDataset与进程队列》由 xkmchenmu 发布于 xkmchenmu Blog。 转载请保留原文链接并注明出处。
支付宝扫一扫