用Pygame和PyTorch复刻经典AI实验:手把手教你搭建自己的Wumpus世界(附完整代码)
·
从零构建Wumpus世界:Pygame可视化与PyTorch强化学习实战指南
在人工智能教学领域,Wumpus世界如同围棋之于棋类游戏,是一个经典的基准测试环境。这个充满危险的洞穴系统不仅考验Agent的感知与决策能力,更为我们理解强化学习的核心机制提供了绝佳沙盒。本文将带你从零开始,用Pygame构建可视化交互界面,结合PyTorch实现深度Q学习算法,完整复现这个AI经典实验。
1. 环境搭建与基础架构
1.1 Wumpus世界核心规则解析
这个4×4网格的洞穴包含三个致命要素:
- Wumpus怪兽:静止的捕食者,会杀死进入同一房间的Agent
- 无底洞:随机分布的致命陷阱
- 黄金:唯一正向奖励源
Agent配备五种传感器:
- 臭气:相邻房间存在Wumpus
- 微风:相邻房间存在无底洞
- 金光:当前房间存在黄金
- 撞击:箭矢击中墙壁
- 嚎叫:Wumpus被射杀
奖励机制设计如下表:
| 行为 | 奖励值 |
|---|---|
| 携带黄金离开洞穴 | +1000 |
| 掉入无底洞/被吞噬 | -1000 |
| 每次移动 | -1 |
| 射箭 | -10 |
1.2 Pygame环境初始化
首先建立游戏窗口和基础类结构:
import pygame
import random
class GameObject(pygame.sprite.Sprite):
def __init__(self, image_path, position, size=(50,50)):
super().__init__()
self.image = pygame.transform.scale(
pygame.image.load(image_path), size)
self.rect = self.image.get_rect()
self.rect.topleft = position
class Room:
def __init__(self, x, y):
self.x = x
self.y = y
self.has_pit = False
self.has_gold = False
self.has_wumpus = False
self.stench = False
self.breeze = False
提示:使用精灵类(Sprite)可以高效处理游戏对象的碰撞检测和渲染
2. 游戏逻辑实现
2.1 世界生成算法
采用Fisher-Yates洗牌算法确保元素分布随机性:
def generate_world(grid_size=4):
positions = [(x,y) for x in range(grid_size) for y in range(grid_size)]
positions.remove((0,0)) # 移除起始点
world = [[Room(x,y) for y in range(grid_size)]
for x in range(grid_size)]
# 随机分配元素
random.shuffle(positions)
wumpus_pos = positions.pop()
world[wumpus_pos[0]][wumpus_pos[1]].has_wumpus = True
gold_pos = positions.pop()
world[gold_pos[0]][gold_pos[1]].has_gold = True
pit_count = min(3, len(positions))
for _ in range(pit_count):
pit_pos = positions.pop()
world[pit_pos[0]][pit_pos[1]].has_pit = True
return world
2.2 传感器系统实现
传感器数据通过位掩码方式编码:
def get_sensors(world, agent_pos):
x, y = agent_pos
current_room = world[x][y]
sensors = {
'stench': False,
'breeze': False,
'glitter': current_room.has_gold,
'bump': False,
'scream': False
}
# 检测相邻房间
for dx, dy in [(-1,0),(1,0),(0,-1),(0,1)]:
nx, ny = x+dx, y+dy
if 0 <= nx < len(world) and 0 <= ny < len(world[0]):
neighbor = world[nx][ny]
sensors['stench'] |= neighbor.has_wumpus
sensors['breeze'] |= neighbor.has_pit
return sensors
3. 强化学习Agent设计
3.1 DQN网络架构
采用双网络结构解决训练稳定性问题:
import torch
import torch.nn as nn
import torch.optim as optim
class DQN(nn.Module):
def __init__(self, input_size, hidden_size, output_size):
super(DQN, self).__init__()
self.fc1 = nn.Linear(input_size, hidden_size)
self.fc2 = nn.Linear(hidden_size, hidden_size)
self.fc3 = nn.Linear(hidden_size, output_size)
def forward(self, x):
x = torch.relu(self.fc1(x))
x = torch.relu(self.fc2(x))
return self.fc3(x)
class Agent:
def __init__(self, state_dim, action_dim):
self.policy_net = DQN(state_dim, 128, action_dim)
self.target_net = DQN(state_dim, 128, action_dim)
self.target_net.load_state_dict(self.policy_net.state_dict())
self.optimizer = optim.Adam(self.policy_net.parameters(), lr=0.001)
self.memory = ReplayBuffer(10000)
3.2 状态编码策略
将游戏状态编码为21维向量:
[agent_x, agent_y, orientation,
has_arrow, has_gold,
wumpus_0_x, wumpus_0_y, ...,
pit_0_x, pit_0_y, ...,
gold_x, gold_y]
注意:相比原始像素输入,这种特征工程大幅提升训练效率
4. 训练流程与可视化
4.1 训练参数配置
关键超参数设置参考:
| 参数 | 值 | 说明 |
|---|---|---|
| batch_size | 64 | 经验回放采样量 |
| gamma | 0.99 | 折扣因子 |
| eps_start | 0.9 | 初始探索率 |
| eps_end | 0.05 | 最小探索率 |
| eps_decay | 200 | 探索率衰减步数 |
| target_update | 10 | 目标网络更新频率 |
4.2 实时训练可视化
通过Pygame实现训练过程动态展示:
def render_training(episode, score, epsilon):
screen.fill((255,255,255))
# 绘制训练曲线
pygame.draw.lines(screen, (0,0,255), False,
[(x,100-y) for x,y in enumerate(scores[-100:])], 2)
# 显示关键指标
font = pygame.font.SysFont(None, 24)
texts = [
f"Episode: {episode}",
f"Avg Score: {np.mean(scores[-10:]):.1f}",
f"Epsilon: {epsilon:.2f}",
f"Best: {max(scores):.1f}"
]
for i, text in enumerate(texts):
text_surface = font.render(text, True, (0,0,0))
screen.blit(text_surface, (10, 10 + i*25))
pygame.display.flip()
5. 进阶优化技巧
5.1 课程学习策略
分阶段训练方案:
- 基础导航:仅包含移动动作,目标找到黄金
- 危险感知:加入无底洞,学习避开陷阱
- 完整挑战:引入Wumpus和射箭机制
- 记忆测试:增大地图尺寸至6×6
5.2 混合探索策略
结合ε-greedy和Boltzmann探索:
def select_action(state, net, epsilon, temp=1.0):
if random.random() < epsilon:
return random.randint(0, NUM_ACTIONS-1)
else:
with torch.no_grad():
q_values = net(state)
probs = torch.softmax(q_values/temp, dim=1)
return torch.multinomial(probs, 1).item()
6. 调试与性能优化
常见问题解决方案:
-
训练不稳定:
- 增加目标网络更新频率
- 减小学习率
- 扩大经验回放缓冲区
-
Agent过于保守:
- 调整奖励函数权重
- 引入探索奖励
- 尝试优先经验回放
# 优先经验回放示例
class PrioritizedReplayBuffer:
def __init__(self, capacity, alpha=0.6):
self.capacity = capacity
self.alpha = alpha
self.buffer = []
self.priorities = np.zeros(capacity)
self.pos = 0
def add(self, transition, error):
max_prio = self.priorities.max() if self.buffer else 1.0
if len(self.buffer) < self.capacity:
self.buffer.append(transition)
else:
self.buffer[self.pos] = transition
self.priorities[self.pos] = (abs(error) + 1e-5) ** self.alpha
self.pos = (self.pos + 1) % self.capacity
在项目开发过程中,最耗时的部分往往是奖励函数的精细调参。通过实践发现,给"发现黄金"设置中间奖励(如+100)能显著加快早期训练速度,而将"射箭惩罚"从固定值改为动态值(基于剩余箭数)则能改善资源管理能力。
更多推荐


所有评论(0)