通俗易懂讲透 SARSA:强化学习 On-Policy 经典算法
·
通俗易懂讲透 SARSA:强化学习 On-Policy 经典算法
SARSA 是基于策略(On-Policy)的时序差分强化学习算法,核心是边执行策略、边学习策略,学得稳、风险低,非常适合动态与安全敏感场景。
一、SARSA 到底是什么?
一句话定位:
SARSA = 跟着自己当前的步子学习,一步一步稳着来的强化学习算法。
名字来源(很好记):
- S:当前状态
sₜ - A:当前动作
aₜ - R:奖励
rₜ - S:下一状态
sₜ₊₁ - A:下一动作
aₜ₊₁
这 5 个量串起来,就是 SARSA 一次更新的全部依据。
二、核心公式(掰开揉碎讲)
更新公式
Q(sₜ,aₜ) ← Q(sₜ,aₜ) + α [ rₜ + γ·Q(sₜ₊₁,aₜ₊₁) − Q(sₜ,aₜ) ]
每个符号是什么意思
Q(s,a):在状态 s 做动作 a 的长期收益打分α:学习率(0~1),控制每次改多少rₜ:执行动作后立刻拿到的奖励γ:折扣因子(0~1),越接近 1 越看重未来sₜ,aₜ:现在的状态和动作sₜ₊₁,aₜ₊₁:下一步的状态和真实会执行的动作
一句话理解更新逻辑
用“下一步真实要走的路”来修正“这一步的判断”,不跳步、不空想最优。
三、SARSA 运行流程(超清晰)
- 初始化 Q 表
所有状态-动作对的 Q 值初始化为 0。 - 每一轮训练(Episode)
- 回到起点,得到初始状态
s - 用 ε-贪心选第一个动作 a
- 执行动作 → 得到奖励 r、新状态 s’
- 在 s’ 再用 ε-贪心选下一个动作 a’
- 用 SARSA 公式更新 Q
- 把 s→s’、a→a’,继续走
- 回到起点,得到初始状态
- 结束一轮
到达终点/撞墙就重置,重复训练直到收敛。
四、SARSA vs Q-Learning(最关键区别)
| 对比项 | SARSA | Q-Learning |
|---|---|---|
| 策略类型 | On-Policy(同策略) | Off-Policy(异策略) |
| 学习依据 | 自己实际会走的下一步动作 | 直接用最优 Q 值(空想最优) |
| 性格 | 稳妥派、保守、安全 | 激进派、追求全局最优 |
| 稳定性 | 高,适合动态环境 | 容易震荡 |
| 风险 | 低,避开危险动作 | 可能铤而走险 |
通俗比喻
- SARSA:自己开车,边开边学,不冒险,稳稳到达。
- Q-Learning:看着攻略开车,总想抄近道,偶尔会翻车。
五、探索与利用:ε-贪心策略
SARSA 和 Q-Learning 都用,但目的不一样:
- 以概率
ε随机走(探索) - 以概率
1−ε选 Q 最大的动作(利用)
训练技巧:
- 刚开始 ε 大(多探索)
- 后期 ε 慢慢减小(多利用)
六、实战代码:5×5 网格世界(可直接跑)
import numpy as np
import matplotlib.pyplot as plt
import torch
import random
# 网格世界参数
GRID_SIZE = 5
ACTIONS = ['上','下','左','右']
ACTION_MAP = {0:(-1,0), 1:(1,0), 2:(0,-1), 3:(0,1)}
# 环境类
class GridWorld:
def __init__(self):
self.start = (0,0)
self.goal = (4,4)
self.obstacles = [(2,2),(3,3)]
self.state = self.start
def reset(self):
self.state = self.start
return self.state
def step(self, action):
r, c = self.state
dr, dc = ACTION_MAP[action]
nr, nc = r+dr, c+dc
# 越界/撞墙 不移动
if nr<0 or nr>=GRID_SIZE or nc<0 or nc>=GRID_SIZE or (nr,nc) in self.obstacles:
nr, nc = r, c
self.state = (nr, nc)
reward = 1 if self.state==self.goal else -0.1
done = self.state==self.goal
return self.state, reward, done
# SARSA 智能体
class SARSA_Agent:
def __init__(self, lr=0.1, gamma=0.9, epsilon=0.1):
self.lr = lr
self.gamma = gamma
self.epsilon = epsilon
self.q_table = torch.zeros(GRID_SIZE, GRID_SIZE, 4)
def choose_action(self, state):
if random.random() < self.epsilon:
return random.randint(0,3)
return torch.argmax(self.q_table[state[0], state[1]]).item()
def update(self, s, a, r, s_next, a_next):
q_old = self.q_table[s[0], s[1], a]
q_target = r + self.gamma * self.q_table[s_next[0], s_next[1], a_next]
self.q_table[s[0], s[1], a] += self.lr * (q_target - q_old)
# 训练函数
def train(episodes=1000):
env = GridWorld()
agent = SARSA_Agent()
reward_list = []
q_mean_list = []
for epi in range(episodes):
s = env.reset()
a = agent.choose_action(s)
total_r = 0
qs = []
while True:
s_next, r, done = env.step(a)
a_next = agent.choose_action(s_next)
agent.update(s, a, r, s_next, a_next)
total_r += r
qs.append(agent.q_table[s[0], s[1]].mean().item())
if done: break
s, a = s_next, a_next
reward_list.append(total_r)
q_mean_list.append(np.mean(qs))
return agent, reward_list, q_mean_list
# 开始训练
agent, rewards, q_means = train(1000)
# 画图:奖励曲线 + 平均Q值 + 策略图 + Q热力图
plt.rcParams['font.sans-serif'] = ['SimHei']
fig, axs = plt.subplots(2,2,figsize=(12,10))
fig.suptitle('SARSA 训练结果可视化', fontsize=16)
axs[0,0].plot(rewards, color='r')
axs[0,0].set_title('累计奖励')
axs[0,0].grid(True)
axs[0,1].plot(q_means, color='b')
axs[0,1].set_title('平均Q值')
axs[0,1].grid(True)
# 策略图
policy = np.zeros((GRID_SIZE,GRID_SIZE), dtype=int)
for i in range(GRID_SIZE):
for j in range(GRID_SIZE):
if (i,j) in [(4,4),(2,2),(3,3)]: continue
policy[i,j] = torch.argmax(agent.q_table[i,j]).item()
axs[1,0].imshow(policy, cmap='coolwarm')
for i in range(GRID_SIZE):
for j in range(GRID_SIZE):
if (i,j) not in [(4,4),(2,2),(3,3)]:
axs[1,0].text(j,i, ACTIONS[policy[i,j]], ha='center',va='center')
axs[1,0].set_title('最优策略')
# Q值热力图
heat = np.zeros((GRID_SIZE,GRID_SIZE))
for i in range(GRID_SIZE):
for j in range(GRID_SIZE):
heat[i,j] = agent.q_table[i,j].max().item()
im = axs[1,1].imshow(heat, cmap='jet')
fig.colorbar(im, ax=axs[1,1])
axs[1,1].set_title('Q值热力图')
plt.tight_layout()
plt.show()
代码说明
GridWorld:5×5 迷宫,有起点、终点、障碍物SARSA_Agent:实现选动作、Q 更新- 输出 4 张图:奖励曲线、Q 值趋势、策略图、Q 热力图
七、SARSA 优点与缺点
优点
- 稳定安全:On-Policy 学习,不冒进
- 动态环境友好:环境变了也能稳健适应
- 风险低:不会像 Q-Learning 强行走最优而踩坑
- 实现简单:逻辑清晰,易调试
缺点
- 收敛偏慢:稳是稳,但学起来慢
- 易陷局部最优:太保守,不敢大胆探索
- 依赖 ε:调不好就学习失效
- 大状态空间不行:Q 表会爆炸(要用深度 SARSA)
八、适用场景(读研/做项目必看)
SARSA 特别适合:
- 机器人路径规划、避障
- 自动驾驶决策
- 股票/交易策略(风险敏感)
- 动态网络负载均衡
- 游戏 AI(需要稳定策略)
不适合:追求极致最优、环境静止、不怕风险的场景(优先 Q-Learning)。
九、总结(3 句背会)
- SARSA 是 On-Policy,用真实下一步动作更新
- 更新必须凑齐 S-A-R-S-A
- 稳、安全、慢,适合动态与风险场景
更多推荐

所有评论(0)