机器学习————Q-Learning/SARSA算法
Q-Learning/SARSA算法是强化学习方向的算法,核心特点是智能体通过与环境交互获取奖励信号,学习最优决策策略,目标是最大化长期累积奖励。
一、核心概念
1.前置条件
- 智能体(Agent):执行决策的主体,比如游戏中的玩家。
- 环境(Environment):智能体所处的外部场景,比如游戏地图。
- 状态(State, S):智能体在环境中的当前位置或情况,比如玩家坐标。
- 动作(Action, A):智能体可执行的操作,比如上下左右移动。
- 奖励(Reward, R):环境对智能体动作的反馈,比如走到目标得 + 10,掉悬崖得 - 100。
- 策略(Policy, π):智能体选择动作的规则,比如 ε- 贪心:大概率选当前最优动作,小概率随机探索。
- Q 值(Q (S,A)):“状态 S 下执行动作 A” 的长期累计奖励期望,是算法优化的核心目标。
- 学习率(α):控制每次更新的幅度(0<α≤1),α 越大更新越激进,越小越保守。
- 折扣因子(γ):衡量未来奖励的重要性(0≤γ≤1),γ=0 只关注即时奖励,γ=1 完全重视未来奖励。
2. Q-Learning 核心概念
- 离策略(Off-Policy):更新 Q 值时,不依赖当前正在执行的策略,而是用 “下一个状态的最优 Q 值” 来更新,即 “学最优,做探索”。
- 核心特点:贪心式更新,只关注 “下一步的最好可能”,不考虑自己实际会走哪条路,更激进、收敛速度可能更快,但在危险环境中容易因贪心探索导致失误。
3. SARSA 核心概念
- 在策略(On-Policy):更新 Q 值时,使用 “当前策略下实际会执行的下一个动作” 的 Q 值,即 “学自己,做自己”。
- 核心特点:更新与实际执行的动作绑定,更保守、更注重 “安全”,在危险环境中,能避免盲目探索导致的高惩罚。
二、数学概念
1.符号定义
,第 t 步时,状态 S 下执行动作 A 的 Q 值估计;
,学习率(0 < α ≤ 1);
,折扣因子(0 ≤ γ ≤ 1);
,执行动作 A 后获得的即时奖励;
,执行动作 A 后进入的新状态;
,在新状态
下选择的动作(SARSA 专用);
2. 时序差分(TD)误差核心框架
两种算法的本质都是用 “TD 误差” 更新 Q 值,TD 误差的核心思想是:
用 “当前估计的 Q 值” 与 “实际奖励 + 未来 Q 值期望” 的差值,修正当前 Q 值。
基础 TD 误差公式:
;
Q 值更新的通用公式:
;
3. SARSA 更新公式
SARSA 的取当前策略下实际选择的下一个动作的 Q 值:
;
关键: 表示不考虑实际选什么动作,只取新状态下最优动作的 Q 值。
4. SARSA 更新公式
SARSA 的取当前策略下实际选择的下一个动作的 Q 值:
;
关键:是智能体根据当前策略(如 ε- 贪心)在
下实际选的动作,而非最优动作。
三、代码解释
模块一:导入核心库
import gym # 强化学习环境库
import numpy as np # 数值计算
import matplotlib.pyplot as plt # 可视化
模块二:定义通用工具函数
def epsilon_greedy_policy(Q_table, state, epsilon):
"""
ε-贪心策略:平衡探索与利用
参数:
Q_table: Q值表(二维数组,行=状态,列=动作)
state: 当前状态
epsilon: 探索概率(0<ε<1)
返回:
选择的动作
"""
# 生成0-1随机数,若小于ε则随机探索(选任意动作)
if np.random.uniform(0, 1) < epsilon:
action = np.random.choice(Q_table.shape[1]) # 随机选动作
# 否则利用(选当前状态下Q值最大的动作)
else:
action = np.argmax(Q_table[state, :]) # argmax返回最大值索引(即最优动作)
return action
模块三:实现Q-Learning算法
class QLearning:
def __init__(self, n_states, n_actions, alpha=0.1, gamma=0.9, epsilon=0.1):
"""
初始化Q-Learning智能体
参数:
n_states: 状态总数
n_actions: 动作总数
alpha: 学习率
gamma: 折扣因子
epsilon: ε-贪心的探索概率
"""
# 初始化Q表:n_states行(状态)×n_actions列(动作),初始值全为0
self.Q_table = np.zeros((n_states, n_actions))
self.alpha = alpha # 学习率
self.gamma = gamma # 折扣因子
self.epsilon = epsilon # 探索概率
self.n_actions = n_actions # 动作总数
def update(self, state, action, reward, next_state, done):
"""
Q-Learning核心更新逻辑(离策略)
参数:
state: 当前状态S_t
action: 当前动作A_t
reward: 即时奖励r_{t+1}
next_state: 新状态S_{t+1}
done: 是否到达终止状态(True/False)
"""
# 步骤1:计算Q_target = 即时奖励 + γ * 新状态下最大Q值(若终止则无未来奖励)
if done:
target = reward # 终止状态,未来无奖励,target=即时奖励
else:
# Q-Learning核心:取next_state下所有动作的最大Q值(不依赖实际选的动作)
target = reward + self.gamma * np.max(self.Q_table[next_state, :])
# 步骤2:计算TD误差
td_error = target - self.Q_table[state, action]
# 步骤3:更新Q值(通用更新公式)
self.Q_table[state, action] += self.alpha * td_error
模块四:实现SARSA算法
class SARSA:
def __init__(self, n_states, n_actions, alpha=0.1, gamma=0.9, epsilon=0.1):
"""初始化SARSA智能体(参数与Q-Learning一致)"""
self.Q_table = np.zeros((n_states, n_actions))
self.alpha = alpha
self.gamma = gamma
self.epsilon = epsilon
self.n_actions = n_actions
def update(self, state, action, reward, next_state, next_action, done):
"""
SARSA核心更新逻辑(在策略)
参数:
state: 当前状态S_t
action: 当前动作A_t
reward: 即时奖励r_{t+1}
next_state: 新状态S_{t+1}
next_action: 新状态下实际选的动作A_{t+1}(SARSA特有)
done: 是否到达终止状态
"""
# 步骤1:计算Q_target = 即时奖励 + γ * 新状态下实际选的动作的Q值
if done:
target = reward
else:
# SARSA核心:取next_state下实际选的next_action的Q值(而非最大值)
target = reward + self.gamma * self.Q_table[next_state, next_action]
# 步骤2:计算TD误差
td_error = target - self.Q_table[state, action]
# 步骤3:更新Q值
self.Q_table[state, action] += self.alpha * td_error
模块五:定义训练函数
def train_agent(env, agent, n_episodes=500):
"""
训练智能体
参数:
env: gym环境
agent: 智能体(QLearning/SARSA实例)
n_episodes: 训练轮数(每轮从起点到终点为1个episode)
返回:
rewards_history: 每轮的累计奖励(用于可视化)
"""
rewards_history = [] # 记录每轮的累计奖励
# 遍历每一轮训练
for episode in range(n_episodes):
state = env.reset() # 重置环境,回到初始状态
done = False # 标记是否终止
total_reward = 0 # 记录本轮累计奖励
# ===== 针对不同算法,初始化第一个动作 =====
if isinstance(agent, SARSA):
# SARSA需要先选第一个动作(因为更新需要next_action)
action = epsilon_greedy_policy(agent.Q_table, state, agent.epsilon)
# 单轮训练循环(直到终止)
while not done:
# ===== 1. 选择当前动作(Q-Learning在这一步选,SARSA已提前选)=====
if isinstance(agent, QLearning):
action = epsilon_greedy_policy(agent.Q_table, state, agent.epsilon)
# ===== 2. 执行动作,获取环境反馈 =====
next_state, reward, done, _ = env.step(action) # gym环境step返回:新状态、奖励、是否终止、额外信息
total_reward += reward # 累计奖励
# ===== 3. 选择下一个动作(SARSA需要,Q-Learning不需要)=====
if isinstance(agent, SARSA) and not done:
next_action = epsilon_greedy_policy(agent.Q_table, next_state, agent.epsilon)
else:
next_action = None # Q-Learning无需next_action
# ===== 4. 更新Q值(核心步骤)=====
if isinstance(agent, QLearning):
agent.update(state, action, reward, next_state, done)
else: # SARSA
agent.update(state, action, reward, next_state, next_action, done)
# ===== 5. 状态/动作迭代 =====
state = next_state # 更新状态
if isinstance(agent, SARSA) and not done:
action = next_action # SARSA更新动作(Q-Learning下一轮重新选)
# 记录本轮累计奖励
rewards_history.append(total_reward)
# 每50轮打印进度
if (episode + 1) % 50 == 0:
avg_reward = np.mean(rewards_history[-50:]) # 最近50轮平均奖励
print(f"算法:{type(agent).__name__} | 轮数:{episode+1} | 最近50轮平均奖励:{avg_reward:.2f}")
return rewards_history
模块六:主函数(运行训练+可视化)
if __name__ == "__main__":
# 1. 创建悬崖行走环境
env = gym.make('CliffWalking-v0')
n_states = env.observation_space.n # 状态总数:48(4行12列)
n_actions = env.action_space.n # 动作总数:4(0=上,1=右,2=下,3=左)
# 2. 初始化智能体
q_learning_agent = QLearning(n_states, n_actions)
sarsa_agent = SARSA(n_states, n_actions)
# 3. 训练智能体
print("===== 开始训练Q-Learning =====")
q_learning_rewards = train_agent(env, q_learning_agent, n_episodes=500)
print("\n===== 开始训练SARSA =====")
sarsa_rewards = train_agent(env, sarsa_agent, n_episodes=500)
# 4. 可视化训练结果(平滑处理,更易看趋势)
def smooth_rewards(rewards, window_size=10):
"""奖励曲线平滑(移动平均)"""
return np.convolve(rewards, np.ones(window_size)/window_size, mode='valid')
# 平滑奖励曲线
q_smooth = smooth_rewards(q_learning_rewards)
s_smooth = smooth_rewards(sarsa_rewards)
# 绘制曲线
plt.figure(figsize=(10, 6))
plt.plot(q_smooth, label='Q-Learning', color='red')
plt.plot(s_smooth, label='SARSA', color='blue')
plt.xlabel('训练轮数(平滑窗口=10)')
plt.ylabel('累计奖励(越高越好)')
plt.title('Q-Learning vs SARSA 悬崖行走训练曲线')
plt.legend()
plt.grid(True, alpha=0.3)
plt.show()
# 5. 关闭环境
env.close()
运行结果

- 红色曲线(Q-Learning):离策略算法,更新时直接用下一步的最大 Q 值,更激进。
- 蓝色曲线(SARSA):在策略算法,更新时用下一步实际选择的动作的 Q 值,更保守。
===== 开始训练Q-Learning =====
算法:QLearning | 轮数:50 | 最近50轮平均奖励:-236.80
算法:QLearning | 轮数:100 | 最近50轮平均奖励:-69.94
算法:QLearning | 轮数:150 | 最近50轮平均奖励:-61.86
算法:QLearning | 轮数:200 | 最近50轮平均奖励:-47.86
算法:QLearning | 轮数:250 | 最近50轮平均奖励:-43.38
算法:QLearning | 轮数:300 | 最近50轮平均奖励:-50.24
算法:QLearning | 轮数:350 | 最近50轮平均奖励:-46.74
算法:QLearning | 轮数:400 | 最近50轮平均奖励:-40.20
算法:QLearning | 轮数:450 | 最近50轮平均奖励:-33.98
算法:QLearning | 轮数:500 | 最近50轮平均奖励:-36.22
===== 开始训练SARSA =====
算法:SARSA | 轮数:50 | 最近50轮平均奖励:-208.72
算法:SARSA | 轮数:100 | 最近50轮平均奖励:-65.28
算法:SARSA | 轮数:150 | 最近50轮平均奖励:-44.50
算法:SARSA | 轮数:200 | 最近50轮平均奖励:-32.96
算法:SARSA | 轮数:250 | 最近50轮平均奖励:-31.64
算法:SARSA | 轮数:300 | 最近50轮平均奖励:-22.30
算法:SARSA | 轮数:350 | 最近50轮平均奖励:-20.50
算法:SARSA | 轮数:400 | 最近50轮平均奖励:-23.18
算法:SARSA | 轮数:450 | 最近50轮平均奖励:-23.96
算法:SARSA | 轮数:500 | 最近50轮平均奖励:-19.86
四、总结
- 核心差异:Q-Learning 是 “离策略”,SARSA 是 “在策略”;Q-Learning 激进,SARSA 保守。
- 适用场景:危险环境,如机器人避障。优先选 SARSA安全,无风险环境,如游戏,可选 Q-Learning,收敛更快。
感谢大家的观看,有不足请大家的批评指正!接下来,我将更新更加具体,有现实意义的机器学习方法!
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐


所有评论(0)