Deep Q Network训练贪吃蛇时Q值固定且无性能提升问题排查
你的问题里藏着几个关键的错误点,正是这些问题导致智能体无法学习,甚至反复自杀。我来逐一拆解并给出修复方案:
1. 网络输出层完全用错了激活函数
DQN的输出是动作的Q值(预期累积奖励),属于连续值范畴,但你最后一层用了softmax——这个激活函数会把输出归一化成概率分布,完全不符合DQN的设计逻辑。而且你输出层是Dense(5),但你的动作空间只有4个(上/右/下/左),维度也不匹配!
修复方案:
# 替换最后两层 self.policy_model.add(Dense(16, activation="relu")) self.policy_model.add(Dense(4, activation=None)) # 用线性激活输出Q值
2. Bellman方程实现错误,完全破坏了Q值更新逻辑
你现在的代码把所有动作的Q值都直接替换成了max_q,这意味着不管智能体选了什么动作,所有动作的价值都被强制设为同一个值,网络根本学不到不同动作的差异。正确的做法是:只更新当前选中动作对应的Q值,其他动作的Q值保持原预测结果。
另外,你还忽略了done状态的处理——当智能体死亡时,没有下一个状态,Bellman方程应该直接等于当前奖励,不能加折扣后的下一个Q值。
修复后的Bellman方程代码:
for memory in memory_sample: memstate = memory.state action = memory.action next_state = memory.next_state reward = memory.reward done = memory.done # 注意:你需要在Experience类里加入done字段 # 1. 预测当前状态的Q值 current_q = self.policy_model.predict(memstate)[0] # 2. 初始化目标Q值为当前Q值的副本 target_q = current_q.copy() if done: # 死亡时,直接用当前奖励作为目标Q值 target_q[action] = reward else: # 非死亡时,计算下一个状态的最大Q值 next_q = self.replay_model.predict(next_state)[0] target_q[action] = reward + self.discount_rate * np.max(next_q) X.append(memstate) Y.append(target_q)
3. 死亡的关键经验没有存入回放缓冲区
当experience['done'] == True时,你直接break,没有把这个"死亡-惩罚"的关键经验存入memory。这导致智能体永远学不到"死亡会带来负奖励"这个核心规则,自然会反复自杀。
修复方案:
if experience['done'] == True: episode_reward += experience['reward'] # 把死亡经验存入回放缓冲区 self.push_memory(Experience(experience['state'], experience['action'], experience['reward'], experience['next_state'], experience['done'])) break
4. 状态输入过于冗余,训练效率极低
你用600x600的截图作为状态,这个尺寸太大了!卷积层需要学习大量无关信息(比如背景像素、蛇的皮肤细节),不仅训练慢,还容易让网络陷入局部最优。
推荐两种优化方向:
- 优先用结构化状态:直接用蛇头坐标、食物坐标、蛇身体节点坐标、蛇当前移动方向作为输入,状态维度从几十万降到几十维,网络能快速学习核心规律。
- 如果坚持用视觉输入:把截图resize到84x84的小尺寸,转成灰度图(单通道),减少计算量。
5. 超参数设置不合理,训练稳定性差
- batch_size=2:太小了,小batch会导致梯度噪声极大,训练波动剧烈。建议改成32或64,这是DQN的标准配置。
- epsilon衰减逻辑无效:你的
eps_decay=1e-5,训练1000轮后epsilon仅降到0.99,智能体几乎全程在随机探索,根本没机会利用学到的策略。换成指数衰减:# 每episode衰减,eps_decay可以设为200 epsilon = self.eps_end + (self.eps_start - self.eps_end) * np.exp(-1. * episode / self.eps_decay) - target_update=100:这个间隔有点大,可以改成每50轮更新一次目标网络,让目标Q值更稳定。
关于帧堆叠的问题
如果用结构化状态(包含蛇的移动方向),贪吃蛇是完全可观测的,不需要帧堆叠。如果用视觉输入,帧堆叠可以帮助网络捕捉运动趋势,但优先解决前面的核心问题比这个更重要。
把这些修改都落实后,重新训练应该就能看到智能体开始学习避开死亡、主动吃食物了。
内容的提问来源于stack exchange,提问作者achandra03

