PyTorch实现DQN训练CartPole-v1失效问题求助
DQN训练CartPole-v1失效问题排查
问题描述
我之前训练DQN能在约65000次迭代后解决CartPole-v1环境,但现在完全失效,达不到原有效果。按过往经验调优超参数后还是没用,训练时损失会爆炸,延长训练时长、调整硬更新tau也解决不了问题。相关代码如下:
完整代码
主训练脚本
import gym import numpy as np import torch from torch import nn from torch.nn import functional as F from torch import optim from models import DQN from memory import Memory from utils import wrap_input, epsilon_greedy def main() -> int: env = gym.make("CartPole-v1") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # Online and offline model for learning model = DQN(env.observation_space, env.action_space, 24).to(device) target = DQN(env.observation_space, env.action_space, 24).to(device) target.eval() # Optimizer and loss function optimizer = optim.Adam(model.parameters(), lr=.001) loss_fn = F.smooth_l1_loss memory = Memory(10_000) obs, info = env.reset() for it in range(65_000): # Do this for the batch norm model.eval() # Maybe explore if np.random.random() <= epsilon_greedy(1.0, .01, 15_000, it): state = wrap_input(obs, device).unsqueeze(0) action = model(state).argmax().item() else: action = env.action_space.sample() # Act in environment and store the memory next_state, reward, done, truncated, info = env.step(action) if truncated or done: next_state = np.zeros(env.observation_space.shape) memory.store([obs, action, reward, int(done), next_state]) done = done or truncated if done: obs, info = env.reset() # Train if len(memory) > 32: model.train() states, actions, rewards, dones, next_states = memory.sample(32) # Wrap and move all values to the cpu states = wrap_input(states, device) actions = wrap_input(actions, device, torch.int64, reshape=True) next_states = wrap_input(next_states, device) rewards = wrap_input(rewards, device, reshape=True) dones = wrap_input(dones, device, reshape=True) # Get current q-values qs = model(states) qs = torch.gather(qs, dim=1, index=actions) # Compute target q-values with torch.no_grad(): next_qs, _ = target(next_states).max(dim=1) next_qs = next_qs.reshape(-1, 1) target_qs = rewards + .9 * (1 - dones) * next_qs.reshape(-1, 1) # Compute loss loss = loss_fn(qs, target_qs) optimizer.zero_grad() loss.backward() # Clip gradients nn.utils.clip_grad_norm_(model.parameters(), 1) # Backprop optimizer.step() # soft update with torch.no_grad(): for target_param, local_param in zip(target.parameters(), model.parameters()): target_param.data.copy_(1e-2 * local_param.data + (1 - 1e-2) * target_param.data) if it % 200 == 0: target.load_state_dict(model.state_dict())
models.py
class FlatExtractor(nn.Module): '''Does nothing but pass the input on''' def __init__(self, obs_space): super(FlatExtractor, self).__init__() self.n_flatten = obs_space.shape[0] def forward(self, obs): return obs class DQN(nn.Module): def __init__(self, obs_space, act_space, layer_size): super(DQN, self).__init__() # Feature extractor if len(obs_space.shape) == 1: self.feature_extractor = FlatExtractor(obs_space) elif len(obs_space.shape) == 3: self.feature_extractor = NatureCnn(obs_space) else: raise NotImplementedErorr("This type of environment is not supported") # Neural network self.net = nn.Sequential( nn.Linear(self.feature_extractor.n_flatten, layer_size), nn.BatchNorm1d(layer_size), nn.ReLU(), nn.Linear(layer_size, layer_size), nn.BatchNorm1d(layer_size), nn.ReLU(), nn.Linear(layer_size, act_space.n), ) def forward(self, obs): return self.net(self.feature_extractor(obs))
memory.py
import random from collections import deque class Memory(object): def __init__(self, maxlen): self.memory = deque(maxlen=maxlen) def store(self, experience): self.memory.append(experience) def sample(self, n_samples): return zip(*random.sample(self.memory, n_samples)) def __len__(self): return len(self.memory)
utils.py
def wrap_input(arr, device, dtype=torch.float, reshape=False): output = torch.from_numpy(np.array(arr)).type(dtype).to(device) if reshape: output = output.reshape(-1, 1) return output def epsilon_greedy(start, end, n_steps, it): return max(start - (start - end) * (it / n_steps), end)
关键问题排查
- ε-贪婪策略逻辑完全颠倒:当前代码中,当随机数小于等于ε时使用模型选动作,否则随机探索。正确逻辑应该是:随机数≤ε时随机探索,否则用模型选最优动作。这会导致前期几乎不探索(ε初始为1.0,此时全用模型选动作,但模型完全没训练,输出随机),后期ε降到0.01时反而大量探索,彻底破坏探索-利用平衡,是训练失效的核心原因。
- 同时混用软更新与硬更新:代码中既做了每步的软更新(tau=0.01),又每隔200次迭代强制硬更新target网络。两种更新方式只需二选一:要么固定频率硬更新,要么每步软更新。同时用会导致target网络更新混乱,破坏DQN目标值的稳定性,引发损失爆炸。
- BatchNorm1d不适合当前场景:CartPole推理时是单样本输入,BatchNorm需要批量数据才能计算有效的均值和方差。训练时的批量统计量在前期样本不足时误差极大,推理时
model.eval()固定的统计量会进一步导致动作选择异常,污染经验池。对于这种简单环境,建议直接去掉BatchNorm,或改用LayerNorm。 - 终止状态next_state处理可优化:将终止状态的next_state设为全0的做法,结合BatchNorm可能加剧数值不稳定。更稳妥的方式是保留原始next_state,通过
(1 - dones)掩码让终止状态的target_qs直接等于rewards,符合DQN目标公式的原始定义。 - 经验采样维度需确认:
memory.sample()返回的是zip后的元组,wrap_input中np.array(arr)是否能正确将多个obs转换为(32,4)的张量?虽然这不是核心问题,但如果维度错误也会导致训练失败,可添加打印语句验证。
内容的提问来源于stack exchange,提问作者Squeemos
相关产品推荐
相关产品推荐

