悬崖行走环境中DQN网络持续选右移动作致反复失败的问题
解决DQN在悬崖行走环境中始终选择右移动作的问题
你遇到的核心问题是网络无法正确学习到悬崖动作的惩罚价值,导致反复选择右移动作坠落。以下是代码中的关键错误及修正方案:
1. 输出层使用ReLU激活函数导致Q值失真
DQN的输出是动作的Q值,Q值可以为负(对应惩罚性动作),但你在输出层用了F.relu(self.fc3(state)),会把所有负Q值截断为0。这意味着网络无法区分"掉悬崖(-100)"和"普通步骤(-1)"的差异,自然会一直选看起来"无惩罚"的右移动作。
修正:移除输出层的ReLU,直接返回线性输出:
def forward(self, state): state = torch.Tensor(state) state = F.relu(self.fc1(state)) state = F.relu(self.fc2(state)) actions = self.fc3(state).squeeze() # 去掉ReLU return actions
2. 目标Q值计算逻辑错误
在计算Q_target时,你用了torch.max(q_next),这会取整个batch中所有样本的最大Q值,而不是每个样本自身的最大下一个状态Q值。这会导致所有样本的目标值被错误地统一为同一个最大值,网络无法学习到每个状态的正确价值。
同时还要注意:终止状态(到达目标或掉悬崖)没有后续状态,目标值应该直接等于当前奖励,不需要叠加后续Q值。
修正:
# 替换原Q_target计算部分 q_next_max = torch.max(q_next, dim=1)[0] # 每个样本的最大下一个状态Q值 rewards = torch.Tensor(list(mini_batch[:,2])) is_terminal = torch.Tensor(list(mini_batch[:,4])) Q_target = q_pred.clone() indices = np.arange(batch_size) actions_batch = mini_batch[:,1].astype(int) # 终止状态目标值=当前奖励,非终止状态=奖励+gamma*max_next_q Q_target[indices, actions_batch] = rewards + gamma * q_next_max * (1 - is_terminal)
3. 网络输入批量处理错误
你在训练时给网络传入列表形式的单个one-hot向量,会导致网络无法正确处理批量输入,输出形状混乱。
修正:将批量状态编码为二维张量:
# 批量编码状态 states_batch = torch.Tensor(np.array([encode_vector(int(s), state_size) for s in mini_batch[:,0]]).squeeze(1)) q_pred = q_net(states_batch) next_states_batch = torch.Tensor(np.array([encode_vector(int(s), state_size) for s in mini_batch[:,3]]).squeeze(1)) q_next = target_net(next_states_batch)
4. 目标网络更新时机优化
原代码在每个回合开始时更新目标网络,会导致第0回合就更新(此时在线网络还没开始学习)。建议改为每N个回合结束后更新:
# 移到回合循环的末尾更新 if (e + 1) % steps_for_tgt_update == 0: target_net.update(q_net) print("Target updated")
其他优化建议
- 学习率调整:0.01的学习率对于DQN过高,建议改为0.001,避免训练震荡。
- 探索率衰减:eps_decay=0.999,10000回合后eps仍高于0.3,可改为0.995加快衰减,让智能体更快进入利用阶段。
- 状态编码简化:悬崖行走的离散状态可以用嵌入层替代one-hot编码,效率更高:
class NeuralNet(nn.Module): def __init__(self, state_size, action_size): super(NeuralNet, self).__init__() self.embedding = nn.Embedding(state_size, 16) # 把状态索引转为16维向量 self.fc1 = nn.Linear(16, 64) self.fc2 = nn.Linear(64, 20) self.fc3 = nn.Linear(20, action_size) def forward(self, state): state = self.embedding(state.long()) state = F.relu(self.fc1(state)) state = F.relu(self.fc2(state)) actions = self.fc3(state).squeeze() return actions
使用时直接传入状态索引张量即可,无需one-hot编码:
# 选择动作时 action = torch.argmax(q_net(torch.tensor([state]))).item() # 训练时 states_batch = torch.tensor(mini_batch[:,0].astype(int)) q_pred = q_net(states_batch)
内容的提问来源于stack exchange,提问作者olejonas
相关产品推荐
相关产品推荐

