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.符号定义

Q_{t}\left ( S,A \right ),第 t 步时,状态 S 下执行动作 A 的 Q 值估计;

\alpha,学习率(0 < α ≤ 1);

\gamma,折扣因子(0 ≤ γ ≤ 1);

r_{t+1},执行动作 A 后获得的即时奖励;

S_{t+1},执行动作 A 后进入的新状态;

A_{t+1},在新状态​S_{t+1}下选择的动作(SARSA 专用);

2. 时序差分(TD)误差核心框架

        两种算法的本质都是用 “TD 误差” 更新 Q 值,TD 误差的核心思想是:

用 “当前估计的 Q 值” 与 “实际奖励 + 未来 Q 值期望” 的差值,修正当前 Q 值。

基础 TD 误差公式:

TD_{error}=\left [r _{t+1}+\gamma \cdot Q_{target}\left ( S_{t+1} ,A_{t+1}\right ) \right ]-Q_{t}\left ( S_{t} ,A_{t}\right );

Q 值更新的通用公式:

Q_{t+1}\left ( S_{t},A_{t} \right )=Q_{t}\left ( S_{t},A_{t} \right )+\alpha \cdot TD_{error};

3. SARSA 更新公式

SARSA 的Q_{target}取当前策略下实际选择的下一个动作的 Q 值:

Q_{t+1}\left (S _{t} ,A_{t}\right )=Q_{t}\left (S _{t} ,A_{t}\right )+\alpha \cdot \left [R _{t+1}+\gamma \cdot ma x_{a} Q_{t}\left ( S_{t+1},a \right )-Q_{t}\left (S _{t} ,A_{t}\right )\right ];

关键:ma x_{a} Q_{t}\left ( S_{t+1},a \right ) 表示不考虑实际选什么动作,只取新状态下最优动作的 Q 值。

4. SARSA 更新公式

SARSA 的​Q_{target}取当前策略下实际选择的下一个动作的 Q 值:

Q_{t+1}\left (S _{t} ,A_{t}\right )=Q_{t}\left (S _{t} ,A_{t}\right )+\alpha \cdot \left [R _{t+1}+\gamma \cdot Q_{t}\left ( S_{t+1},a \right )-Q_{t}\left (S _{t} ,A_{t}\right )\right ];

关键:A_{t+1}是智能体根据当前策略(如 ε- 贪心)在S_{t+1}下实际选的动作,而非最优动作。

三、代码解释

模块一:导入核心库

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,收敛更快。

        感谢大家的观看,有不足请大家的批评指正!接下来,我将更新更加具体,有现实意义的机器学习方法!

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐