0. 出租车任务¶
先渲染一个初始状态,确认出租车、乘客和目的地分别在哪里;再看动作含义、TD 更新和训练后路线。
In [3]:
# 出租车调度:先看一个初始状态和六个可选动作。
taxi_env = gym.make("Taxi", render_mode="ansi")
taxi_env.action_space.seed(11)
action_labels = ["向南", "向北", "向东", "向西", "接乘客", "放下乘客"]
location_names = ["红色站点", "绿色站点", "黄色站点", "蓝色站点"]
preview_state, _ = taxi_env.reset(seed=42)
taxi_row, taxi_col, passenger_idx, dest_idx = taxi_env.unwrapped.decode(preview_state)
print(taxi_env.render())
display(pd.DataFrame([{
"状态编号": preview_state,
"出租车位置": f"第 {taxi_row} 行,第 {taxi_col} 列",
"乘客": "在车上" if passenger_idx == 4 else location_names[passenger_idx],
"目的地": location_names[dest_idx],
}]))
display(pd.DataFrame({
"动作编号": range(len(action_labels)),
"动作": action_labels,
"含义": ["向南移动", "向北移动", "向东移动", "向西移动", "接乘客", "放下乘客"],
}))
taxi_sites = {
"红色站点": (0, 0),
"绿色站点": (0, 4),
"黄色站点": (4, 0),
"蓝色站点": (4, 3),
}
site_colors = {
"红色站点": "#fecaca",
"绿色站点": "#bbf7d0",
"黄色站点": "#fef08a",
"蓝色站点": "#bfdbfe",
}
def draw_taxi_board(ax, row, col, passenger_idx, dest_idx, title, path=None):
for r in range(5):
for c in range(5):
ax.add_patch(plt.Rectangle((c - 0.5, r - 0.5), 1, 1, color="#f8fafc", ec="#cbd5e1"))
for name, (sr, sc) in taxi_sites.items():
ax.scatter([sc], [sr], s=520, color=site_colors[name], edgecolors="#334155", linewidths=1.2, zorder=2)
ax.text(sc, sr, name[0], ha="center", va="center", fontweight="bold", color="#0f172a", zorder=3)
if path:
ys = [item[0] for item in path]
xs = [item[1] for item in path]
ax.plot(xs, ys, color="#2563eb", linewidth=2.4, marker="o", markersize=5, zorder=4)
passenger_label = "车上" if passenger_idx == 4 else location_names[passenger_idx]
if passenger_idx != 4:
pr, pc = taxi_sites[passenger_label]
ax.scatter([pc], [pr], s=170, marker="o", color="#f97316", edgecolors="#7c2d12", linewidths=1.1, zorder=5)
dest_name = location_names[dest_idx]
dr, dc = taxi_sites[dest_name]
ax.scatter([dc], [dr], s=210, marker="*", color="#16a34a", edgecolors="#14532d", linewidths=1.0, zorder=5)
ax.scatter([col], [row], s=230, marker="s", color="#2563eb", edgecolors="#1e3a8a", linewidths=1.2, zorder=6)
ax.text(col, row, "车", ha="center", va="center", color="white", fontweight="bold", zorder=7)
ax.set_title(f"{title}\n乘客:{passenger_label};目的地:{dest_name}", loc="left", fontweight="bold")
ax.set_xlim(-0.5, 4.5)
ax.set_ylim(4.5, -0.5)
ax.set_xticks(range(5))
ax.set_yticks(range(5))
ax.grid(True, color="#e2e8f0", linewidth=0.7)
fig, ax = plt.subplots(figsize=(5.4, 5.2))
draw_taxi_board(ax, taxi_row, taxi_col, passenger_idx, dest_idx, "出租车初始状态")
plt.tight_layout()
plt.show()
# Q-learning:每一步用奖励和下一状态更新 Q(s,a)。
n_states_taxi = taxi_env.observation_space.n
n_actions_taxi = taxi_env.action_space.n
Q_taxi = np.zeros((n_states_taxi, n_actions_taxi))
alpha = 0.15
gamma = 0.95
epsilon_start = 1.0
epsilon_end = 0.05
episodes = 1600
rng = np.random.default_rng(11)
training_rows = []
td_samples = []
for episode in range(1, episodes + 1):
state, _ = taxi_env.reset(seed=episode)
total_reward = 0
steps = 0
epsilon = max(epsilon_end, epsilon_start * (0.995 ** episode))
terminated = truncated = False
while not (terminated or truncated) and steps < 200:
if rng.random() < epsilon:
action = taxi_env.action_space.sample()
else:
action = int(np.argmax(Q_taxi[state]))
next_state, reward, terminated, truncated, _ = taxi_env.step(action)
before_q = Q_taxi[state, action]
td_target = reward + gamma * np.max(Q_taxi[next_state]) * (not (terminated or truncated))
td_error = td_target - Q_taxi[state, action]
Q_taxi[state, action] += alpha * td_error
if len(td_samples) < 12:
td_samples.append({
"回合": episode,
"状态": state,
"动作": action_labels[action],
"奖励": reward,
"下一状态": next_state,
"TD 目标": td_target,
"TD 误差": td_error,
"更新前 Q": before_q,
"更新后 Q": Q_taxi[state, action],
})
state = next_state
total_reward += reward
steps += 1
if episode % 100 == 0:
training_rows.append({"episode": episode, "reward": total_reward, "steps": steps, "epsilon": epsilon})
taxi_trace = pd.DataFrame(training_rows)
td_trace = pd.DataFrame(td_samples).round(3)
display(td_trace)
display(taxi_trace.tail(8).rename(columns={
"episode": "回合",
"reward": "回合奖励",
"steps": "步数",
"epsilon": "探索率",
}).round(3))
+---------+ |R: | : :G| | : | : : | | : : : : | | | : | : | |Y| : |B: | +---------+
| 状态编号 | 出租车位置 | 乘客 | 目的地 | |
|---|---|---|---|---|
| 0 | 386 | 第 3 行,第 4 列 | 绿色站点 | 黄色站点 |
| 动作编号 | 动作 | 含义 | |
|---|---|---|---|
| 0 | 0 | 向南 | 向南移动 |
| 1 | 1 | 向北 | 向北移动 |
| 2 | 2 | 向东 | 向东移动 |
| 3 | 3 | 向西 | 向西移动 |
| 4 | 4 | 接乘客 | 接乘客 |
| 5 | 5 | 放下乘客 | 放下乘客 |
| 回合 | 状态 | 动作 | 奖励 | 下一状态 | TD 目标 | TD 误差 | 更新前 Q | 更新后 Q | |
|---|---|---|---|---|---|---|---|---|---|
| 0 | 1 | 252 | 向南 | -1 | 352 | -1.0 | -1.00 | 0.00 | -0.150 |
| 1 | 1 | 352 | 向南 | -1 | 452 | -1.0 | -1.00 | 0.00 | -0.150 |
| 2 | 1 | 452 | 接乘客 | -10 | 452 | -10.0 | -10.00 | 0.00 | -1.500 |
| 3 | 1 | 452 | 向东 | -1 | 452 | -1.0 | -1.00 | 0.00 | -0.150 |
| 4 | 1 | 452 | 向西 | -1 | 432 | -1.0 | -1.00 | 0.00 | -0.150 |
| 5 | 1 | 432 | 向西 | -1 | 432 | -1.0 | -1.00 | 0.00 | -0.150 |
| 6 | 1 | 432 | 接乘客 | -10 | 432 | -10.0 | -10.00 | 0.00 | -1.500 |
| 7 | 1 | 432 | 向南 | -1 | 432 | -1.0 | -1.00 | 0.00 | -0.150 |
| 8 | 1 | 432 | 向东 | -1 | 452 | -1.0 | -1.00 | 0.00 | -0.150 |
| 9 | 1 | 452 | 向南 | -1 | 452 | -1.0 | -1.00 | 0.00 | -0.150 |
| 10 | 1 | 452 | 向东 | -1 | 452 | -1.0 | -0.85 | -0.15 | -0.277 |
| 11 | 1 | 452 | 放下乘客 | -10 | 452 | -10.0 | -10.00 | 0.00 | -1.500 |
| 回合 | 回合奖励 | 步数 | 探索率 | |
|---|---|---|---|---|
| 8 | 900 | 4 | 8 | 0.05 |
| 9 | 1000 | 6 | 15 | 0.05 |
| 10 | 1100 | 6 | 15 | 0.05 |
| 11 | 1200 | 12 | 9 | 0.05 |
| 12 | 1300 | -3 | 15 | 0.05 |
| 13 | 1400 | 8 | 13 | 0.05 |
| 14 | 1500 | 7 | 14 | 0.05 |
| 15 | 1600 | 9 | 12 | 0.05 |
In [4]:
# 训练后执行一条路线:每一步都按当前 Q 表选择价值最高的动作。
rollout_state, _ = taxi_env.reset(seed=42)
initial_row, initial_col, initial_passenger_idx, initial_dest_idx = taxi_env.unwrapped.decode(rollout_state)
rollout_rows = []
rollout_coords = [(initial_row, initial_col)]
rollout_reward = 0
for step in range(1, 31):
row, col, passenger_idx, dest_idx = taxi_env.unwrapped.decode(rollout_state)
q_values = Q_taxi[rollout_state]
action = int(np.argmax(q_values))
next_state, reward, terminated, truncated, _ = taxi_env.step(action)
next_row, next_col, next_passenger_idx, next_dest_idx = taxi_env.unwrapped.decode(next_state)
rollout_rows.append({
"步数": step,
"出租车位置": f"({row},{col})",
"乘客": "在车上" if passenger_idx == 4 else location_names[passenger_idx],
"目的地": location_names[dest_idx],
"选择动作": action_labels[action],
"动作价值": q_values[action],
"下一位置": f"({next_row},{next_col})",
"下一乘客状态": "在车上" if next_passenger_idx == 4 else location_names[next_passenger_idx],
"奖励": reward,
})
rollout_coords.append((next_row, next_col))
rollout_reward += reward
rollout_state = next_state
if terminated or truncated:
break
taxi_rollout_df = pd.DataFrame(rollout_rows)
display(taxi_rollout_df.round(3))
print("路线总奖励:", rollout_reward)
print("是否完成:", bool(terminated))
fig, ax = plt.subplots(figsize=(5.6, 5.4))
last_row, last_col, last_passenger_idx, last_dest_idx = taxi_env.unwrapped.decode(rollout_state)
draw_taxi_board(
ax,
last_row,
last_col,
last_passenger_idx,
last_dest_idx,
"训练后执行路线",
path=rollout_coords,
)
plt.tight_layout()
plt.show()
| 步数 | 出租车位置 | 乘客 | 目的地 | 选择动作 | 动作价值 | 下一位置 | 下一乘客状态 | 奖励 | |
|---|---|---|---|---|---|---|---|---|---|
| 0 | 1 | (3,4) | 绿色站点 | 黄色站点 | 向北 | -0.023 | (2,4) | 绿色站点 | -1 |
| 1 | 2 | (2,4) | 绿色站点 | 黄色站点 | 向北 | 2.366 | (1,4) | 绿色站点 | -1 |
| 2 | 3 | (1,4) | 绿色站点 | 黄色站点 | 向北 | 3.948 | (0,4) | 绿色站点 | -1 |
| 3 | 4 | (0,4) | 绿色站点 | 黄色站点 | 接乘客 | 5.210 | (0,4) | 在车上 | -1 |
| 4 | 5 | (0,4) | 在车上 | 黄色站点 | 向西 | 6.537 | (0,3) | 在车上 | -1 |
| 5 | 6 | (0,3) | 在车上 | 黄色站点 | 向南 | 7.933 | (1,3) | 在车上 | -1 |
| 6 | 7 | (1,3) | 在车上 | 黄色站点 | 向南 | 9.404 | (2,3) | 在车上 | -1 |
| 7 | 8 | (2,3) | 在车上 | 黄色站点 | 向西 | 10.951 | (2,2) | 在车上 | -1 |
| 8 | 9 | (2,2) | 在车上 | 黄色站点 | 向西 | 12.580 | (2,1) | 在车上 | -1 |
| 9 | 10 | (2,1) | 在车上 | 黄色站点 | 向西 | 14.295 | (2,0) | 在车上 | -1 |
| 10 | 11 | (2,0) | 在车上 | 黄色站点 | 向南 | 16.100 | (3,0) | 在车上 | -1 |
| 11 | 12 | (3,0) | 在车上 | 黄色站点 | 向南 | 18.000 | (4,0) | 在车上 | -1 |
| 12 | 13 | (4,0) | 在车上 | 黄色站点 | 放下乘客 | 20.000 | (4,0) | 黄色站点 | 20 |
路线总奖励: 8 是否完成: True
In [5]:
# 绘制训练曲线和一个起始状态的动作价值。
start_state, _ = taxi_env.reset(seed=42)
fig, axes = plt.subplots(1, 2, figsize=(10.0, 4.4))
axes[0].plot(taxi_trace["episode"], taxi_trace["reward"], marker="o", color="#2563eb", linewidth=2.0)
axes[0].set_title("出租车调度 Q-learning 回报", loc="left", fontweight="bold")
axes[0].set_xlabel("训练回合")
axes[0].set_ylabel("回合奖励")
axes[0].grid(True, color="#e2e8f0", linewidth=0.8)
axes[1].bar(action_labels, Q_taxi[start_state], color="#f97316")
axes[1].set_title(f"状态 {start_state} 的 Q(s,a)", loc="left", fontweight="bold")
axes[1].tick_params(axis="x", rotation=30)
axes[1].grid(True, axis="y", color="#e2e8f0", linewidth=0.8)
plt.tight_layout()
plt.show()
print(taxi_env.render())
print("当前渲染状态的贪心动作:", action_labels[int(np.argmax(Q_taxi[start_state]))])
taxi_env.close()
+---------+ |R: | : :G| | : | : : | | : : : : | | | : | : | |Y| : |B: | +---------+ 当前渲染状态的贪心动作: 向北