tf.keras.optimizers.Adam.apply_gradients触发tf.function重追踪的内存泄漏问题
我遇到内存泄漏问题,怀疑和下面的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

