PyTorch Snake AI循环移动不收集食物问题排查求助
贪吃蛇AI训练异常问题
基于Pygame实现贪吃蛇游戏,蛇初始位于屏幕中央,收集食物后变长;使用PyTorch(支持CUDA)结合Q-Learning算法训练AI玩该游戏,实现了自定义并行环境管理器加速训练,并添加GUI可视化训练过程。但训练时AI始终循环移动,既不探索地图也不收集食物。已尝试调整Q-Learning算法、删除__pycache__重置训练、适配CPU/GPU设备,问题仍未解决。
完整代码
贪吃蛇游戏代码(snake_game.py)
import pygame import random from enum import Enum from collections import namedtuple import numpy as np pygame.init() font = pygame.font.Font(None, 36) class Direction(Enum): RIGHT = 1 LEFT = 2 UP = 3 DOWN = 4 Point = namedtuple('Point', 'x y') # RGB colors WHITE = (255, 255, 255) RED = (200, 0, 0) BLUE1 = (0, 0, 255) BLUE2 = (0, 100, 255) BLACK = (0, 0, 0) BLOCK_SIZE = 20 SPEED = 10 class SnakeGameAI: def __init__(self, w=640, h=480): self.w = w self.h = h self.reset() def reset(self): self.direction = Direction.RIGHT self.head = Point(self.w // 2, self.h // 2) self.snake = [self.head, Point(self.head.x - BLOCK_SIZE, self.head.y), Point(self.head.x - (2 * BLOCK_SIZE), self.head.y)] self.score = 0 self.food = None self._place_food() self.frame_iteration = 0 return self.get_state() def _place_food(self): while True: x = random.randint(0, (self.w - BLOCK_SIZE) // BLOCK_SIZE) * BLOCK_SIZE y = random.randint(0, (self.h - BLOCK_SIZE) // BLOCK_SIZE) * BLOCK_SIZE self.food = Point(x, y) if self.food not in self.snake: break def play_step(self, action): self.frame_iteration += 1 self._move(action) self.snake.insert(0, self.head) reward = 0 done = False if self.is_collision() or self.frame_iteration > 100 * len(self.snake): done = True reward = -10 return self.get_state(), reward, done if self.head == self.food: self.score += 1 reward = 10 self._place_food() else: self.snake.pop() return self.get_state(), reward, done def is_collision(self, pt=None): if pt is None: pt = self.head if pt.x >= self.w or pt.x < 0 or pt.y >= self.h or pt.y < 0: return True if pt in self.snake[1:]: return True return False def _move(self, action): clock_wise = [Direction.RIGHT, Direction.DOWN, Direction.LEFT, Direction.UP] idx = clock_wise.index(self.direction) if np.array_equal(action, [1, 0, 0]): # Move straight new_dir = clock_wise[idx] elif np.array_equal(action, [0, 1, 0]): # Turn right new_dir = clock_wise[(idx + 1) % 4] else: # Turn left new_dir = clock_wise[(idx - 1) % 4] # Prevent the snake from reversing if (self.direction == Direction.RIGHT and new_dir == Direction.LEFT) or \ (self.direction == Direction.LEFT and new_dir == Direction.RIGHT) or \ (self.direction == Direction.UP and new_dir == Direction.DOWN) or \ (self.direction == Direction.DOWN and new_dir == Direction.UP): new_dir = self.direction self.direction = new_dir x = self.head.x y = self.head.y if self.direction == Direction.RIGHT: x += BLOCK_SIZE elif self.direction == Direction.LEFT: x -= BLOCK_SIZE elif self.direction == Direction.DOWN: y += BLOCK_SIZE elif self.direction == Direction.UP: y -= BLOCK_SIZE self.head = Point(x, y) def get_state(self): head_x, head_y = self.head.x, self.head.y food_x, food_y = self.food.x, self.food.y direction = self.direction.value # Features: normalized distances and direction return np.array([ (food_x - head_x) / self.w, (food_y - head_y) / self.h, direction / 4, len(self.snake) / ((self.w // BLOCK_SIZE) * (self.h // BLOCK_SIZE)) ])
Q-Learning模型代码(q_learning.py)
import torch import torch.nn as nn import torch.optim as optim import os import numpy as np # Device configuration device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') class Linear_QNet(nn.Module): def __init__(self, input_size, hidden_size, output_size): super().__init__() self.linear1 = nn.Linear(input_size, hidden_size) self.linear2 = nn.Linear(hidden_size, output_size) def forward(self, x): x = torch.relu(self.linear1(x)) x = self.linear2(x) return x def save(self, file_name='model.pth'): model_folder_path = './model' if not os.path.exists(model_folder_path): os.makedirs(model_folder_path) file_name = os.path.join(model_folder_path, file_name) torch.save(self.state_dict(), file_name) class QTrainer: def __init__(self, model, lr, gamma, epsilon_start=1.0, epsilon_end=0.01, epsilon_decay=0.995): self.lr = lr self.gamma = gamma self.epsilon = epsilon_start # Initial exploration rate self.epsilon_end = epsilon_end # Minimum exploration rate self.epsilon_decay = epsilon_decay # Decay factor for epsilon self.model = model.to(device) # Ensure model is on the correct device self.optimizer = optim.Adam(model.parameters(), lr=self.lr) self.criterion = nn.MSELoss() def get_action(self, state): if np.random.rand() < self.epsilon: # Explore: select a random action return np.random.randint(0, 3) # Assuming 3 possible actions else: # Exploit: select the best action based on the model state = torch.tensor(state, dtype=torch.float).to(device) with torch.no_grad(): q_values = self.model(state.unsqueeze(0)) # Add batch dimension return torch.argmax(q_values).item() def train_step(self, state, action, reward, next_state, done): state = torch.tensor(state, dtype=torch.float).to(device) next_state = torch.tensor(next_state, dtype=torch.float).to(device) action = torch.tensor(action, dtype=torch.long).to(device) reward = torch.tensor(reward, dtype=torch.float).to(device) done = torch.tensor(done, dtype=torch.float).to(device) if len(state.shape) == 1: state = torch.unsqueeze(state, 0) next_state = torch.unsqueeze(next_state, 0) action = torch.unsqueeze(action, 0) reward = torch.unsqueeze(reward, 0) done = (done, ) pred = self.model(state) target = pred.clone() for idx in range(len(done)): Q_new = reward[idx] if not done[idx]: Q_new = reward[idx] + self.gamma * torch.max(self.model(next_state[idx])) target[idx][torch.argmax(action[idx]).item()] = Q_new self.optimizer.zero_grad() loss = self.criterion(target, pred) loss.backward() self.optimizer.step() # Update epsilon (exploration rate) self.epsilon = max(self.epsilon_end, self.epsilon * self.epsilon_decay)
主训练代码(main.py)
import pygame import numpy as np import torch from snake_game import SnakeGameAI from q_learning import Linear_QNet, QTrainer # Hyperparameters INPUT_SIZE = 4 HIDDEN_SIZE = 256 OUTPUT_SIZE = 3 LEARNING_RATE = 0.001 GAMMA = 0.9 NUM_EPISODES = 1000 NUM_ENVS = 1 # Adjust this for parallel environments if needed # Device configuration device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # Colors BLACK = (0, 0, 0) WHITE = (255, 255, 255) RED = (200, 0, 0) BLUE1 = (0, 0, 255) BLUE2 = (0, 100, 255) BLOCK_SIZE = 20 class ParallelEnvManager: def __init__(self, env_class, num_envs): self.envs = [env_class() for _ in range(num_envs)] def reset(self): return np.array([env.reset() for env in self.envs]) def step(self, actions): next_state, rewards, done = [], [], [] for env, action in zip(self.envs, actions): state, reward, d = env.play_step(action) next_state.append(state) rewards.append(reward) done.append(d) return np.array(next_state), np.array(rewards), np.array(done) def close(self): pygame.quit() def main(): pygame.init() width, height = 640, 480 screen = pygame.display.set_mode((width, height)) pygame.display.set_caption('Snake Q-Learning') clock = pygame.time.Clock() model = Linear_QNet(INPUT_SIZE, HIDDEN_SIZE, OUTPUT_SIZE).to(device) trainer = QTrainer(model, LEARNING_RATE, GAMMA) env_manager = ParallelEnvManager(SnakeGameAI, NUM_ENVS) font = pygame.font.Font(None, 36) running = True episode = 0 while running and episode < NUM_EPISODES: state = env_manager.reset() done = np.zeros(NUM_ENVS, dtype=bool) while not np.all(done): actions = [trainer.get_action(s) for s in state] next_state, rewards, dones = env_manager.step(actions) trainer.train_step(state, actions, rewards, next_state, dones) state = next_state # Update display screen.fill(BLACK) for env in env_manager.envs: # Draw snake for pt in env.snake: pygame.draw.rect(screen, BLUE1, pygame.Rect(pt.x, pt.y, BLOCK_SIZE, BLOCK_SIZE)) pygame.draw.rect(screen, BLUE2, pygame.Rect(pt.x + 4, pt.y + 4, 12, 12)) # Draw food pygame.draw.rect(screen, RED, pygame.Rect(env.food.x, env.food.y, BLOCK_SIZE, BLOCK_SIZE)) score_text = font.render(f"Episode: {episode} | Score: {env_manager.envs[0].score}", True, WHITE) screen.blit(score_text, (10, 10)) pygame.display.flip() for event in pygame.event.get(): if event.type == pygame.QUIT: running = False break episode += 1 pygame.display.set_caption(f'Snake Q-Learning - Episode {episode}') env_manager.close() model.save('snake_model.pth') pygame.quit() if __name__ == '__main__': main()
问题根源与修复方案
1. 状态特征缺失关键信息
当前get_state()仅提供食物相对距离、方向和蛇身长度,缺少障碍物/边界的危险感知和食物相对方向的明确标识,AI无法判断移动风险和食物位置,导致盲目循环。扩展状态特征:
def get_state(self): head_x, head_y = self.head.x, self.head.y food_x, food_y = self.food.x, self.food.y dir_right = self.direction == Direction.RIGHT dir_left = self.direction == Direction.LEFT dir_up = self.direction == Direction.UP dir_down = self.direction == Direction.DOWN # 危险感知:前方、右方、左方是否有障碍物/边界 danger_straight = ( (dir_right and self.is_collision(Point(head_x + BLOCK_SIZE, head_y))) or (dir_left and self.is_collision(Point(head_x - BLOCK_SIZE, head_y))) or (dir_up and self.is_collision(Point(head_x, head_y - BLOCK_SIZE))) or (dir_down and self.is_collision(Point(head_x, head_y + BLOCK_SIZE))) ) danger_right = ( (dir_up and self.is_collision(Point(head_x + BLOCK_SIZE, head_y))) or (dir_right and self.is_collision(Point(head_x, head_y + BLOCK_SIZE))) or (dir_down and self.is_collision(Point(head_x - BLOCK_SIZE, head_y))) or (dir_left and self.is_collision(Point(head_x, head_y - BLOCK_SIZE))) ) danger_left = ( (dir_up and self.is_collision(Point(head_x - BLOCK_SIZE, head_y))) or (dir_left and self.is_collision(Point(head_x, head_y + BLOCK_SIZE))) or (dir_down and self.is_collision(Point(head_x + BLOCK_SIZE, head_y))) or (dir_right and self.is_collision(Point(head_x, head_y - BLOCK_SIZE))) ) # 食物方向:相对蛇头的位置 food_left = food_x < head_x food_right = food_x > head_x food_up = food_y < head_y food_down = food_y > head_y return np.array([ danger_straight, danger_right, danger_left, dir_left, dir_right, dir_up, dir_down, food_left, food_right, food_up, food_down ], dtype=int)
同时更新main.py中的INPUT_SIZE = 11。
2. 动作与游戏逻辑不匹配
get_action()返回整数索引,但游戏_move()方法期望长度为3的one-hot数组,导致动作解析完全错误。修复get_action():
def get_action(self, state): if np.random.rand() < self.epsilon: # 生成随机one-hot动作 action_idx = np.random.randint(0, 3) action = np.zeros(3) action[action_idx] = 1 return action else: state = torch.tensor(state, dtype=torch.float).to(device) with torch.no_grad(): q_values = self.model(state.unsqueeze(0)) action_idx = torch.argmax(q_values).item() action = np.zeros(3) action[action_idx] = 1 return action
3. 训练逻辑的批量处理错误
QTrainer.train_step()中对动作的处理错误,整数动作无法用torch.argmax()解析。修改为直接用动作索引更新Q值:
def train_step(self, state, action, reward, next_state, done): state = torch.tensor(state, dtype=torch.float).to(device) next_state = torch.tensor(next_state, dtype=torch.float).to(device) # 动作转为索引张量 action = torch.tensor([np.argmax(a) for a in action], dtype=torch.long).to(device) reward = torch.tensor(reward, dtype=torch.float).to(device) done = torch.tensor(done, dtype=torch.bool).to(device) if len(state.shape) == 1: state = torch.unsqueeze(state, 0) next_state = torch.unsqueeze(next_state, 0) action = torch.unsqueeze(action, 0) reward = torch.unsqueeze(reward, 0) done = torch.unsqueeze(done, 0) pred = self.model(state) target = pred.clone() for idx in range(len(done)):
相关产品推荐
相关产品推荐

