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

训练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转换。


解决方案
  1. 抑制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
      )
      
  2. 手动释放GradientTape资源

    • 在计算完梯度后,手动将tape变量置为None,触发垃圾回收,避免旧实例占用内存:
      def get_gradients(...):
          with tf.GradientTape() as tape:
              # 计算逻辑
          gradients = tape.gradient(...)
          tape = None  # 释放资源
          # 后续优化器操作
      
  3. 替换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)
      
  4. 指定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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 04:39:35