You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Atari PongNoFrameskip环境下DQN模型无法学习的问题排查与技术求助

Alright, let's dig into why your DQN isn't learning to play Pong effectively even after 1000 episodes. I've gone through your code line by line and spotted several critical issues that are preventing your agent from improving. Here's a breakdown of the problems and actionable fixes:

1. Network Output Activation is Breaking Q-Value Estimation

This is the biggest red flag: your final fully connected layer uses a ReLU activation. DQN outputs Q-values, which can be positive or negative (to represent the value of bad actions). ReLU clamps all negative values to zero, so your network can never learn to avoid harmful moves.

Fix it by removing the ReLU from the last layer:

# Original code
output = F.relu(self.fc3(output))
# Updated code
output = self.fc3(output)

2. Target Network Updates Are Way Too Frequent

You're updating the target network every single time you call train(), which completely defeats its purpose. Target networks are meant to stay fixed for hundreds/thousands of steps to stabilize Q-value estimates.

Fix:

  1. Add a counter and update frequency to your Agent class:
    def __init__(self, s_space, a_space, device) -> None:
        # ... existing code ...
        self.target_update_counter = 0
        self.target_update_freq = 1000  # Update every 1000 steps
    
  2. Remove this line from your train() method:
    self.tgt_net.load_state_dict(self.evl_net.state_dict())
    
  3. Add this to your main training loop's step iteration:
    while not done:
        # ... existing step logic ...
        agent.target_update_counter += 1
        if agent.target_update_counter >= agent.target_update_freq:
            agent.tgt_net.load_state_dict(self.evl_net.state_dict())
            agent.target_update_counter = 0
    

3. Action Selection Logic Has a Bug

Your current action selection code is misinterpreting the network's output:

actions = self.evl_net(state).data.tolist()
action = actions.index(max(actions))

Since state is a batch of 1 (unsqueeze(0)), actions becomes a 2D list like [[q1, q2, q3]]. Using index(max(actions)) will always return 0 (the index of the inner list), not the index of the maximum Q-value.

Fix:

def select(self, state, a_space):
    # Remove batch dimension, move to CPU for numpy operations
    q_values = self.evl_net(state).squeeze(0).data.cpu().numpy()
    if random.random() <= self.epsilon:
        action = random.randint(0, a_space-1)
    else:
        action = np.argmax(q_values)  # Correctly get index of max Q-value
    return action

4. Memory Preprocessing Has Device and Indexing Issues

  • You're using a global device instead of self.device for the reward tensor, which can cause device mismatch errors.
  • Filtering out None next states breaks the index alignment with your nonDone_index logic, leading to incorrect target Q-values.

Fix:

Rewrite your data_pre_process method to handle done states without filtering:

def data_pre_process(self, batch_size):
    s_v = []
    a_v = []
    next_s_v = []
    r_v = []
    dones = []
    materials = random.sample(self.memory, batch_size)
    for t in materials:
        s_v.append(t[0])
        a_v.append(t[1])
        # Replace None next states with zeros (matches state shape)
        next_s_v.append(t[2] if t[2] is not None else np.zeros_like(t[0]))
        r_v.append(t[3])
        dones.append(t[4])
    
    # Use self.device consistently
    s_v = th.FloatTensor(s_v).to(self.device)
    a_v = th.LongTensor(a_v).unsqueeze(1).to(self.device)
    r_v = th.FloatTensor(r_v).to(self.device)
    next_s_v = th.FloatTensor(next_s_v).to(self.device)
    dones = th.BoolTensor(dones).to(self.device)
    
    return s_v, a_v, next_s_v, r_v, dones

Then update your train() method's target Q-value calculation:

def train(self, batch_size):
    s_v, a_v, next_s_v, r_v, dones = self.data_pre_process(batch_size)
    
    evl_Q_value = self.evl_net(s_v).gather(1, a_v)
    
    # Calculate target Q-values
    tgt_Q = self.tgt_net(next_s_v).max(1)[0].detach()
    # Set target Q to 0 for done states (no future reward)
    tgt_Q[dones] = 0.0
    tgt_Q_value = r_v + self.gamma * tgt_Q
    tgt_Q_value = tgt_Q_value.reshape(batch_size, 1)
    
    # ... existing loss and optimization logic ...

5. Training Trigger Timing Is Too Restrictive

You're only training when the episode number is a multiple of 4, which means you're missing most training opportunities. Instead, train every step once the memory buffer is large enough:

# In main loop's while not done:
if len(agent.memory) > batch_size:
    agent.train(batch_size)

Also, remove the unused state parameter from your train() method.

6. Minor Tweaks to Boost Convergence

  • Scale observations: Set scale_obs=True in AtariPreprocessing to normalize pixel values to 0-1, which helps the network learn faster.
  • Fixed batch size: Stop incrementing batch size every episode—stick to a fixed size (32 or 64) for stable training.
  • Remove unnecessary pool layers: Your current pooling after every conv layer shrinks the feature map too quickly. Consider removing some pooling steps or adjusting kernel/stride sizes to preserve more spatial information.

Once you fix these core issues, you should start seeing your agent's score improve over time. Pong typically takes 2000-5000 episodes to start consistently winning, but you'll notice upward trends long before that.

内容的提问来源于stack exchange,提问作者speedhawk1

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.28 18:47:49