训练Atari Breakout强化学习代理时RAM占用持续增长求助
问题描述
最近训练Atari Breakout强化学习代理,运行约1.5小时后电脑卡顿,鼠标交互困难。监控发现RAM占用随运行时长持续增长。一开始怀疑是replay buffer的问题,优化后无效;即使停止向replay buffer添加数据(已达50000条),RAM仍继续增长,最终定位到get_gradients函数,同时伴随tf.function重追踪警告。
定位到的代码片段
get_gradients函数
def get_gradients(self, target_q_values, importance, states, actions): with tf.GradientTape() as tape: q_values_current_state_dqn = self.dqn_architecture(states) one_hot_actions = tf.keras.utils.to_categorical(actions, self.num_legal_actions, dtype=np.float32) # e.g. [[0,0,1,0],[1,0,0,0],...] Q = tf.reduce_sum(tf.multiply(q_values_current_state_dqn, one_hot_actions), axis=1) error = Q - tf.cast(target_q_values, tf.float32) loss = tf.keras.losses.Huber()(target_q_values, Q) if self.use_prioritized_experience_replay: loss = tf.reduce_mean(loss * importance) # Gradient is scaled -> loss = lower at begining -> reduces bias against situataions that are sampled more frequently dqn_architecture_gradients = tape.gradient(loss, self.dqn_architecture.trainable_variables) # Computes the gradient using operations recorded in context of this tape. self.dqn_architecture.optimizer.apply_gradients(zip(dqn_architecture_gradients, self.dqn_architecture.trainable_variables)) return loss, error
tf.function重追踪警告日志
2023-02-16 22:48:32,045 5 out of the last 5 calls to <function Agent.get_gradients at 0x7fb3ec66e830> triggered tf.function retracing. Tracing is expensive and the excessive number of tracings could be due to (1) creating @tf.function repeatedly in a loop, (2) passing tensors with different shapes, (3) passing Python objects instead of tensors. For (1), please define your @tf.function outside of the loop. For (2), @tf.function has reduce_retracing=True option that can avoid unnecessary retracing. For (3), please refer to https://www.tensorflow.org/guide/function#controlling_retracing and https://www.tensorflow.org/api_docs/python/tf/function for more details. 2023-02-16 22:48:32,217 6 out of the last 6 calls to <function Agent.get_gradients at 0x7fb3ec66e830> triggered tf.function retracing. Tracing is expensive and the excessive number of tracings could be due to (1) creating @tf.function repeatedly in a loop, (2) passing tensors with different shapes, (3) passing Python objects instead of tensors. For (1), please define your @tf.function outside of the loop. For (2), @tf.function has reduce_retracing=True option that can avoid unnecessary retracing. For (3), please refer to https://www.tensorflow.org/guide/function#controlling_retracing and https://www.tensorflow.org/api_docs/python/tf/function for more details.
函数调用关系
train_network函数(调用get_gradients)
def train_network(self, batch_size, gamma, frame_number, priority_scale): importance = 0 if self.use_prioritized_experience_replay: (states, actions, rewards, new_states, terminal_flags), importance, indices = self.replay_buffer.sample_buffer(self.batch_size, priority_scale) importance = importance ** (1-self.calculate_epsilon(frame_number)) # recently started training = low frame number = high epsilon = low power = largely decreased importance = lower importance, later in training is slightly decreased importance. Increases importance of newer frames else: states, actions, rewards, new_states, terminal_flags = self.replay_buffer.sample_buffer(self.batch_size, priority_scale) best_action_in_next_state_dqn = self.dqn_architecture.predict(new_states, verbose=0).argmax(axis=1) target_q_network_q_values = self.target_dqn_architecture.predict(new_states, verbose=0) optimal_q_value_in_next_state_target_dqn = target_q_network_q_values[range(batch_size), best_action_in_next_state_dqn] target_q_values = rewards + (gamma*optimal_q_value_in_next_state_target_dqn * (1-terminal_flags)) # makes 0 if terminal flag set # Calculate loss and perform gradfient descent # TensorFlow "records" relevant operations executed inside the context of a tf. GradientTape onto a "tape". TensorFlow then uses that tape to compute the gradients of a "recorded" computation using reverse mode differentiation. loss, error = self.get_gradients(target_q_values, importance, states, actions) if self.use_prioritized_experience_replay: self.replay_buffer.set_priorities(indices, error) return float(loss.numpy()), error
循环调用train_network的代码
while frame_number < NUM_FRAMES_AGENT_TRAINED_OVER: breakout_environment.reset_env() episode_reward_sum = 0 for _ in range(MAX_EPISODE_LENGTH): # Get action action = breakout_agent.take_action(frame_number, breakout_environment.state) # Take step frame, reward, terminal, life_lost = breakout_environment.step(action) frame_number += 1 episode_reward_sum += reward # Add experience to replay memory action, frame, reward, terminal, clip_reward breakout_agent.add_experience_to_replay_buffer(action, frame[:, :, 0], reward, life_lost, CLIP_REWARD) # Train the network every 4 additions to the replay buffer if frame_number % UPDATE_FREQUENCY == 0 and breakout_agent.replay_buffer.total_indexes_written_to > REPLAY_BUFFER_START_SIZE: loss, _ = breakout_agent.train_network(BATCH_SIZE, DISCOUNT_FACTOR, frame_number, PRIORITY_SCALE) # batch_size, gamma, frame_number, priority_scale loss_list.append(loss) # Update target network if frame_number % TARGET_UPDATE_FREQ == 0 and frame_number > REPLAY_BUFFER_START_SIZE: breakout_agent.update_target_network() # Break the loop when the game is over if terminal: break rewards_list.append(episode_reward_sum)
补充调研
发现Python标量或列表传入tf.function会不断生成新图,应尽量传Tensor类型,但optimizer.apply_gradients的参数维度和嵌套深度不固定,无法用tf.convert_to_tensor或tf.ragged.constant转换。
解决方案
抑制tf.function重复追踪
- 给
get_gradients添加@tf.function(reduce_retracing=True)装饰器,开启自动减少重追踪机制;同时在train_network中将所有传入get_gradients的非Tensor参数转为Tensor:# 在train_network中调用get_gradients前添加 target_q_values = tf.convert_to_tensor(target_q_values, dtype=tf.float32) if self.use_prioritized_experience_replay: importance = tf.convert_to_tensor(importance, dtype=tf.float32) actions = tf.convert_to_tensor(actions, dtype=tf.int32) - 用
tf.cond替代Python原生if判断处理use_prioritized_experience_replay,让tf.function能稳定追踪计算图:# 替换get_gradients中的if判断 loss = tf.cond( tf.constant(self.use_prioritized_experience_replay), lambda: tf.reduce_mean(loss * importance), lambda: loss )
- 给
手动释放GradientTape资源
- 在计算完梯度后,手动将
tape变量置为None,触发垃圾回收,避免旧实例占用内存:def get_gradients(...): with tf.GradientTape() as tape: # 计算逻辑 gradients = tape.gradient(...) tape = None # 释放资源 # 后续优化器操作
- 在计算完梯度后,手动将
替换predict调用为模型直接调用
train_network中用模型的__call__方法(model(inputs, training=False))替代predict,复用tf.function的计算图,避免重复追踪:best_action_in_next_state_dqn = tf.argmax(self.dqn_architecture(new_states, training=False), axis=1) target_q_network_q_values = self.target_dqn_architecture(new_states, training=False)
指定tf.function输入签名
- 给
@tf.function添加input_signature,明确参数的Tensor形状和类型,减少无意义的重追踪(需根据实际state形状调整):@tf.function(input_signature=[ tf.TensorSpec(shape=(None,), dtype=tf.float32), tf.TensorSpec(shape=(None,), dtype=tf.float32), tf.TensorSpec(shape=(None, 84, 84), dtype=tf.float32), tf.TensorSpec(shape=(None,), dtype=tf.int32) ], reduce_retracing=True) def get_gradients(self, target_q_values, importance, states, actions): # 函数逻辑
- 给
内容的提问来源于stack exchange,提问作者Bryan Carty
相关产品推荐
相关产品推荐

