tf.function动态输入导致图缓存无限增长并非内存泄漏,而是因新shape/dtype触发新图生成且不自动清理;模型全局持有、Dataset迭代器未释放、map闭包捕获大对象等会叠加加剧显存持续上涨。

tf.function 动态输入导致图缓存无限增长
不是内存泄漏,是图缓存没被清理——每次 tf.function 接收到新 shape 或 dtype 的输入(比如 batch size 变化、文本长度不固定),就会生成并缓存一个全新计算图。这些图不会自动释放,显存持续上涨,tf.keras.backend.clear_session() 也无效。
常见现象:nvidia-smi 显存线性上升;len(tf.get_default_graph().get_operations()) 随请求次数增加;tf.autograph.trace 日志里反复出现 new trace。
- 强制统一签名:
@tf.function(input_signature=[tf.TensorSpec(shape=[None, 224, 224, 3], dtype=tf.float32)]) - 预处理动态输入:对图像
resize、对文本pad_sequences或clip_by_value - 禁用自动追踪(仅限无控制流函数):
@tf.function(autograph=False)
模型对象长期持有 graph 引用无法 GC
tf.keras.models.load_model() 返回的对象内部绑定了 ConcreteFunction 和底层 Graph。一旦赋值给模块级变量、self.model 或闭包,整个图结构就卡在 C++ 层,Python 的 del model 只删引用,不触发放。
典型错误:把 model = tf.keras.models.load_model("path") 写在文件顶层,或在 Flask/FastAPI 的全局 scope 初始化模型。
立即学习“Python免费学习笔记(深入)”;
- 长周期服务中改用函数封装,每次推理前显式卸载:
with tf.device("/CPU:0"): model = ...,再del model - 紧接调用
tf.keras.backend.clear_session()(注意:必须在del之后) - 避免用
global或类属性长期持有模型实例
tf.data.Dataset 迭代器未耗尽残留 CUDA 资源
tf.data.Dataset.from_generator() 或复杂 pipeline 构建的迭代器,其 Iterator 对象会强引用图节点和 GPU 缓冲区。HTTP 请求中断、异常退出、或只取 dataset.take(n) 就丢弃,底层资源就卡住不释放。
验证方式:tf.data.experimental.cardinality(dataset).numpy() 返回 -2(unknown),且显存不降。
- 显式管理迭代器生命周期:
it = iter(dataset); try: next(it) finally: del it - 优先用
dataset.batch().map().prefetch(),避免from_generator - 检查是否误用了
repeat()放太前,导致 dataset 变成 infinite,prefetch永远填不满
Dataset.map() 闭包捕获大对象引发隐式驻留
map() 函数定义在外部作用域里,如果它闭包捕获了大数组、模型、h5py.File 或数据库连接,这些对象会被 tf.data 内部缓存机制长期持有,GC 无法回收。
例如:scale_factor = np.random.rand(1000, 1000) 在外层定义,再传进 def preprocess(x): return x * scale_factor —— 整个大数组就被锁死。
- 改用
functools.partial(preprocess, scale_factor=1.0)或显式参数传入标量 - 禁止在
map()内打开文件、创建临时大数组;若必须,末尾加del arr - 用
tf.py_function时务必声明Tout,如tf.py_function(func, [x], Tout=[tf.float32]),否则可能保留 Python 对象引用链
真正难处理的是叠加:比如 map() 闭包捕获模型 + @tf.function 动态追踪 + 全局 list append 中间 tf.Tensor —— 单次调用看不出问题,高并发长周期服务里才逐步崩盘。这类混合驻留,clear_session() 基本失效,得逐层切开查。


















