当前位置:首页>python>Python奖惩算法实战:实现Q-learning,让智能体学会“避坑寻宝”

Python奖惩算法实战:实现Q-learning,让智能体学会“避坑寻宝”

  • 2026-10-04 23:52:34
Python奖惩算法实战:实现Q-learning,让智能体学会“避坑寻宝”

奖惩算法的核心思想非常直观:智能体(Agent)在环境中采取动作,环境会返回一个奖励或惩罚信号。智能体的目标就是通过不断试错,找到一套策略,使得长期累积的奖励最大。

举个生活中的例子:你教小狗握手,做对了给零食(奖励),做错了不给(惩罚)。小狗慢慢就会学会握手。在程序中,我们让一个智能体在网格世界里寻找宝藏,路上有陷阱,走错路会扣分,最终找到宝藏获得高分。

本文用 Python 实现一个经典的 Q-learning 奖惩算法,展示智能体是如何“学会”最优路径的。

一、环境设计:一个 5×5 的网格世界

我们设计一个 5×5 的网格:

  • 起点:左上角 (0, 0)

  • 终点(宝藏):右下角 (4, 4),到达奖励 +10

  • 陷阱:三个位置,踩到惩罚 -10

  • 普通格子:每走一步惩罚 -1,鼓励智能体尽快到达终点

  • 出界:如果智能体试图走出网格,惩罚 -2,并停在原地

这样智能体需要学会:避开陷阱,用最短路径到达终点。

二、Q-learning 奖惩算法原理

Q-learning 是一种无模型的强化学习算法,它维护一张 Q 表,记录在每个状态 s 下采取动作 a 的预期价值。

更新公式如下:

其中:

  • s:当前状态

  • a:当前动作

  • r:环境返回的即时奖励/惩罚

  • s':下一状态

  • α:学习率,控制每次更新幅度

  • γ:折扣因子,表示未来奖励的重要性

简单理解:如果某个动作带来了正奖励,那么 Q 值会增大,智能体以后更倾向于选择这个动作;如果某个动作带来了惩罚,Q 值会减小,智能体以后会尽量避免它。

三、完整代码实现

下面是一个完整的 Python 程序,包含环境定义、Q-learning 训练和可视化输出。

import numpy as npimport matplotlib.pyplot as pltplt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'PingFang SC', 'Noto Sans CJK SC', 'DejaVu Sans']plt.rcParams['axes.unicode_minus'] = Falsenp.random.seed(42)class GridWorld:    """5x5 网格世界环境"""    def __init__(self, size=5, start=(0,0), goal=(4,4),                 traps=None, step_penalty=-1, goal_reward=10,                 trap_penalty=-10, out_of_bounds_penalty=-2):        self.size = size        self.start = start        self.goal = goal        self.traps = traps if traps else []        self.step_penalty = step_penalty        self.goal_reward = goal_reward        self.trap_penalty = trap_penalty        self.out_of_bounds_penalty = out_of_bounds_penalty        self.num_states = size * size        self.num_actions = 4        self.agent_pos = start        # 动作: 0上 1下 2左 3右        self.action_moves = [(-1, 0), (1, 0), (0, -1), (0, 1)]    def reset(self):        """重置智能体到起点"""        self.agent_pos = self.start        return self.state_index(self.agent_pos)    def state_index(self, pos):        """将 (row, col) 转换为状态索引"""        return pos[0] * self.size + pos[1]    def index_to_pos(self, index):        """将状态索引转换回 (row, col)"""        return (index // self.size, index % self.size)    def is_terminal(self, state):        """判断是否为终止状态(终点或陷阱)"""        pos = self.index_to_pos(state)        return pos == self.goal or pos in self.traps    def step(self, action):        """执行动作,返回下一状态、奖励、是否终止"""        row, col = self.agent_pos        dr, dc = self.action_moves[action]        new_row, new_col = row + dr, col + dc        # 检查是否出界        if not (0 <= new_row < self.size and 0 <= new_col < self.size):            self.agent_pos = (row, col)  # 保持原地            return self.state_index(self.agent_pos), self.out_of_bounds_penalty, False        self.agent_pos = (new_row, new_col)        next_state = self.state_index(self.agent_pos)        # 判断奖励与终止        if self.agent_pos == self.goal:            return next_state, self.goal_reward, True        elif self.agent_pos in self.traps:            return next_state, self.trap_penalty, True        else:            return next_state, self.step_penalty, Falsedef get_optimal_path(env, q_table):    """根据学习到的 Q 表,从起点出发走最优路径"""    path = []    state = env.reset()    path.append(env.index_to_pos(state))    for _ in range(env.num_states):        if env.is_terminal(state):            break        action = np.argmax(q_table[state])        row, col = env.index_to_pos(state)        dr, dc = env.action_moves[action]        new_pos = (row + dr, col + dc)        # 防止走出边界        if not (0 <= new_pos[0] < env.size and 0 <= new_pos[1] < env.size):            break        state = env.state_index(new_pos)        path.append(new_pos)        if env.is_terminal(state):            break    return pathdef moving_average(data, window=20):    """计算滑动平均"""    if len(data) < window:        return np.array(data)    return np.convolve(data, np.ones(window)/window, mode='valid')def plot_results(env, q_table, rewards_history):    fig, axes = plt.subplots(2, 2, figsize=(13, 10))    # 计算状态价值 V(s) = max_a Q(s,a)    values = np.max(q_table, axis=1).reshape(env.size, env.size)    values[env.goal] = env.goal_reward    for trap in env.traps:        values[trap] = env.trap_penalty    # --- 1. 学习到的最优策略 ---    ax = axes[0, 0]    im = ax.imshow(values, cmap='coolwarm', origin='upper', alpha=0.85)    ax.set_title('学习到的最优策略(箭头方向)', fontsize=13)    for r in range(env.size):        for c in range(env.size):            state = env.state_index((r, c))            if (r, c) == env.goal:                ax.text(c, r, 'G', ha='center', va='center', fontsize=16,                        fontweight='bold', color='green')            elif (r, c) in env.traps:                ax.text(c, r, 'X', ha='center', va='center', fontsize=16,                        fontweight='bold', color='red')            else:                action = np.argmax(q_table[state])                dr, dc = env.action_moves[action]                ax.arrow(c, r, dc*0.4, dr*0.4, head_width=0.15,                         head_length=0.15, fc='black', ec='black',                         length_includes_head=True)    ax.set_xticks(np.arange(env.size))    ax.set_yticks(np.arange(env.size))    ax.set_xticklabels([str(i) for i in range(env.size)])    ax.set_yticklabels([str(i) for i in range(env.size)])    ax.set_xlabel('列')    ax.set_ylabel('行')    plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04)    # --- 2. 状态价值热力图 ---    ax = axes[0, 1]    im2 = ax.imshow(values, cmap='viridis', origin='upper')    ax.set_title('状态价值热力图(max Q)', fontsize=13)    for r in range(env.size):        for c in range(env.size):            ax.text(c, r, f'{values[r, c]:.1f}', ha='center', va='center',                    color='white', fontsize=9)    ax.set_xticks(np.arange(env.size))    ax.set_yticks(np.arange(env.size))    ax.set_xticklabels([str(i) for i in range(env.size)])    ax.set_yticklabels([str(i) for i in range(env.size)])    ax.set_xlabel('列')    ax.set_ylabel('行')    plt.colorbar(im2, ax=ax, fraction=0.046, pad=0.04)    # --- 3. 学习曲线 ---    ax = axes[1, 0]    ax.plot(rewards_history, alpha=0.3, color='steelblue', label='单轮奖励')    window = 20    if len(rewards_history) >= window:        smoothed = moving_average(rewards_history, window)        ax.plot(range(window-1, len(rewards_history)), smoothed,                color='darkorange', linewidth=2, label='滑动平均(20)')    ax.set_xlabel('训练轮数(Episode)')    ax.set_ylabel('总奖励')    ax.set_title('学习曲线', fontsize=13)    ax.legend()    ax.grid(alpha=0.3)    # --- 4. 最优路径 ---    ax = axes[1, 1]    ax.imshow(np.zeros((env.size, env.size)), cmap='Greys', origin='upper', alpha=0.1)    ax.set_title('智能体找到的最优路径', fontsize=13)    for trap in env.traps:        ax.text(trap[1], trap[0], 'X', ha='center', va='center', fontsize=16,                fontweight='bold', color='red')    ax.text(env.goal[1], env.goal[0], 'G', ha='center', va='center', fontsize=16,            fontweight='bold', color='green')    ax.text(env.start[1], env.start[0], 'S', ha='center', va='center', fontsize=16,            fontweight='bold', color='blue')    path = get_optimal_path(env, q_table)    if path:        rows = [p[0] for p in path]        cols = [p[1] for p in path]        ax.plot(cols, rows, 'o-', color='lime', linewidth=2.5, markersize=8,                markerfacecolor='yellow', markeredgecolor='black')    ax.set_xticks(np.arange(env.size))    ax.set_yticks(np.arange(env.size))    ax.set_xticklabels([str(i) for i in range(env.size)])    ax.set_yticklabels([str(i) for i in range(env.size)])    ax.set_xlabel('列')    ax.set_ylabel('行')    plt.tight_layout()    plt.savefig('reward_punishment_qlearning.png', dpi=150, bbox_inches='tight')    plt.show()if __name__ == '__main__':    # 创建环境    env = GridWorld(        size=5,        start=(0, 0),        goal=(4, 4),        traps=[(1, 3), (3, 1), (2, 2)],        step_penalty=-1,        goal_reward=10,        trap_penalty=-10,        out_of_bounds_penalty=-2    )    # 初始化 Q 表    q_table = np.zeros((env.num_states, env.num_actions))    # 超参数    epsilon = 1.0          # 初始探索率    epsilon_min = 0.01     # 最小探索率    epsilon_decay = 0.995  # 探索率衰减    alpha = 0.1            # 学习率    gamma = 0.9            # 折扣因子    episodes = 800         # 训练轮数    max_steps = 100        # 每轮最大步数    rewards_history = []    print("开始训练 Q-learning 智能体...")    for episode in range(episodes):        state = env.reset()        total_reward = 0        for step in range(max_steps):            # ε-贪心策略选择动作            if np.random.rand() < epsilon:                action = np.random.randint(env.num_actions)            else:                action = np.argmax(q_table[state])            # 执行动作            next_state, reward, done = env.step(action)            # Q-learning 核心更新            if done:                best_next = 0            else:                best_next = np.max(q_table[next_state])            q_table[state, action] += alpha * (reward + gamma * best_next - q_table[state, action])            state = next_state            total_reward += reward            if done:                break        rewards_history.append(total_reward)        epsilon = max(epsilon_min, epsilon * epsilon_decay)        if (episode + 1) % 100 == 0:            avg_reward = np.mean(rewards_history[-100:])            print(f"Episode {episode+1}/{episodes}, 平均奖励: {avg_reward:.2f}")    plot_results(env, q_table, rewards_history)

四、运行结果解读

1. 左上:学习到的最优策略

每个格子中的箭头表示智能体在该格子的最优动作方向。背景颜色越红代表状态价值越高,越蓝代表状态价值越低。可以看到智能体学会了在陷阱附近绕行。

2. 右上:状态价值热力图

数字表示每个状态的最大 Q 值。终点 G 的价值最高(+10),陷阱 X 的价值最低(-10),越靠近终点的格子价值越高。

3. 左下:学习曲线

横轴是训练轮数,纵轴是每轮总奖励。蓝色细线是每轮的实际奖励,橙色粗线是 20 轮滑动平均。可以看到训练初期智能体经常踩陷阱,奖励很低;随着训练进行,奖励逐渐上升并趋于稳定。

4. 右下:智能体找到的最优路径

绿色线条展示了智能体从起点 S 出发,避开所有陷阱 X,最终到达终点 G 的路径。

最新文章

随机文章