使用DQN求解Gym Taxi-v3问题遇性能瓶颈的技术问询
DQN求解Gym Taxi-v3问题的性能异常排查与优化
问题概述
用表格型Q-learning求解Taxi-v3问题时,10000次迭代后平均奖励达8.x,成功率100%,效果理想。但改用DQN时,训练约100次迭代后,评估的episode_reward_mean收敛在-210左右,episode_len_mean稳定在200左右,完全未达到预期性能。
评估结果截图

当前使用的DQN训练与评估代码
from ray.rllib.algorithms.ppo import PPOConfig from ray.rllib.algorithms.dqn.dqn import DQN, DQNConfig from ray.rllib.algorithms.a2c import A2CConfig import ray import csv import datetime import os ray.init(local_mode=True) # ray.init(address='auto') # connect to Ray cluster # config = DQNConfig() num_rollout_workers = 62 max_train_iter_times = 20000 config = DQNConfig() config = config.environment("Taxi-v3") config = config.rollouts(num_rollout_workers=num_rollout_workers) config = config.framework("torch") # Update exploration_config exploration_config={ "type": "EpsilonGreedy", "initial_epsilon": 1.0, "final_epsilon": 0.02, "epsilon_timesteps": max_train_iter_times } config = config.exploration(exploration_config=exploration_config) config.evaluation_config = { "evaluation_interval": 10, "evaluation_num_episodes": 10, } # Update replay_buffer_config replay_buffer_config = { "_enable_replay_buffer_api": True, "type": "MultiAgentPrioritizedReplayBuffer", "capacity": 1000, "prioritized_replay_alpha": 0.5, "prioritized_replay_beta": 0.5, "prioritized_replay_eps": 3e-6, } config = config.training( model={"fcnet_hiddens": [50, 50, 50]}, lr=0.001, gamma=0.99, replay_buffer_config=replay_buffer_config, target_network_update_freq=500, double_q=True, dueling=True, num_atoms=1, noisy=False, n_step=3, ) algo = DQN(config=config) # algo = config.build() # 2. build the algorithm, no_improvement_counter = 0 prev_reward = None # Get the current date current_date = datetime.datetime.now().strftime('%Y%m%d') # Open the csv file in write mode with open(f'train_{current_date}.csv', 'w', newline='') as file: writer = csv.writer(file) # Write the header row writer.writerow(["Iteration", "Reward_Mean", "Episode_Length_Mean"]) for i in range(max_train_iter_times): print(f'#{i}: {algo.train()}\n') # 3. train it, # Save the model every 5 iterations if (i + 1) % 10 == 0: checkpoint = algo.save() print("Model checkpoint saved at", checkpoint) eval_result = algo.evaluate() print(f'to evaluate model: {eval_result}') # 4. and evaluate it. cur_reward = eval_result['evaluation']['sampler_results']['episode_reward_mean'] cur_episode_len_mean = eval_result['evaluation']['sampler_results']['episode_len_mean'] # Write the iteration, reward and episode length to csv writer.writerow([i + 1, cur_reward, cur_episode_len_mean]) # Force the file to be written to disk immediately file.flush() os.fsync(file.fileno()) if prev_reward is not None and cur_reward <= prev_reward: no_improvement_counter += 1 else: no_improvement_counter = 0 print(f'evaluated episode_reward_mean: {cur_reward}, no improvement counter: {no_improvement_counter}\n') if no_improvement_counter >= 20: print(f"Training stopped as the episode_reward_mean did not improve for 20 consecutive evaluations. totalIterNum: {i + 1}") break prev_reward = cur_reward
已尝试的调整
- 将回放缓冲区容量从1000调整为10000
- 将
n_step从3调整为20
上述调整均未改善模型性能。
问题分析与优化建议
1. 探索策略参数不合理
当前epsilon_timesteps设置为20000(训练迭代次数),但Ray RLlib中该参数的单位是环境步长而非训练迭代次数。Taxi-v3每个训练迭代会生成大量步长,导致epsilon衰减极慢,后期仍在高概率随机探索,无法利用已学习的策略。
- 优化:将
epsilon_timesteps改为50000(环境步长),或根据实际步长调整,确保epsilon在合理时间内衰减到0.02。 - 额外:评估时需关闭探索,在
evaluation_config中添加"explore": False,避免评估时的随机行为影响结果。
2. 回放缓冲区配置不当
使用MultiAgentPrioritizedReplayBuffer对于单Agent的Taxi-v3问题过于复杂,优先回放的参数(alpha/beta)可能导致训练不稳定;初始缓冲区容量1000过小,不足以存储足够多样的经验。
- 优化:改用基础的
ReplayBuffer,将容量设置为50000;若坚持使用优先回放,需调优prioritized_replay_alpha(建议0.6)和prioritized_replay_beta(从0.4逐步提升到1.0)。
3. 网络结构过度复杂
Taxi-v3的状态空间仅为500个离散状态,3层50神经元的全连接网络过于庞大,容易导致过拟合或训练效率低下。
- 优化:简化网络结构,例如设置
model={"fcnet_hiddens": [64, 64]}或甚至单层64神经元。
4. 目标网络更新频率过高
target_network_update_freq=500表示每500步更新一次目标网络,对于小问题来说更新间隔太长,导致Q值估计偏差无法及时修正。
- 优化:将该值调整为100或200步,加快目标网络的更新速度。
5. 训练参数与资源配置不合理
- 学习率:0.001对于小网络来说偏大,建议降低到0.0001或0.0005,避免训练震荡。
- Rollout Workers数量:62个worker远超Taxi-v3的需求,过多worker会导致数据分布混乱,训练不稳定。建议降低到2-4个,或设置为0(仅用主进程)。
- 终止条件:当前的
no_improvement_counter逻辑过于严格,可能在模型尚未收敛时提前终止训练。建议放宽到50次无提升,或改为监控奖励是否达到阈值(如5以上)再终止。
6. 简化DQN变体配置
当前启用了double_q和dueling,虽然这些变体通常能提升性能,但在小问题中可能增加训练复杂度,导致收敛变慢。
- 优化:先禁用
double_q和dueling,确保基础DQN能正常收敛后,再逐步添加这些变体。
内容的提问来源于stack exchange,提问作者Aaron
相关产品推荐
相关产品推荐

