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

基于DQL的强化学习贪吃蛇项目无学习效果求助

贪吃蛇DQL训练问题排查与改进建议

问题背景

使用带经验回放的Deep Q-Learning(DQL)训练贪吃蛇,状态由苹果坐标、蛇头及蛇身坐标组成的向量表示。经过一周多的尝试,更换过多种网络架构、超参数和奖励函数,模型始终无有效学习:蛇要么频繁撞墙,要么原地振荡。以下是对代码的问题排查及改进方案。

现有DQL实现代码

from collections import deque
import numpy as np
import random
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense

class DQL:
    def __init__(self, model, actions, discount_factor=0.95, exploration_rate=0.2, memory_size=100000, batch_size=20, decay_rate=0.995, base_exploration_rate = 0.1):
        # 神经网络模型
        self.model = model
        self.actions = actions
        # 折扣因子gamma
        self.discount_factor = discount_factor
        # ε-贪心探索率
        self.exploration_rate = exploration_rate
        # 经验回放缓冲区
        self.memory = deque(maxlen=memory_size)
        self.batch_size = batch_size
        # 探索率衰减系数
        self.decay_rate = decay_rate
        self.base_exploration_rate = base_exploration_rate

    def get_action(self, state, direction,length):
        if np.random.rand() < self.base_exploration_rate + self.exploration_rate:
            # 随机选择动作(排除当前方向)
            l = list(range(0,len(self.actions)))
            l.remove(direction)
            action = l[np.random.choice(len(self.actions)-1)]
        else:
            # 根据模型选择最优动作
            q_values = self.model.predict(state,verbose = 0)
            sorted = q_values.argsort()
            action = sorted[0][0]
            if(action == direction and length>1):
                action= sorted[0][1]
        return action
    
    def add_memory(self, state, action, reward, next_state, done):
        # 冗余代码,已移除
        self.memory.append((state, action, reward, next_state, done))

    
    def train(self):
        if len(self.memory) < self.batch_size:
            # 经验不足,跳过训练
            return
        # 随机采样经验
        batch = random.sample(self.memory, self.batch_size)
        states, actions, rewards, next_states, dones = [], [], [], [], []
        for state, action, reward, next_state, done in batch:
            states.append(state[0])
            actions.append(action)
            rewards.append(reward)
            next_states.append(next_state[0])
            dones.append(done)
        states = np.array(states)
        next_states = np.array(next_states)

        # 计算目标Q值
        next_q_values = self.model.predict(next_states,verbose = 0)
        target_q_values = np.zeros((self.batch_size,len(self.actions)))

        for i in range(self.batch_size):
            if dones[i]:
                target_q_values[i][actions[i]] = rewards[i]
            else:
                target_q_values[i][actions[i]] = rewards[i] + self.discount_factor * max(next_q_values[i])

        # 训练模型
        self.model.fit(states, target_q_values, verbose=0)

现有游戏实现代码

import pygame
import sys
import random
import numpy as np
from copy import deepcopy
import os
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense
import keras

# 蛇块大小
block_size = 10
# 游戏窗口尺寸
width = 150
height = 150

pygame.init()
# 创建窗口
screen = pygame.display.set_mode((width, height))
pygame.display.set_caption("Snake Game")

# 颜色定义
white = (255, 255, 255)
black = (0, 0, 0)
red = (255, 0, 0)

# 帧率控制
clock = pygame.time.Clock()
font = pygame.font.Font(None, 30)
fps = 10

def game_over():
    text = font.render("Game Over!", True, red)
    screen.blit(text, [width/2 - text.get_width()/2, height/2 - text.get_height()/2])

def display_score(score,gen,s):
    text = font.render(f"Gen: {gen} Length : {score} Score: {s}", True, black)
    screen.blit(text, [0,0])

def draw_snake(snake_list):
    for block in snake_list:
        pygame.draw.rect(screen, black, [block[0], block[1], block_size, block_size])

def generate_food(snake_list):
    # 在非蛇身位置生成食物
    food_x, food_y = None, None
    while food_x is None or food_y is None or (food_x, food_y) in snake_list:
        food_x = round(random.randrange(0, width - block_size) / 10.0) * 10.0
        food_y = round(random.randrange(0, height - block_size) / 10.0) * 10.0
    return food_x, food_y

# 动作映射(索引对应:0=up,1=down,2=left,3=right)
actions = {"up":(-1,0),"down":(1,0),"left":(0,-1),"right":(0,1)}
# 原状态输入维度(冗余,需优化)
input_size = 2*width*height//block_size**2+2

def initNN():
    # 初始化全连接神经网络
    model = Sequential()
    model.add(Dense(256, input_shape=(input_size,), activation='sigmoid'))
    model.add(Dense(128 , activation = 'tanh'))
    model.add(Dense(len(actions), activation='linear'))
    model.compile(loss='mse', optimizer='adam', metrics=['accuracy'])
    print(model.summary())
    return model

# 生成状态向量(原实现冗余,需优化)
def state(snake_list,apple):
    s = np.array(snake_list)
    input = np.zeros(input_size)
    input[0],input[1] = apple[0]/width,apple[1]/height
    for u in range(len(snake_list)):
        if(2*u+2>=input_size):
            break
        input[2*u+2],input[2*u+3]= s[len(snake_list)-1-u][0]/width,s[len(snake_list)-1-u][1]/height
    input=input.reshape(1,input_size)
    return input

def normalized_distance(u,v,food_x,food_y):
    return np.sqrt((((u-food_x)/width)**2+((v-food_y)/height)**2)/2)

def inBounds(u,v):
    return 0 <= u < width and 0 <= v < height

def gaussian_aroundone(x,alpha):
    return(np.exp(-alpha*(x-1)**2))

# 奖励函数(原设计混乱,需优化)
def reward(action, snake_list):
    copy = deepcopy(snake_list)
    p = copy[-1]
    a = list(actions.values())
    (u,v)=(a[action][1]*block_size+p[0],a[action][0]*block_size+p[1])
    penalty_touch_self = -1 if [u,v] in snake_list else 0
    copy.append([u,v])
    del copy[0]
    
    global food_x, food_y
    reward_distance = 1-normalized_distance(u,v,food_x,food_y)
    gass_reward = gaussian_aroundone(reward_distance,10)
    reward_eat = 1 if u == food_x and v == food_y else 0
    penalty_distance = -1 if normalized_distance(u,v,food_x,food_y) > normalized_distance(p[0], p[1], food_x, food_y) else 1
    penalty_wall = -1 if not inBounds(u,v) else 0
    
    c = np.ones(5)
    c[1] = 10
    c[3]=5
    total_reward =(c[0]* gass_reward + c[1]*reward_eat + c[2]*penalty_distance + c[3]*penalty_wall + c[4]*penalty_touch_self)/c.sum()
    return total_reward

# 模型加载/初始化
filename = f"{width} {height} DeepQ.h5"
if(os.path.exists(f"./{filename}")):
    print("model already exists ")
    model = keras.models.load_model(f"./{filename}")
else: 
    model = initNN()

# 初始化DQL
dql = DQL(model,actions.values())
food_x, food_y = generate_food([])   

def main(gen,length):
    # 初始化蛇位置
    snake_x = (width//block_size)//2 * block_size   
    snake_y = (height//block_size)//2 * block_size   
    snake_list = [[snake_x,snake_y]]
    
    global food_x,food_y
    episode_length =0 
    snake_length =length
    # 初始动作索引(对应right)
    current_action_idx = 3
    done = False
    score = 0
    acts = list(actions.values())
    
    while True:
        for event in pygame.event.get():
            if event.type == pygame.KEYDOWN and event.key == pygame.K_SPACE:
                # 避免递归,改为重启循环
                return main(gen,length)
            if event.type == pygame.QUIT:
                pygame.quit()
                sys.exit()

        St1 = state(snake_list,[food_x,food_y]) 
        current_action_idx = dql.get_action(St1, current_action_idx, snake_length)
        r = reward(current_action_idx, snake_list)

        # 移动蛇
        snake_x+=acts[current_action_idx][1]*block_size
        snake_y+=acts[current_action_idx][0]*block_size
        snake_list.append([snake_x, snake_y])
        if len(snake_list) > snake_length:
            del snake_list[0]

        # 碰撞检测
        if not inBounds(snake_x, snake_y):
            done = True
        for block in snake_list[:-1]:
            if block[0] == snake_x and block[1] == snake_y:
                done = True
                break

        # 画面渲染
        screen.fill(white)
        pygame.draw.rect(screen, red, [food_x, food_y, block_size, block_size])
        draw_snake(snake_list)
        St2 = state(snake_list,[food_x,food_y])
        # 存入经验回放
        dql.add_memory(St1, current_action_idx, r, St2, done)
        pygame.display.update()

        # 吃食物逻辑
        if snake_x == food_x and snake_y == food_y:
            food_x, food_y = generate_food(snake_list)
            snake_length += 1
            score+=1
        episode_length+=1

        if done:
            return [snake_length,episode_length,score]
        
        clock.tick(fps)

# 训练参数
num_episodes =10000000
max_score =0 
initial_max_length = 2
max_possible_length = width*height//block_size**2

# 训练循环
for i in range(num_episodes):
    # 启动一局游戏
    episode_result = main(i, np.random.randint(1, initial_max_length+1))
    max_score = max(episode_result[2], max_score)
    # 每200局重置食物位置
    if i%200 ==0:
        food_x, food_y = generate_food([])   
    # 每1000局保存模型并增加初始蛇长度
    if (i+1)%1000==0:
        model.save(f"{width} {height} DeepQ.h5")
        if initial_max_length < max_possible_length:
            initial_max_length+=1
    # 每局结束后训练一次
    dql.train()
    # 衰减探索率(移到此处,避免每次训练都衰减)
    dql.exploration_rate = max(dql.base_exploration_rate, dql.exploration_rate * dql.decay_rate)

核心问题与改进方案

1. 状态表示冗余低效

原状态将苹果+整条蛇的所有坐标拼接,输入维度高达452,模型难以提取有效特征。优化方向:

  • 只保留关键信息:苹果相对蛇头的方向/距离、蛇头上下左右是否有障碍物(墙/自身)、蛇的移动方向
  • 示例优化后的状态函数:
    def state(snake_list, apple):
        head_x, head_y = snake_list[-1]
        food_x, food_y = apple
        # 相对方向(上下左右是否有食物)
        food_up = 1 if head_y > food_y else 0
        food_down = 1 if head_y < food_y else 0
        food_left = 1 if head_x > food_x else 0
        food_right = 1 if head_x < food_x else 0
        # 障碍物检测(上下左右是否是墙/自身)
        dirs = [(-1,0),(1,0),(0,-1),(0,1)]
        obstacles = []
        for dx, dy in dirs:
            check_x = head_x + dx*block_size
            check_y = head_y + dy*block_size
            hit_wall = not inBounds(check_x, check_y)
            hit_self = [check_x, check_y] in snake_list[:-1]
            obstacles.append(1 if hit_wall or hit_self else 0)
        # 蛇的移动方向
        if len(snake_list) >1:
            prev_x, prev_y = snake_list[-2]
            dir_up = 1 if head_y < prev_y else 0
            dir_down = 1 if head_y > prev_y else 0
            dir_left = 1 if head_x < prev_x else 0
            dir_right = 1 if head_x > prev_x else 0
        else:
            dir_up, dir_down, dir_left, dir_right = 0,0,0,1
        # 拼接成状态向量
        state_vec = np.array([food_up, food_down, food_left, food_right] + obstacles + [dir_up, dir_down, dir_left, dir_right])
        return state_vec.reshape(1, -1)
    

2. ε-贪心策略逻辑错误

  • 原代码探索率计算base_exploration_rate + exploration_rate可能超过1,导致全程随机探索,需限制最大值
  • 随机动作应排除反方向而非当前方向(比如当前向右,不能直接向左,但可以继续向右/向上/向下)
  • 选最优动作时,argsort()返回升序索引,原代码取最小Q值,应改为取最大Q值

3. 训练稳定性优化

  • 增加固定目标网络:每隔N步复制当前模型作为目标网络,计算目标Q值时使用目标网络,避免训练震荡
  • 探索率衰减应在每局结束后执行,而非每次训练步骤
  • 经验回放可增加优先级采样(可选),提升重要经验的利用率

4. 奖励函数简化清晰

原奖励函数混合过多冗余项,信号模糊,优化方向:

def reward(action, snake_list):
    head_x, head_y = snake_list[-1]
    a = list(actions.values())
    new_x = head_x + a[action][1]*block_size
    new_y = head_y + a[action][0]*block_size
    global food_x, food_y

    # 撞墙/撞自身:大惩罚
    if not inBounds(new_x, new_y) or [new_x, new_y] in snake_list[:-1]:
        return -10
    # 吃食物:大奖励
    if new_x == food_x and new_y == food_y:
        return 20
    # 靠近食物:小奖励;远离食物:小惩罚
    prev_dist = normalized_distance(head_x, head_y, food_x, food_y)
    new_dist = normalized_distance(new_x, new_y, food_x, food_y)
    return 1 if new_dist < prev_dist else -1

5. 训练流程优化

  • 初始阶段固定蛇的初始长度为2,让模型先学习基础移动和找食物的行为
  • 增加训练日志:每100局输出当前探索率、平均得分、最长存活时间,方便监控训练进度
  • 避免递归重启游戏,改用循环结构

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 04:18:19