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

基于DQN的Pong实现问题咨询:奖励函数与训练优化

Pong强化学习环境问题排查与优化方案

问题概述

基于OpenAI Gym搭建Pong游戏RL环境,训练AI控制玩家 paddle 时,智能体频繁获得负奖励。核心问题是奖励函数中「远离球的惩罚」触发过于频繁,甚至在球朝向玩家 paddle 移动时也错误生效;同时需要排查DQN实现及游戏逻辑中的潜在问题,提升训练稳定性与性能。


一、奖励函数修正(核心问题解决)

原奖励逻辑中,判断「朝向球移动」的条件完全颠倒,导致正确动作被惩罚、错误动作被奖励。同时需严格限定:仅当球朝向玩家 paddle 移动时(ball_dx > 0),才对移动方向进行奖惩;球朝AI移动时不做干预,让智能体专注于准备接球的站位。

修正后的奖励逻辑代码:

# Reward for moving towards the ball and penalty for moving away
if self.ball_dx > 0:  # Only apply when ball is moving towards player
    # Calculate direction to ball: if ball is above paddle, need to move up (action 0); if below, move down (action 1)
    should_move_up = self.ball.centery < self.player_paddle.centery
    should_move_down = self.ball.centery > self.player_paddle.centery
    
    # Correct action: move towards ball
    if (should_move_up and action == 0) or (should_move_down and action == 1):
        reward += 0.1  # Small positive reward for correct movement
    # Wrong action: move away from ball, or stay when need to move
    elif (should_move_up and action == 1) or (should_move_down and action == 0):
        reward -= 0.1  # Small penalty for wrong movement
    # No reward/penalty for staying when already aligned with ball

额外优化:奖励量级统一

将不同事件的奖励缩放至相近范围,避免大奖励主导训练:

  • 击球奖励从+5改为+1
  • 得分奖励从+10改为+2
  • 失球惩罚从-10改为-2

二、DQNAgent实现优化

1. 简化神经网络结构

原网络6层48神经元加Dropout,过度复杂易导致过拟合与训练缓慢。简化为2-3层隐藏层,去掉不必要的Dropout:

def _build_model(self):
    model = Sequential()
    model.add(Input(shape=(self.state_size,)))
    model.add(Dense(64, activation='relu'))
    model.add(Dense(64, activation='relu'))
    model.add(Dense(self.action_size, activation='linear'))
    model.compile(loss='mse', optimizer=Adam(learning_rate=self.learning_rate))
    return model

2. 批量处理经验回放

原replay方法逐样本训练,效率极低。改为批量计算目标值,一次性更新模型:

def replay(self, batch_size):
    if len(self.memory) < batch_size:
        return 0
    
    minibatch = random.sample(self.memory, batch_size)
    states = np.array([x[0][0] for x in minibatch])
    actions = np.array([x[1] for x in minibatch])
    rewards = np.array([x[2] for x in minibatch])
    next_states = np.array([x[3][0] for x in minibatch])
    dones = np.array([x[4] for x in minibatch])
    
    # Calculate target Q-values
    target_q = rewards + self.gamma * np.amax(self.target_model.predict(next_states, verbose=0), axis=1)
    target_q[dones] = rewards[dones]  # Terminal states have no future reward
    
    # Get current Q-values and update the action taken
    current_q = self.model.predict(states, verbose=0)
    current_q[np.arange(batch_size), actions] = target_q
    
    # Train in batch
    history = self.model.fit(states, current_q, epochs=1, verbose=0)
    avg_loss = history.history['loss'][0]
    
    # Decay epsilon
    if self.epsilon > self.epsilon_min:
        self.epsilon *= self.epsilon_decay
    
    return avg_loss

3. 定期更新Target模型

原仅初始化时更新Target模型,改为每固定步数或episode更新:

# 在训练循环中添加:
if episode % 10 == 0:  # 每10个episode更新一次
    agent.update_target_model()

4. 优化Epsilon衰减逻辑

改为按步数衰减,而非每轮回放衰减,更贴合训练进度:

# 在DQNAgent初始化中添加
self.step_count = 0
self.epsilon_decay_step = 0.99995  # 每步衰减

# 在act方法中更新
def act(self, state):
    self.step_count += 1
    if np.random.rand() <= self.epsilon:
        return np.random.randint(self.action_size)
    else:
        act_values = self.model.predict(state, verbose=0)
        return np.argmax(act_values[0])
    
# 在replay中替换原衰减逻辑为:
if self.epsilon > self.epsilon_min:
    self.epsilon = max(self.epsilon_min, self.epsilon * (self.epsilon_decay_step ** self.step_count))

三、游戏逻辑修复

1. 修正观测空间范围

原observation_space的low=0, high=255与实际状态值不符(如ball.centerx最大为640),导致Gym环境检查报错:

self.observation_space = spaces.Box(
    low=np.array([0, 0, 0, 0, -self.ball_speed, -self.ball_speed, -self.height, -self.height, 0, 0, 0, 0, -self.paddle_speed]),
    high=np.array([self.height, self.height, self.width, self.height, self.ball_speed, self.ball_speed, self.height, self.height, self.height, self.height, self.height, self.height, self.paddle_speed]),
    dtype=np.float32
)

2. 优化Paddle移动边界处理

原先调整centery再修正top/bottom,可能导致位置误差,改为直接限制centery范围:

# Move player paddle
if action == 0:
    self.player_paddle.centery -= self.paddle_speed
elif action == 1:
    self.player_paddle.centery += self.paddle_speed

# Ensure paddle stays within screen
self.player_paddle.centery = np.clip(self.player_paddle.centery, self.player_paddle.height//2, self.height - self.player_paddle.height//2)

3. 修复渲染函数重复初始化问题

原每次render都初始化pygame,导致窗口闪烁,将初始化移至__init__:

def __init__(self):
    # ... 原有代码 ...
    self.screen = None
    self.clock = None

def render(self, mode='human'):
    if mode == 'human':
        if self.screen is None:
            pygame.init()
            self.screen = pygame.display.set_mode((self.width, self.height))
            self.clock = pygame.time.Clock()
        
        self.screen.fill((0, 0, 0))
        pygame.draw.rect(self.screen, (255, 255, 255), self.player_paddle)
        pygame.draw.rect(self.screen, (255, 255, 255), self.ai_paddle)
        pygame.draw.ellipse(self.screen, (255, 255, 255), self.ball)
        pygame.draw.aaline(self.screen, (255, 255, 255), (self.width // 2, 0), (self.width // 2, self.height))
        pygame.display.flip()
        self.clock.tick(60)  # 控制帧率

四、训练性能提升建议

  • 状态归一化:将所有状态值缩放到[0,1]区间,例如:
    def get_state(self):
        state = np.array([...])  # 原有状态数组
        # 归一化
        state[0] /= self.height
        state[1] /= self.height
        state[2] /= self.width
        state[3] /= self.height
        state[4] /= self.ball_speed
        state[5] /= self.ball_speed
        state[6] /= self.height
        state[7] /= self.height
        state[8] /= self.height
        state[9] /= self.height
        state[10] /= self.height
        state[11] /= self.height
        state[12] /= self.paddle_speed
        return state.astype(np.float32)
    
  • 固定随机种子:设置numpy、tensorflow、random的随机种子,确保训练可复现。
  • 监控训练指标:记录每episode的总奖励、平均损失,绘制曲线观察收敛趋势。
  • 调整批量大小:使用32或64的batch_size,平衡训练稳定性与效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 07:29:49