You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

tf.function结合自定义训练函数引发内存泄漏的解决方法咨询

TensorFlow 2.x 自定义FRAE模型训练内存泄漏解决方案

问题背景

基于tf.keras.Model实现的FRAE模型可正常运行,但训练阶段内存持续增长,预测阶段无此问题。排查确认是@tf.function图模式下,非训练变量self.buffer的更新操作导致内存泄漏,且无法移除@tf.function以保留训练加速能力。

核心原因

  1. 训练时GradientTape默认追踪所有参与计算的张量,包括self.buffer更新过程中生成的临时拼接张量,这些张量在图模式下未被正确回收,导致内存累积。
  2. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.07 20:05:01