首页
/ ML-For-Beginners:基于 OpenAI Gym 与 Q-Learning 训练 CartPole 平衡的完整实战指南

ML-For-Beginners:基于 OpenAI Gym 与 Q-Learning 训练 CartPole 平衡的完整实战指南

2026-09-06 09:14:24作者:凌朦慧Richard

本篇指南基于 ML-For-Beginners 课程《8-Reinforcement / 2-Gym》(CartPole Skating)编写,系统讲解如何把上一课的 Q-Learning 算法从离散状态环境迁移到连续状态环境:使用 OpenAI Gym 模拟 CartPole 物理系统,完成状态离散化、基于字典的 Q-Table 构建、训练循环与超参数调优的全流程。读完后,你将能够独立编写并在 Gym 中训练一个保持杆子平衡的智能体,掌握处理连续观测空间这一强化学习核心难点的通用方法。

1. 从离散状态到连续状态:为什么是 CartPole

上一课(见 Q-Learning 入门)中解决的问题看似是"玩具问题",但实际上它代表了相当多真实场景——包括国际象棋和围棋:棋盘有给定规则与离散状态。本课则进入更贴近现实的一类问题:状态由一个或多个实数给出,即连续状态(continuous state)

本课的问题设定是:

问题:如果 Peter 想从狼那里逃脱,他就需要跑得更快。我们将看到 Peter 如何用 Q-Learning 学会滑冰——尤其是学会保持平衡。

为了简化"平衡"这个目标,课程采用了著名的 CartPole(小车-单摆) 问题:在 cartpole 世界中,有一个可以左右移动的水平滑块(小车),目标是在滑块上方平衡一根竖直的杆。

CartPole 小车-单摆系统示意

前提说明:本课使用 OpenAI Gym 库来模拟不同的 environment。Gym 是一个专为训练强化学习算法而维护的模拟环境集合,从 cartpole 这样的经典控制问题到 Atari 游戏都涵盖在内。上一课中游戏规则与状态由我们自己定义的 Board 类给出,而这里则使用 Gym 提供的模拟环境来模拟平衡杆背后的真实物理。你可以在本地(例如 Visual Studio Code)运行本课代码,此时模拟窗口会在新窗口中弹出;如果在 Binder、Google Colab 等在线环境中运行,需要对渲染代码做少量调整。

2. 安装 Gym 与初始化 CartPole 环境

2.1 安装并导入依赖

import sys
!{sys.executable} -m pip install gym

import gym
import matplotlib.pyplot as plt
import numpy as np
import random

完整可运行版本见仓库中的 notebook.ipynb,代码块编号与下文 code block 1~13 一一对应。

2.2 观察空间与动作空间

使用 cartpole 平衡问题前,需要先初始化对应的环境。Gym 中每个环境都关联两个核心抽象:

  • Observation space(观察空间):定义从环境接收到的信息结构。对 cartpole 而言,会收到杆的角度、速度等若干数值;
  • Action space(动作空间):定义可执行的动作。本课的动作空间是离散的,仅包含两个动作:left(向左推)right(向右推)

(code block 2) 初始化环境并查看两个空间:

env = gym.make("CartPole-v1")
print(env.action_space)
print(env.observation_space)
print(env.action_space.sample())

2.3 随机策略跑 100 步

为了直观感受环境行为,先跑一段 100 步的短模拟。每一步都从 action_space 中随机采样一个动作——由于策略完全随机,小车几乎立刻失衡:

✅ 建议在本地 Python 环境中运行以下代码(在线 notebook 需要额外配置渲染)。

(code block 3)

env.reset()

for i in range(100):
   env.render()
   env.step(env.action_space.sample())
env.close()

2.4 step() 的返回值:观察、奖励与终止标志

模拟过程中,我们需要获取观察值来决定如何行动。step 函数返回四个值:当前观察 obs、奖励 rew、以及指示模拟是否应该继续的终止标志 done(外加 info 字典):

(code block 4)

env.reset()

done = False
while not done:
   env.render()
   obs, rew, done, info = env.step(env.action_space.sample())
   print(f"{obs} -> {rew}")
env.close()

输出形如:

[ 0.03403272 -0.24301182  0.02669811  0.2895829 ] -> 1.0
[ 0.02917248 -0.04828055  0.03248977  0.00543839] -> 1.0
[ 0.02820687  0.14636075  0.03259854 -0.27681916] -> 1.0
[ 0.03113408  0.34100283  0.02706215 -0.55904489] -> 1.0
[ 0.03795414  0.53573468  0.01588125 -0.84308041] -> 1.0
...
[ 0.17299878  0.15868546 -0.20754175 -0.55975453] -> 1.0
[ 0.17617249  0.35602306 -0.21873684 -0.90998894] -> 1.0

每个模拟步返回的观察向量包含 4 个值:

序号 含义
0 小车位置(position of cart)
1 小车速度(velocity of cart)
2 杆的角度(angle of pole)
3 杆的旋转速率(rotation rate of pole)

可以查看各观察值的取值范围(code block 5):

print(env.observation_space.low)
print(env.observation_space.high)

仓库 notebook.ipynb 中保存的实际运行输出为:

[-4.8000002e+00 -3.4028235e+38 -4.1887903e-01 -3.4028235e+38]
[ 4.8000002e+00  3.4028235e+38  4.1887903e-01  3.4028235e+38]

这印证了一个关键事实:4 个观察量中有 2 个(小车速度、杆的旋转速率)没有实际上下界(输出为浮点极值 ±3.4e38),而小车位置限定在约 ±4.8、杆角度限定在约 ±0.419 弧度。这个"半有界"的特性正是后文状态离散化方案的出发点。

另外可以注意到:每个模拟步的奖励值恒为 1.0,因为目标是尽可能久地存活——即让杆保持在接近竖直的状态尽可能长的时间。

提示:CartPole 模拟的官方判定标准是——在连续 100 次试跑中平均获得 195 分的累计奖励,即视为"解决问题"。

3. 状态离散化:把连续观测映射到有限状态集

Q-Learning 需要构建一张 Q-Table,为每个状态规定该做什么。这就要求状态是离散的——更准确地说,状态应只包含有限个离散取值。因此,必须以某种方式**离散化(discretize)**观察向量,将其映射到有限状态集合上。常用方法有:

  • 分箱(Divide into bins):如果知道某个值的取值区间,就可以把区间切成若干 bin,再把值替换为它所属的 bin 编号。这一步可用 numpy 的 np.digitize 方法完成。采用这种方式时,状态空间大小是精确已知的,因为它完全取决于所选的 bin 数。
  • 线性插值到有限区间后取整:用线性变换把值带进某个有限区间(比如 -20 到 20),再通过取整把数字转换为整数。这种方式对状态空间大小的控制稍弱,尤其是当不清楚输入值的确切范围时。例如本课 4 个观察量中有 2 个没有上下界,理论上可能导致状态数无穷。

本例采用第二种方式。后文会看到:尽管上下界未定义,这些值实际上很少落在某个有限区间之外,因此取极端值的 state 极为罕见。

(code block 6) 将模型观察转换为 4 个整数组成的元组的函数:

def discretize(x):
    return tuple((x/np.array([0.25, 0.25, 0.01, 0.1])).astype(np.int))

这里 4 个分母分别是各观察量的"分辨率步长":位置与速度每 0.25 一档、角度每 0.01 弧度一档、旋转速率每 0.1 一档。

(code block 7) 同时探索基于 bin 的另一种离散化方式:

def create_bins(i,num):
    return np.arange(num+1)*(i[1]-i[0])/num+i[0]

print("Sample bins for interval (-5,5) with 10 bins\n",create_bins((-5,5),10))

ints = [(-5,5),(-2,2),(-0.5,0.5),(-2,2)] # intervals of values for each parameter
nbins = [20,20,10,10] # number of bins for each parameter
bins = [create_bins(ints[i],nbins[i]) for i in range(4)]

def discretize_bins(x):
    return tuple(np.digitize(x[i],bins[i]) for i in range(4))

仓库 notebook 中保存的运行结果显示,对区间 (-5,5) 取 10 个 bin 时,create_bins 生成的边界为 [-5. -4. -3. -2. -1. 0. 1. 2. 3. 4. 5.],即 11 个分界点、10 个等宽区间。

(code block 8) 跑一段短模拟,观察离散后的环境值。可以随意切换 discretizediscretize_bins 对比两者差异:

env.reset()

done = False
while not done:
   #env.render()
   obs, rew, done, info = env.step(env.action_space.sample())
   #print(discretize_bins(obs))
   print(discretize(obs))
env.close()

两种离散化方式的语义差异值得留意:

  • discretize_bins 返回从 0 开始的 bin 编号,因此输入变量在 0 附近的值会返回区间中间位置的编号(如 10);
  • discretize 不关心输出值范围、允许负数,状态值没有偏移,0 仍对应 0。

提示:想看到环境动画就把 env.render() 一行取消注释;否则可以在后台无渲染运行,速度更快。后续 Q-Learning 训练过程就使用这种"隐身"执行。

4. Q-Table 结构:为什么这里用字典而不是张量

上一课中状态是 0 到 8 的简单数字对,因此用形状 8x8x2 的 numpy 张量表示 Q-Table 很方便。如果采用 bin 离散化,状态向量大小也是已知的,同样可以用张量,形状为 20x20x10x10x2——其中最后的 2 是动作空间维度,前面各维对应观察空间各参数所选的 bin 数(与 code block 7 中 nbins = [20,20,10,10] 一致)。

但有时观察空间的精确维度是未知的。对 discretize 函数而言,部分原始值无界,无法确信状态始终落在某个限制范围内。因此本课采用另一种更稳健的方式:用 Python 字典表示 Q-Table,以 (state, action) 元组为键,Q 值为值:

(code block 9)

Q = {}
actions = (0,1)

def qvalues(state):
    return [Q.get((state,a),0) for a in actions]

qvalues() 函数返回给定状态下所有可能动作对应的 Q-Table 值列表;若条目尚不存在于 Q-Table 中,则默认返回 0。这种惰性建表的方式天然适配无界状态:只在实际访问到的状态-动作对上分配条目。

5. 开始 Q-Learning:训练循环与超参数

现在可以教 Peter 保持平衡了。

5.1 超参数设置

(code block 10)

# hyperparameters
alpha = 0.3
gamma = 0.9
epsilon = 0.90

三个超参数的作用:

  • alpha(学习率 learning rate):定义每一步对 Q-Table 当前值调整的力度。上一课从 1 开始并在训练中逐渐调低 alpha;本例为简化起见保持恒定,后续可自行实验调整 alpha 的变化策略;
  • gamma(折扣因子 discount factor):表示相对当前奖励,应在多大程度上优先未来奖励;
  • epsilon(探索/利用因子 exploration/exploitation factor):决定算法偏向探索还是利用。在本算法中,epsilon 比例的步骤按 Q-Table 值选择下一个动作,剩余比例执行随机动作,从而得以探索从未见过的搜索空间区域。

提示:就平衡问题而言——随机动作(exploration)相当于朝错误方向随机"打一拳",而杆需要从这些"错误"中重新学会恢复平衡。这正是探索的价值所在。

5.2 对上一课算法的两处改进

  • 计算平均累计奖励:对多次模拟做平均。每 5000 次迭代打印一次进度,并在该时间窗口内对累计奖励取平均。如果平均超过 195 分,则意味着以高于官方标准的质量解决了问题。
  • 跟踪最大平均结果 Qmax:记录最佳平均累计结果,并保存对应的 Q-Table(Qbest)。训练中可以观察到平均累计奖励偶尔开始下滑——此时希望保留训练中观察到的最佳模型对应的 Q-Table 值。

5.3 完整训练代码

(code block 11) 所有模拟的累计奖励都收集进 rewards 向量,供后续绘图使用:

def probs(v,eps=1e-4):
    v = v-v.min()+eps
    v = v/v.sum()
    return v

Qmax = 0
cum_rewards = []
rewards = []
for epoch in range(100000):
    obs = env.reset()
    done = False
    cum_reward=0
    # == do the simulation ==
    while not done:
        s = discretize(obs)
        if random.random()<epsilon:
            # exploitation - chose the action according to Q-Table probabilities
            v = probs(np.array(qvalues(s)))
            a = random.choices(actions,weights=v)[0]
        else:
            # exploration - randomly chose the action
            a = np.random.randint(env.action_space.n)

        obs, rew, done, info = env.step(a)
        cum_reward+=rew
        ns = discretize(obs)
        Q[(s,a)] = (1 - alpha) * Q.get((s,a),0) + alpha * (rew + gamma * max(qvalues(ns)))
    cum_rewards.append(cum_reward)
    rewards.append(cum_reward)
    # == Periodically print results and calculate average reward ==
    if epoch%5000==0:
        print(f"{epoch}: {np.average(cum_rewards)}, alpha={alpha}, epsilon={epsilon}")
        if np.average(cum_rewards) > Qmax:
            Qmax = np.average(cum_rewards)
            Qbest = Q
        cum_rewards=[]

代码逐段解读:

  1. probs(v, eps=1e-4) 把 Q 值向量做 min-max 平移归一化(加 eps 防止全 0 时除以零),转成一个合法的概率分布,供 random.choices 按 Q-Table 概率采样动作——这就是"按 Q 值概率 exploitation"的实现;
  2. 内层 while not done 是标准 Gym 回合循环:env.reset() 开始新回合,env.step(a) 推进物理模拟并获得奖励;
  3. 关键的一行是 Q-Learning 更新公式:Q[(s,a)] = (1 - alpha) * Q.get((s,a),0) + alpha * (rew + gamma * max(qvalues(ns))),即对目标值 rew + gamma * max Q(s', a') 做指数平滑。注意这里取 max 而非期望——这正是 Q-Learning(而非 SARSA)的本质;
  4. 外层每 5000 个回合统计一次滑动窗口内的平均累计奖励,刷新 Qmax 并把当时的 Q 存为 Qbest(注意这是引用保存,后文挑战 Task 3 会讨论其局限性)。

训练结果通常可以观察到两点:

  • 接近目标:可能已经非常接近"连续 100+ 次模拟平均 195 分"的目标,甚至已经达成。即使数字略低也无法直接下结论——因为这里对 5000 次取平均,而官方标准只要求 100 次;
  • 奖励开始下滑:有时累计奖励会掉下去,意味着新写入的值可能"破坏"了 Q-Table 中已经学好的值。

把训练过程画出来,这个现象会更直观。

6. 绘制训练进程:原始曲线与滑动平均

训练过程中,每个迭代的累计奖励值被收集进 rewards 向量。直接按迭代序号绘制:

plt.plot(rewards)

原始训练进程曲线:每次模拟的累计奖励

由于随机训练过程的固有特性,每次训练回合(episode)长度差异很大,这条原始曲线几乎看不出趋势。让图表变得可读的办法是计算一段训练序列(比如 100 个)的滑动平均(running average),用 np.convolve 可以很方便地实现:

(code block 12)

def running_average(x,window):
    return np.convolve(x,np.ones(window)/window,mode='valid')

plt.plot(running_average(rewards,100))

滑动平均后的训练进程曲线,趋势清晰可见

平滑后可以明显看出:奖励先震荡上升,随后出现回落——这正是 5.2 节所说"奖励开始下滑"的可视化证据,也是引入 Qbest 快照机制的动机。

7. 超参数调优:让训练更稳定

为了让学习更稳定,训练中动态调整部分超参数是合理的做法:

  • 学习率 alpha:可以从接近 1 的值开始,然后持续衰减。随着训练进行,Q-Table 中逐渐积累起好的值,此时应该只做小幅调整,而不是用新值完全覆盖旧值;
  • 增大 epsilon:为了少探索、多利用,可以慢慢增大 epsilon——从较低的值开始,逐步升到接近 1。

课程给出了两个动手任务:

Task 1:调整超参数取值,看能否达到更高的累计奖励——能否拿到 195 以上?

Task 2:要"正式"解决问题,需要在连续 100 次运行中获得 195 的平均奖励。请在训练中实测这一指标,确认问题已被正式解决!

8. 运行训练好的模型:按 Q-Table 概率分布执行

真正观察训练好的模型行为会很有趣。运行模拟时采用与训练相同的动作选择策略——按 Q-Table 概率分布采样(code block 13):

obs = env.reset()
done = False
while not done:
   s = discretize(obs)
   env.render()
   v = probs(np.array(qvalues(s)))
   a = random.choices(actions,weights=v)[0]
   obs,_,done,_ = env.step(a)
env.close()

训练成功后,你会看到小车平稳往返、杆长时间保持竖直的平衡动画。

9. 进阶挑战

课程最后给出两个挑战性任务:

Task 3:这里使用的是 Q-Table 的最终副本,它未必是最佳副本。还记得我们把最优表现的 Q-Table 存进了 Qbest 变量吗?请把 Qbest 复制回 Q,用最优 Q-Table 重复上面的运行示例,看看有没有差别。

Task 4:上文每一步并非选择最优动作,而是按对应概率分布采样。是否应该总是选择 Q-Table 值最高的动作?可以用 np.argmax 找到最高 Q-Table 值对应的动作编号。请实现这一贪心策略,看它能否改善平衡效果。

此外,本课的正式作业是 Train a Mountain Car:把 Q-Learning 算法迁移到 Gym 的 MountainCar 环境——该环境中所有环境共享 reset/step/render 同一套 API 以及动作空间、观察空间抽象,因此只需替换环境、修改状态离散化函数,用最小代码改动让现有算法在新环境中训练收敛。

10. 小结

到这里,我们已经学会了如何仅通过一个定义"游戏期望状态"的奖励函数训练智能体,并给予其智能探索搜索空间的机会,从而得到好的结果。本课成功地把 Q-Learning 算法应用到了离散与连续两种状态环境(配合离散动作):状态连续、动作离散的场景,用"离散化 + 字典 Q-Table"即可解决。

值得继续研究的方向是:当动作空间也是连续的、观察空间也复杂得多(例如 Atari 游戏屏幕图像)时,往往需要神经网络等更强的机器学习技术才能达到良好效果。这些更进阶的主题是更高级 AI 课程的内容。

本课关键文件索引

文件 说明
2-Gym 讲义(英文原文) 本指南对应的英文原文
本课 notebook 与 code block 1-13 对应的完整代码,含保存的运行输出
作业:训练 Mountain Car 跨环境迁移 Q-Learning 的实战作业
Q-Learning 入门课 前置课程:离散状态下的 Q-Learning 基础
登录后查看全文
热门项目推荐
相关项目推荐