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

tf.keras.optimizers.Adam.apply_gradients触发tf.function重追踪的内存泄漏问题

问题:TensorFlow重追踪警告与内存泄漏修复

我遇到内存泄漏问题,怀疑和下面的TensorFlow警告有关:

WARNING:tensorflow:最近6次调用<function _BaseOptimizer._update_step_xla at 0x7fa9f8074c20>均触发了tf.function重追踪。重追踪开销较大,频繁重追踪可能源于:(1) 在循环中重复创建@tf.function;(2) 传入形状不同的张量;(3) 传入Python对象而非张量。针对(1),请在循环外定义@tf.function;针对(2),@tf.function的reduce_retracing=True选项可避免不必要的重追踪;针对(3),请参考TensorFlow官方文档获取更多细节。

警告出现在learn函数中,代码如下:

def learn(self):
    for _ in range(self.n_epochs):
        state_arr, additional_info, action_arr, old_prob_arr, values, reward_arr, _, trades_complete, env_states, batches = self.memory.generate_batches() # generate batches
        reward_diff = reward_arr[:-1] + values[1:] * (1 - tf.cast(trades_complete[:-1], dtype=tf.float32)) - values[:-1]
        advantage = tf.concat([tf.cumsum(reward_diff, reverse=True), self.zero_tensor], axis=0)
        with tf.GradientTape(persistent=True) as tape:
            new_probs, new_val = self.cnn_actor_critic([state_arr, additional_info])
            masked_new_probs = ENVIRONMENT.mass_apply_mask(new_probs.numpy(), env_states)
            rows = tf.range(tf.shape(masked_new_probs)[0])
            index_arr = tf.add(tf.cast(action_arr, dtype=tf.int32), self.one_val)
            gather_indices = tf.stack([rows, index_arr], axis=1)
            chosen_probs = tf.gather_nd(masked_new_probs, gather_indices)
            new_log_probs_of_old_actions = tf.negative(tf.math.log(chosen_probs))
            critic_value = tf.squeeze(new_val, 1) # removes dimensions of size 1 from the tensor
            values_list = tf.convert_to_tensor(values)
            returns = tf.add(advantage, values_list)
            critic_loss = tf.keras.losses.MSE(critic_value, returns)
            prob_ratio = tf.math.exp(tf.add(new_log_probs_of_old_actions,tf.cast(tf.negative(old_prob_arr), dtype=tf.float32)))
            weighted_probs = tf.multiply(advantage, prob_ratio)
            clipped_probs = tf.clip_by_value(prob_ratio, 0.8, 1.2)
            weighted_clipped_probs = tf.multiply(clipped_probs, advantage)
            l_clip = tf.math.minimum(weighted_probs, weighted_clipped_probs) # prviously actor_loss
            entropy_term = tf.negative(tf.reduce_sum(tf.multiply(new_log_probs_of_old_actions, tf.math.log(new_log_probs_of_old_actions))))
            l_q_extension = tf.add(tf.multiply(self.c1,critic_loss), tf.negative(tf.multiply(self.c2,entropy_term)))
            l_q = tf.negative(tf.add(l_clip,l_q_extension))
            actor_critic_cnn_loss = tf.math.reduce_mean(l_q)
        cnn_actor_critic_params = self.cnn_actor_critic.trainable_variables
        actor_critic_grads = tape.gradient([actor_critic_cnn_loss, critic_loss], cnn_actor_critic_params, unconnected_gradients=tf.UnconnectedGradients.ZERO)
        self.cnn_actor_critic.optimizer.apply_gradients(zip(actor_critic_grads, cnn_actor_critic_params))
    self.memory.clear_memory()

注释掉最后一行self.cnn_actor_critic.optimizer.apply_gradients(zip(actor_critic_grads, cnn_actor_critic_params))后警告消失,求修改方案避免警告和内存泄漏。


解决方案

1. 把优化器更新逻辑封装到循环外的@tf.function中

问题核心是每次循环调用apply_gradients时,TensorFlow会隐式创建并追踪新的tf.function,导致重复重追踪。将梯度更新逻辑抽离,在类初始化阶段或循环外定义带reduce_retracing=True的tf.function:

# 在类内部(循环外)定义,比如__init__方法之后
@tf.function(reduce_retracing=True)
def update_model(self, grads, params):
    self.cnn_actor_critic.optimizer.apply_gradients(zip(grads, params))

然后在learn函数的循环中替换原代码行:

self.update_model(actor_critic_grads, cnn_actor_critic_params)

2. 修复张量形状不一致问题

检查generate_batches返回的所有张量(如state_arr、additional_info等)是否每次形状固定。如果存在动态形状变化,可在生成批次后用tf.ensure_shape显式声明形状,帮助TensorFlow减少重追踪:

# 示例:根据你的实际张量形状调整参数
state_arr = tf.ensure_shape(state_arr, [None, 84, 84, 4])
additional_info = tf.ensure_shape(additional_info, [None, 10])

3. 避免在计算图中混合numpy操作

代码中masked_new_probs = ENVIRONMENT.mass_apply_mask(new_probs.numpy(), env_states)将张量转成numpy数组再处理,打破了TensorFlow计算图的连续性,容易导致动态形状或Python对象传入tf.function。建议把mass_apply_mask改写成纯TensorFlow操作,比如用tf.where、tf.tensor_scatter_nd_update实现掩码逻辑:

# 示例:将掩码逻辑替换为TensorFlow版本
def mass_apply_mask(probs, env_states):
    # 根据实际掩码逻辑调整,这里假设env_states是布尔型掩码
    mask = tf.cast(env_states, tf.float32)
    return probs * mask

修改后直接传入张量处理,无需转numpy:

masked_new_probs = ENVIRONMENT.mass_apply_mask(new_probs, env_states)

4. 清理GradientTape的persistent属性

代码中使用了persistent=True的GradientTape,但仅调用了一次tape.gradient,该属性会保留tape资源导致内存占用。去掉persistent=True可减少内存泄漏风险:

with tf.GradientTape() as tape: # 移除persistent=True
    # ... 原计算逻辑 ...

内容的提问来源于stack exchange,提问作者Bryan Carty

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 20:05:14