根本原因是num_parallel_calls默认为none导致map()串行阻塞,应设为tf.data.autotune;正确流水线顺序是map()→cache()→shuffle()→batch()→prefetch()。

tf.data 流水线卡在 map() 里,GPU 利用率忽高忽低
常见现象是训练刚开始几轮很快,之后 map() 耗时陡增,prefetch() 像没起作用。根本原因不是 CPU 不够,而是 num_parallel_calls 默认为 None(即单线程执行),尤其当预处理含 tf.io.read_file、tf.image.decode_jpeg 或自定义 Python 函数时,会退化成串行阻塞。
- 显式设
num_parallel_calls=tf.data.AUTOTUNE,别硬写数字(不同机器核数不同) - 若用了
tf.py_function,必须加num_parallel_calls,否则自动降级为单线程 - 避免在
map()里做重 IO 操作(比如每次读 config 文件),提前加载到内存或用tf.lookup.StaticHashTable -
cache()放在map()后、shuffle()前——对小数据集有效;若内存吃紧,改用cache('/path/to/cache')写磁盘
tf.function 没编译训练循环,还在用 eager 模式逐行解释
TensorFlow 2.x 默认启用 eager execution,写起来方便,但每步都过 Python 解释器,CPU 开销大。真正提速靠 @tf.function 把训练步骤编译成静态图,跳过解释过程。
- 把
@tf.function加在训练 step 函数上(如train_step),不是加在 episode 循环里 - 不要反复调用
@tf.function装饰的函数却传不同 shape 的输入,否则触发多次 tracing,生成冗余图 - 若模型含动态控制流(如
tf.while_loop),XLA 编译可能失效,需检查concrete_function是否被复用
clear_session() 没调用,计算图持续膨胀
在多轮训练(如 Q-learning 的 episode 循环)中反复 save_model() 却不清理状态,会导致默认图、变量引用、梯度路径不断累积,后续 train_on_batch() 或 fit() 调用变慢,表现为“越往后越卡”。
- 在每次
model.save()后立即调用tf.keras.backend.clear_session() - 必须在保存完成之后调用,否则模型权重和图结构丢失,后续无法继续训练
-
del model不起作用,它只删 Python 引用,不释放 TensorFlow 后端状态 - 多智能体场景下(如 left_ball/right_ball),每个 agent 保存后都要单独清一次
batch_size 和 prefetch() 配置错位,流水线断层
典型错误是 dataset.prefetch().batch(),这会让 prefetch 提前拉取未 batch 的原始样本,GPU 等着喂数据时反而在等 CPU 拼 batch。正确链路必须是:map() → cache() → shuffle() → batch() → prefetch()。
-
prefetch(1)够用,设更大(如tf.data.AUTOTUNE)未必加速,还可能吃光内存 - 如果模型每步需要多个 batch(如梯度累积),
prefetch()无法替代逻辑,得靠外层循环控制 - 用
dataset.cardinality().numpy()确认数据集大小,避免prefetch()在无限数据集上无意义堆积
GPU 显存没预留,TensorRT 加速静默失败
想用 TensorRT 加速 TensorFlow 推理,但直接对 tf.keras.Model 调用 trt.create_inference_graph 报 NotImplementedError: Cannot convert a symbolic Tensor,本质是没走 frozen graph 路径,也没给 TensorRT 留显存。
- 必须先导出 frozen graph(
.pb),不能直接传 Keras Model - 启动 session 前,用
tf.GPUOptions(per_process_gpu_memory_fraction=0.67)预留至少 30% 显存 - TensorRT 只加速 Conv2D/ReLU/BatchNorm/FullyConnected 类子图,RNN、自定义 op、
tf.while_loop会被跳过,仍走 TF 原生路径
最常被忽略的是:首帧延迟高不是 bug,是 TensorRT 在做 kernel auto-tuning;校准 int8 用错数据分布,精度掉点可能超 5%;trt.calib_mode=True 忘关,会持续统计拖慢首帧。这些细节不处理,再快的硬件也白搭。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











