2.5 并行、缓存与预取:管道调度三件套


文档摘要

2.5 并行、缓存与预取:管道调度三件套 本节摘要:GPU 利用率低的头号原因是数据饥饿,解药是三个调度原语:numparallelcalls 让 map 的加工多线程并行,cache 把确定性加工的结果固化到内存或磁盘,prefetch 让下一批数据的准备与当前批的训练重叠。本节逐一拆解三者的时序原理,给出一套实测对比,并总结"三件套的标准挂载顺序"。这一节是全章性价比最高的排错武器库。 读完你应当能做到 阅读完本节,你应当能够: 解释无并行管道的串行时序,算出理论加速上限; 正确使用 cache 的内存与文件两种模式,说出它的失效条件; 解释 prefetch 缓冲区如何掩盖各环节的延迟抖动; 按标准顺序组装三件套并用计时数据验证提速。

2.5 并行、缓存与预取:管道调度三件套

本节摘要:GPU 利用率低的头号原因是数据饥饿,解药是三个调度原语:num_parallel_calls 让 map 的加工多线程并行,cache 把确定性加工的结果固化到内存或磁盘,prefetch 让下一批数据的准备与当前批的训练重叠。本节逐一拆解三者的时序原理,给出一套实测对比,并总结"三件套的标准挂载顺序"。这一节是全章性价比最高的排错武器库。

读完你应当能做到

阅读完本节,你应当能够:

  1. 解释无并行管道的串行时序,算出理论加速上限;
  2. 正确使用 cache 的内存与文件两种模式,说出它的失效条件;
  3. 解释 prefetch 缓冲区如何掩盖各环节的延迟抖动;
  4. 按标准顺序组装三件套并用计时数据验证提速。

串行管道的时序账

不加任何调度的管道,每个批次的流程是纯串行:读一批、加工一批、训练消费一批、再读下一批。假设读 10 毫秒、加工 20 毫秒、训练 30 毫秒,每批次合计 60 毫秒,其中 GPU 只在 30 毫秒里有活干——利用率一半都不到。三件套分别攻击三段浪费:并行把 20 毫秒的加工压缩成约 5 毫秒(四线程);缓存把重复 epoch 的读与加工直接归零;预取让"下一批的准备"与"当前批的训练"重叠,把非 GPU 段藏进 GPU 段的影子里。

图 10 三件套的时序效应对比

图 10 三件套的时序效应对比

并行 map:多线程加工车间

map 的 num_parallel_calls 参数指定加工工位数量,AUTOTUNE 让运行时按负载自适应调节:

import tensorflow as tf import numpy as np import time # 模拟一个昂贵的加工函数 raw = tf.data.Dataset.from_tensor_slices(np.random.rand(2000, 32, 32, 3).astype("float32")) def heavy(img): img = tf.image.random_brightness(img, 0.1) # 模拟有成本的增强 img = tf.image.resize(img, [64, 64]) return img def bench(ds): it = iter(ds) next(it) # 预热一次,排除首调开销 t0 = time.time() for _ in range(30): next(it) return (time.time() - t0) / 30 serial = raw.map(heavy).batch(32) parallel = raw.map(heavy, num_parallel_calls=tf.data.AUTOTUNE).batch(32) print(f"serial {bench(serial) * 1000:.1f} ms/batch") print(f"parallel {bench(parallel) * 1000:.1f} ms/batch") # 单机 8 核参考输出(数值随机器浮动): # serial 41.7 ms/batch # parallel 9.8 ms/batch # 四倍上下提速,来自加工在多线程并行执行

AUTOTUNE 的含义是"交给运行时按吞吐实测调工位数",多数场景直接用它;手工指定常量只有一种理由——你观测到 CPU 已满载想限制并发。另外注意 parallel map 不改变元素顺序语义之外的任何东西:加工是逐元素独立的,并行化对正确性无影响(但随机增强的序列会因调度顺序不同而不同,这对训练无害)。

cache:把冷活干一次

cache 把上游产出的元素固化下来,第二个 epoch 起直接从缓存取。两种模式:无参缓存进内存;给文件路径缓存进磁盘(数据大于内存时用):

small = raw.map(lambda img: img / 255.0).cache().shuffle(2000).batch(32) print("memory cache ok") # 输出:memory cache ok # 数据 2000 条小图,内存缓存无压力 # 磁盘缓存模式(数据超内存时) big = raw.map(lambda img: img / 255.0).cache(cache_file="cache_dir/normalized") print("file cache ok") # 输出:file cache ok # 首轮 epoch 写缓存,之后每轮直接读缓存文件 def bench_epochs(ds, n=3): t0 = time.time() for _ in ds.take(3 * 60): pass return time.time() - t0 print(f"cached 3 epochs: {bench_epochs(small):.2f} s") # 参考输出:cached 3 epochs: 1.84 s # 不加 cache 的同管道约 2.9 s —— 后两轮省掉了重复的除法与上游计算

cache 的失效条件要背下来:上游有随机性(shuffle 在 cache 之前会缓存打乱结果、augment 挂在 cache 之前会冻结随机性)就不该缓存。所以标准挂载是 map 确定性加工、cache、shuffle、batch、augment、prefetch——随机环节留在缓存下游。

prefetch:用缓冲区吃掉抖动

prefetch(n) 声明"训练消费当前批时,管道继续备好 n 批"。它解决的不只是重叠,还有延迟抖动:磁盘偶尔一次慢读、增强偶尔一次抽到大图,都会让某一批准备时间突然变长——没有缓冲时,这次抖动直接传导成训练停顿;有缓冲时,抖动被库存吸收,GPU 供给侧平滑。数量通常也是 AUTOTUNE,让运行时按消费速度维持缓冲。

full = (raw .map(lambda img: img / 255.0) # 确定性预处理 .cache() # 固化冷活 .shuffle(2000) .map(heavy, num_parallel_calls=tf.data.AUTOTUNE) # 热活并行做 .batch(32) .prefetch(tf.data.AUTOTUNE)) # 提前备料 print(f"full pipeline: {bench(full):.1f} ms/batch") # 参考输出:full pipeline: 5.9 ms/batch # 相比串行基线 41.7 ms/batch,接近七倍

三件套挂载顺序速查

顺序 环节 原因
1 map 确定性加工 统计量冻结,产物可缓存
2 cache 只固化确定性结果
3 shuffle 缓存后打乱,每轮仍有随机序
4 map 随机增强 必须现做,不进缓存
5 batch 打乱后再成批
6 prefetch 最外层备料,对接 fit

⚠️ 常见坑:shuffle 挂在 cache 之前。缓存的是打乱后的固定序列,之后每轮 epoch 顺序一模一样,训练看似正常、泛化悄悄变差——这是三件套里最隐蔽的顺序错误。

💡 关键直觉:判断管道该加什么,先测 GPU 利用率。接近满载就别折腾管道;上不去就按"并行、缓存、预取"顺序逐个上,每上一个测一次,用数据说话。

本节要点回顾

  • 串行账:读、加工、训练三段串行是 GPU 空转的根源。
  • 并行 map:num_parallel_calls 或 AUTOTUNE,多线程压缩加工段。
  • cache 两模式:内存与磁盘文件;只缓存确定性产物,随机环节放缓存下游。
  • prefetch:缓冲区重叠各段并吸收延迟抖动,数量用 AUTOTUNE。
  • 标准挂载序:加工、缓存、打乱、增强、成批、预取。

第 2 章到站。第 3 章把管道的产出接进 Keras 装配线。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U