训练DQN Agent时速度变慢并在约50轮后崩溃的问题排查
训练DQN Agent时出现性能骤降并崩溃的问题
训练DQN Agent时,大约在50个episodes后,replay函数中的拟合操作开始变慢,进而导致电脑卡顿,最终PyCharm直接崩溃。训练轮次耗时从10-13ms逐渐增至数秒,直至完全冻结。
智能体核心代码(输入为含12个不同浮点数的数组)
def __init__(self, agent_config): self.agent_config = agent_config self.state_size = agent_config['state_size'] self.discount_factor = agent_config['discount_factor'] self.action_size = agent_config['action_size'] self.memory = deque(maxlen=agent_config['memory_size']) self.epsilon = agent_config['epsilon'] self.learning_rate = agent_config['learning_rate'] self.model = self._build_model() self.target_model = self._build_model() def _build_model(self): model = Sequential() model.add(Dense(64, input_dim=self.state_size, activation='relu')) model.add(Dense(64, activation='relu')) model.add(Dense(self.action_size, activation='linear')) model.compile(loss=keras.losses.Huber(), optimizer=Adam(learning_rate=self.learning_rate)) return model def act(self, state): if np.random.random() > self.epsilon: actions = self.model.predict(state) return np.argmax(actions[0]) else: return np.random.randint(0, self.action_size) def remember(self, state, action, reward, next_state, done): self.memory.append((state, action, reward, next_state, done)) def replay(self, batch_size): minibatch = random.sample(self.memory, batch_size) X = [] y = [] for index, (state, action, reward, next_state, done) in enumerate(minibatch): current_q = self.model.predict(state) future_q = self.target_model.predict(next_state) if not done: max_future_q = np.max(future_q) new_q = reward + self.discount_factor + max_future_q else: new_q = reward current_q[0][action] = new_q X.append(state[0]) y.append(current_q) self.model.train_on_batch(np.array(X), np.array(y))
训练循环代码
def train_dqn(agent, env: gym.Env, episodes: int, batch_size: int, episode_length: int): ep_rewards = [] aggr_ep_rewards = {'ep': [], 'avg': [], 'max': [], 'min': []} STATS_EVERY = 5 START_EPSILON_DECAYING = 1 END_EPSILON_DECAYING = episodes // 2 epsilon_decay_value = agent.epsilon / (END_EPSILON_DECAYING - START_EPSILON_DECAYING) for episode in range(episodes): total_reward = 0 # reset state in the beginning of each game state = env.reset() state = encode_state(state) # Loop over time steps until the episode is done or the time limit is reached for i in range(episode_length): action = agent.act(state) decoded_action = decode_action(action, state, env.action_space) next_state, reward, done, _, _ = env.step(decoded_action) encoded_next_state = encode_state(next_state) agent.remember(state, action, reward, encoded_next_state, done) state = encoded_next_state total_reward += reward if (episode * 30) % 90 == 0: if len(agent.memory) > batch_size: agent.replay(batch_size) if (episode * 30) % 450 == 0: agent.target_model.set_weights(agent.model.get_weights()) ep_rewards.append(total_reward) if not episode % STATS_EVERY: average_reward = sum(ep_rewards[-STATS_EVERY:]) / STATS_EVERY aggr_ep_rewards['ep'].append(episode) aggr_ep_rewards['avg'].append(average_reward) aggr_ep_rewards['max'].append(max(ep_rewards[-STATS_EVERY:])) aggr_ep_rewards['min'].append(min(ep_rewards[-STATS_EVERY:])) print( f'Episode: {episode:>5d}, average reward: {average_reward:>4.1f}, current epsilon: {agent.epsilon:>1.2f}') if END_EPSILON_DECAYING >= episode >= START_EPSILON_DECAYING: if agent.epsilon - epsilon_decay_value >= 0: agent.epsilon -= epsilon_decay_value print("episode: {}/{}, score: {}, e: {:.2}" .format(episode, episodes, total_reward, agent.epsilon)) if episode % 10 == 0: agent.save(f"models/model_{episode}.h5") plt.plot(aggr_ep_rewards['ep'], aggr_ep_rewards['avg'], label="average rewards") plt.plot(aggr_ep_rewards['ep'], aggr_ep_rewards['max'], label="max rewards") plt.plot(aggr_ep_rewards['ep'], aggr_ep_rewards['min'], label="min rewards") plt.legend(loc=4) plt.show() agent.save(f"model_{episodes}.h5") return agent
已尝试的调整措施
- 调整不同的batch size、memory容量以及replay函数的执行频率,仅能延迟崩溃时间
- 调整神经网络规模,无效果
- 新增16GB内存后,崩溃点延后至约245个episodes
- 使用
model.fit时在210轮崩溃,train_on_batch时在245轮崩溃 - 安装CUDA改用GPU训练,无改善
请问该问题由代码缺陷还是电脑/操作系统问题导致?
内容的提问来源于stack exchange,提问作者willem12
相关产品推荐
相关产品推荐

