tf.function结合自定义训练函数引发内存泄漏的解决方法咨询
TensorFlow 2.x 自定义FRAE模型训练内存泄漏解决方案
问题背景
基于tf.keras.Model实现的FRAE模型可正常运行,但训练阶段内存持续增长,预测阶段无此问题。排查确认是@tf.function图模式下,非训练变量self.buffer的更新操作导致内存泄漏,且无法移除@tf.function以保留训练加速能力。
核心原因
- 训练时
GradientTape默认追踪所有参与计算的张量,包括self.buffer更新过程中生成的临时拼接张量,这些张量在图模式下未被正确回收,导致内存累积。 update_buffer中使用tf.concat生成新张量再赋值的操作,会产生额外的未被释放的中间张量。
解决方案
以下修改可在保留@tf.function加速的前提下解决内存泄漏:
1. 优化update_buffer的内存操作逻辑
将拼接后整体赋值改为切片原地更新,避免生成临时拼接张量:
@tf.function(experimental_compile=True) def update_buffer(self, new_element): n = self.shape[0] # 先将buffer内容向后移动n位,再把新元素写入前n位 self.buffer[:, n:].assign(self.buffer[:, :-n]) self.buffer[:, :n].assign(new_element)
2. 限制GradientTape的追踪范围
仅让梯度磁带追踪可训练变量,排除非训练的self.buffer:
@tf.function(experimental_compile=True) def train_step(self, data): x, y = data # 关闭自动追踪,手动指定需要监控的可训练变量 trainable_vars = self.trainable_variables with tf.GradientTape(watch_accessed_variables=False) as tape: tape.watch(trainable_vars) y_pred = self(x, training=True) loss = self.compute_loss(y=y, y_pred=y_pred) gradients = tape.gradient(loss, trainable_vars) self.optimizer.apply_gradients(zip(gradients, trainable_vars)) # 更新指标 for metric in self.metrics: if metric.name == "loss": metric.update_state(loss) else: metric.update_state(y, y_pred) return {m.name: m.result() for m in self.metrics}
3. 优化call函数中的TensorArray使用
添加自动清理配置并显式关闭,帮助内存回收:
@tf.function(experimental_compile=True) def call(self, x): x = tf.squeeze(x, axis=0) seq_len = tf.shape(x)[0] # 开启读取后自动清理,减少内存占用 decoded = tf.TensorArray(tf.float32, size=seq_len, clear_after_read=True) for i in tf.range(seq_len): xexpand = tf.expand_dims(x[i], axis=0) xin = tf.concat((xexpand, self.buffer), axis=1) encoded = self.ls(self.l2(self.l1(xin))) decin = tf.concat([encoded, self.buffer], axis=1) y = self.l5(self.l4(self.l3(decin))) decoded = decoded.write(i, y) self.update_buffer(y) tmp = tf.transpose(decoded.stack(), [1, 0, 2]) decoded.close() # 显式关闭释放资源 return tmp
4. 简化resetBuffer实现
使用tf.zeros_like避免重复定义形状,优化赋值效率:
@tf.function(experimental_compile=True) def resetBuffer(self): self.buffer.assign(tf.zeros_like(self.buffer))
验证效果
修改后重新启动训练:
- 训练阶段内存不再持续增长,保持稳定
- 模型训练速度与修改前一致(保留
@tf.function编译加速) - 模型输出结果与原逻辑完全一致
内容的提问来源于stack exchange,提问作者user20780
相关产品推荐
相关产品推荐

