【面试必问】强化学习中的时序差分方法:从TD(0)到 eligibility traces 详解

强化学习中的时序差分方法:从TD(0)到 eligibility traces 详解

引言

时序差分(Temporal Difference, TD)学习是强化学习的"黄金中庸之道"——它像蒙特卡洛一样无需环境模型,又像动态规划一样通过自举(Bootstrapping)实现高效学习。TD方法结合了二者的精华,成为现代强化学习(如DQN、A3C等)的核心引擎。

本文将深入剖析TD的理论基础、经典算法及其实现细节。


目录


一、核心思想:TD误差的诞生

1.1 TD学习的核心公式

TD(0)更新规则

V ( S t ) ← V ( S t ) + α [ R t + 1 + γ V ( S t + 1 ) − V ( S t ) ] ⏟ TD误差  δ t V(S_t) \leftarrow V(S_t) + \alpha \underbrace{[R_{t+1} + \gamma V(S_{t+1}) - V(S_t)]}_{\text{TD误差 }\delta_t} V(St)V(St)+αTD误差 δt [Rt+1+γV(St+1)V(St)]

其中:

  • δ t \delta_t δt 是时序差分误差,衡量当前估计与一步目标之间的差距
  • α \alpha α 是学习率
  • 目标值 R t + 1 + γ V ( S t + 1 ) R_{t+1} + \gamma V(S_{t+1}) Rt+1+γV(St+1) 被称为TD目标

1.2 与MC和DP的本质区别

方法目标值更新时机特性
MC G t = ∑ k = 0 T − t − 1 γ k R t + 1 + k G_t = \sum_{k=0}^{T-t-1} \gamma^k R_{t+1+k} Gt=k=0Tt1γkRt+1+kEpisode结束无偏差,高方差
DP$\sum_{s’}P(s’s,a)[R + \gamma V(s’)]$每个step
TD R t + 1 + γ V ( S t + 1 ) R_{t+1} + \gamma V(S_{t+1}) Rt+1+γV(St+1)每个step有偏差,低方差

关键洞察:TD通过"一步自举"将MC的无模型特性与DP的高效更新完美结合。


二、TD预测:估计价值函数

2.1 TD(0)算法

def td0_prediction(env, policy, num_episodes=10000, alpha=0.1, gamma=0.9):
    V = defaultdict(float)
    
    for _ in range(num_episodes):
        state = env.reset()
        done = False
        
        while not done:
            action = policy(state)
            next_state, reward, done, _ = env.step(action)
            
            # TD(0)核心更新
            td_error = reward + gamma * V[next_state] - V[state]
            V[state] += alpha * td_error
            
            state = next_state
    
    return V

2.2 批处理TD(Batch TD)

在固定数据集上重复更新直至收敛:

V ( s ) = ∑ k ∑ t I ( S t k = s ) ⋅ [ R t + 1 k + γ V ( S t + 1 k ) ] ∑ k ∑ t I ( S t k = s ) V(s) = \frac{\sum_{k} \sum_{t} \mathbb{I}(S_t^k = s) \cdot [R_{t+1}^k + \gamma V(S_{t+1}^k)]}{\sum_{k} \sum_{t} \mathbb{I}(S_t^k = s)} V(s)=ktI(Stk=s)ktI(Stk=s)[Rt+1k+γV(St+1k)]

def batch_td(env, policy, num_episodes=100, alpha=0.1, gamma=0.9, epochs=100):
    V = defaultdict(float)
    dataset = []
    
    # 收集数据
    for _ in range(num_episodes):
        episode = []
        state = env.reset()
        done = False
        
        while not done:
            action = policy(state)
            next_state, reward, done, _ = env.step(action)
            episode.append((state, action, reward, next_state))
            state = next_state
        
        dataset.append(episode)
    
    # 批处理更新
    for epoch in range(epochs):
        for episode in dataset:
            for state, action, reward, next_state in episode:
                td_error = reward + gamma * V[next_state] - V[state]
                V[state] += alpha * td_error
    
    return V

三、TD控制:SARSA与Q-learning

3.1 SARSA(On-policy TD控制)

名称来源:State-Action-Reward-State-Action

更新公式

Q ( S t , A t ) ← Q ( S t , A t ) + α [ R t + 1 + γ Q ( S t + 1 , A t + 1 ) − Q ( S t , A t ) ] Q(S_t, A_t) \leftarrow Q(S_t, A_t) + \alpha [R_{t+1} + \gamma Q(S_{t+1}, A_{t+1}) - Q(S_t, A_t)] Q(St,At)Q(St,At)+α[Rt+1+γQ(St+1,At+1)Q(St,At)]

def sarsa(env, num_episodes=10000, alpha=0.1, gamma=0.9, epsilon=0.1):
    Q = defaultdict(lambda: np.zeros(env.action_space.n))
    
    def epsilon_greedy_policy(state):
        if random.random() < epsilon:
            return env.action_space.sample()
        else:
            return np.argmax(Q[state])
    
    for _ in range(num_episodes):
        state = env.reset()
        action = epsilon_greedy_policy(state)
        done = False
        
        while not done:
            next_state, reward, done, _ = env.step(action)
            next_action = epsilon_greedy_policy(next_state)
            
            # SARSA更新
            td_error = reward + gamma * Q[next_state][next_action] - Q[state][action]
            Q[state][action] += alpha * td_error
            
            state, action = next_state, next_action
    
    return Q

3.2 Q-learning(Off-policy TD控制)

核心创新:使用最大Q值更新,直接逼近最优策略

Q ( S t , A t ) ← Q ( S t , A t ) + α [ R t + 1 + γ max ⁡ a Q ( S t + 1 , a ) − Q ( S t , A t ) ] Q(S_t, A_t) \leftarrow Q(S_t, A_t) + \alpha [R_{t+1} + \gamma \max_a Q(S_{t+1}, a) - Q(S_t, A_t)] Q(St,At)Q(St,At)+α[Rt+1+γamaxQ(St+1,a)Q(St,At)]

def q_learning(env, num_episodes=10000, alpha=0.1, gamma=0.9, epsilon=0.1):
    Q = defaultdict(lambda: np.zeros(env.action_space.n))
    
    def epsilon_greedy_policy(state):
        if random.random() < epsilon:
            return env.action_space.sample()
        else:
            return np.argmax(Q[state])
    
    for _ in range(num_episodes):
        state = env.reset()
        done = False
        
        while not done:
            action = epsilon_greedy_policy(state)
            next_state, reward, done, _ = env.step(action)
            
            # Q-learning更新(使用max Q值)
            best_next_action = np.argmax(Q[next_state])
            td_target = reward + gamma * Q[next_state][best_next_action]
            td_error = td_target - Q[state][action]
            Q[state][action] += alpha * td_error
            
            state = next_state
    
    return Q

3.3 SARSA vs Q-learning对比

维度SARSAQ-learning
策略类型On-policyOff-policy
更新目标 Q ( S t + 1 , A t + 1 ) Q(S_{t+1}, A_{t+1}) Q(St+1,At+1) max ⁡ a Q ( S t + 1 , a ) \max_a Q(S_{t+1}, a) maxaQ(St+1,a)
行为特点保守,考虑探索风险激进,直接学最优
适用场景安全关键任务性能优化任务

四、n步TD与TD(λ)

4.1 n步TD

n步回报

G t ( n ) = R t + 1 + γ R t + 2 + ⋯ + γ n − 1 R t + n + γ n V ( S t + n ) G_{t}^{(n)} = R_{t+1} + \gamma R_{t+2} + \cdots + \gamma^{n-1} R_{t+n} + \gamma^n V(S_{t+n}) Gt(n)=Rt+1+γRt+2++γn1Rt+n+γnV(St+n)

n步TD更新

V ( S t ) ← V ( S t ) + α [ G t ( n ) − V ( S t ) ] V(S_t) \leftarrow V(S_t) + \alpha [G_t^{(n)} - V(S_t)] V(St)V(St)+α[Gt(n)V(St)]

def n_step_td(env, policy, n=4, alpha=0.1, gamma=0.9, num_episodes=10000):
    V = defaultdict(float)
    
    for _ in range(num_episodes):
        states, rewards = [], []
        state = env.reset()
        states.append(state)
        T = float('inf')
        t = 0
        
        while True:
            if t < T:
                action = policy(state)
                next_state, reward, done, _ = env.step(action)
                states.append(next_state)
                rewards.append(reward)
                
                if done:
                    T = t + 1
            
            tau = t - n + 1  # 要更新的时间步
            
            if tau >= 0:
                # 计算n步回报
                G = sum(gamma**i * rewards[tau + i] for i in range(min(n, T - tau)))
                
                if tau + n < T:
                    G += gamma**n * V[states[tau + n]]
                
                # 更新价值函数
                V[states[tau]] += alpha * (G - V[states[tau]])
            
            if tau == T - 1:
                break
            
            t += 1
    
    return V

4.2 TD(λ)与资格迹

前向视角TD(λ)

G t λ = ( 1 − λ ) ∑ n = 1 ∞ λ n − 1 G t ( n ) G_t^\lambda = (1-\lambda)\sum_{n=1}^{\infty} \lambda^{n-1} G_t^{(n)} Gtλ=(1λ)n=1λn1Gt(n)

后向视角(更有效):使用** eligibility traces **

E t ( s ) = γ λ E t − 1 ( s ) + 1 ( S t = s ) E_t(s) = \gamma\lambda E_{t-1}(s) + \mathbb{1}(S_t = s) Et(s)=γλEt1(s)+1(St=s)

V ( s ) ← V ( s ) + α δ t E t ( s ) V(s) \leftarrow V(s) + \alpha \delta_t E_t(s) V(s)V(s)+αδtEt(s)

class TDLambda:
    def __init__(self, gamma=0.9, lamb=0.8, alpha=0.1):
        self.gamma = gamma
        self.lamb = lamb
        self.alpha = alpha
        self.V = defaultdict(float)
        self.E = defaultdict(float)
    
    def learn(self, env, policy, num_episodes=10000):
        for _ in range(num_episodes):
            # 重置资格迹
            self.E.clear()
            
            state = env.reset()
            done = False
            
            while not done:
                action = policy(state)
                next_state, reward, done, _ = env.step(action)
                
                # TD误差
                delta = reward + self.gamma * self.V[next_state] - self.V[state]
                
                # 更新资格迹
                self.E[state] += 1
                
                # 对所有状态更新
                for s in list(self.V.keys()):
                    self.V[s] += self.alpha * delta * self.E[s]
                    self.E[s] *= self.gamma * self.lamb
                
                state = next_state
        
        return self.V

五、完整代码实现:悬崖漫步

import gym
import numpy as np
from collections import defaultdict
import matplotlib.pyplot as plt

class CliffWalkingSolver:
    def __init__(self, env_name='CliffWalking-v0'):
        self.env = gym.make(env_name)
        self.n_states = self.env.nS
        self.n_actions = self.env.nA
        self.shape = (4, 12)
    
    def q_learning(self, num_episodes=500, alpha=0.5, gamma=0.95, epsilon=0.1):
        Q = defaultdict(lambda: np.zeros(self.n_actions))
        episode_rewards = []
        
        for episode in range(num_episodes):
            state = self.env.reset()
            total_reward = 0
            done = False
            
            while not done:
                # ε-贪婪策略
                if random.random() < epsilon:
                    action = self.env.action_space.sample()
                else:
                    action = np.argmax(Q[state])
                
                next_state, reward, done, _ = self.env.step(action)
                total_reward += reward
                
                # Q-learning更新
                best_next_action = np.argmax(Q[next_state])
                td_target = reward + gamma * Q[next_state][best_next_action]
                td_error = td_target - Q[state][action]
                Q[state][action] += alpha * td_error
                
                state = next_state
            
            episode_rewards.append(total_reward)
        
        return Q, episode_rewards
    
    def sarsa(self, num_episodes=500, alpha=0.5, gamma=0.95, epsilon=0.1):
        Q = defaultdict(lambda: np.zeros(self.n_actions))
        episode_rewards = []
        
        for episode in range(num_episodes):
            state = self.env.reset()
            # ε-贪婪选择初始动作
            action = self.env.action_space.sample() if random.random() < epsilon else np.argmax(Q[state])
            
            total_reward = 0
            done = False
            
            while not done:
                next_state, reward, done, _ = self.env.step(action)
                total_reward += reward
                
                # ε-贪婪选择下一个动作
                next_action = self.env.action_space.sample() if random.random() < epsilon else np.argmax(Q[next_state])
                
                # SARSA更新
                td_target = reward + gamma * Q[next_state][next_action]
                td_error = td_target - Q[state][action]
                Q[state][action] += alpha * td_error
                
                state, action = next_state, next_action
            
            episode_rewards.append(total_reward)
        
        return Q, episode_rewards
    
    def plot_results(self, sarsa_rewards, q_rewards):
        plt.figure(figsize=(12, 5))
        
        # 原始奖励
        plt.subplot(1, 2, 1)
        plt.plot(sarsa_rewards, label='SARSA', alpha=0.7)
        plt.plot(q_rewards, label='Q-learning', alpha=0.7)
        plt.xlabel('Episode')
        plt.ylabel('Total Reward')
        plt.title('Episode Rewards')
        plt.legend()
        
        # 移动平均
        window = 50
        sarsa_ma = np.convolve(sarsa_rewards, np.ones(window)/window, mode='valid')
        q_ma = np.convolve(q_rewards, np.ones(window)/window, mode='valid')
        
        plt.subplot(1, 2, 2)
        plt.plot(sarsa_ma, label='SARSA')
        plt.plot(q_ma, label='Q-learning')
        plt.xlabel('Episode')
        plt.ylabel(f'{window}-Episode Moving Average')
        plt.title('Smoothed Performance')
        plt.legend()
        
        plt.tight_layout()
        plt.savefig('td_comparison.png', dpi=150)
        plt.show()
    
    def visualize_policy(self, Q, title="Policy"):
        policy = np.zeros(self.shape, dtype=str)
        arrow_map = {0: '↑', 1: '↓', 2: '←', 3: '→'}
        
        for s in range(self.n_states):
            row, col = divmod(s, self.shape[1])
            if self.env.desc[row, col] == b'S':
                policy[row, col] = 'S'
            elif self.env.desc[row, col] == b'G':
                policy[row, col] = 'G'
            elif self.env.desc[row, col] == b'H':
                policy[row, col] = 'H'
            else:
                policy[row, col] = arrow_map[np.argmax(Q[s])]
        
        print(f"\n{title}:")
        print(policy)

# 主执行
if __name__ == "__main__":
    import random
    
    solver = CliffWalkingSolver()
    
    # 训练两种算法
    print("Training SARSA...")
    Q_sarsa, rewards_sarsa = solver.sarsa(num_episodes=500)
    
    print("\nTraining Q-learning...")
    Q_q, rewards_q = solver.q_learning(num_episodes=500)
    
    # 可视化策略
    solver.visualize_policy(Q_sarsa, "SARSA Policy")
    solver.visualize_policy(Q_q, "Q-learning Policy")
    
    # 绘制结果
    solver.plot_results(rewards_sarsa, rewards_q)
    
    # 性能统计
    print(f"\nSARSA最后100集平均奖励: {np.mean(rewards_sarsa[-100:]):.2f}")
    print(f"Q-learning最后100集平均奖励: {np.mean(rewards_q[-100:]):.2f}")

六、收敛性分析

6.1 收敛条件

TD算法收敛需满足:

  1. 学习率条件 ∑ k = 1 ∞ α k = ∞ \sum_{k=1}^{\infty} \alpha_k = \infty k=1αk= ∑ k = 1 ∞ α k 2 < ∞ \sum_{k=1}^{\infty} \alpha_k^2 < \infty k=1αk2<
  2. 充分探索:所有状态-动作对无限次访问
  3. 策略性质:对于On-policy,策略需满足GLIE条件

6.2 收敛速度对比

  • TD(0) O ( 1 ϵ ) O(\frac{1}{\epsilon}) O(ϵ1) 收敛速度
  • Q-learning:在ε-greedy探索下收敛到最优策略
  • SARSA:收敛到ε-贪婪最优策略(更保守)

七、MC vs DP vs TD终极对比

对比维度蒙特卡洛(MC)动态规划(DP)时序差分(TD)
模型需求❌ 无模型✅ 需完整模型❌ 无模型
更新目标完整回报 G t G_t Gt贝尔曼期望一步TD目标
更新时机Episode结束每个step每个step
偏差/方差无偏差/高方差有偏差/低方差有偏差/中低方差
收敛速度最快
内存效率
在线学习❌ 离线✅ 在线✅ 在线
实现复杂度
适用任务Episodic任意MDP任意MDP

选择指南

  • 需要模型? → DP
  • 必须无偏? → MC
  • 效率优先?TD(首选)

八、优缺点与实战技巧

✅ 优点

  1. 无需模型:适用于真实世界问题
  2. 在线学习:可增量更新,边学边用
  3. 低方差:比MC更稳定
  4. 实现简单:代码量小,易于调试
  5. 收敛快速:通常比MC快得多

❌ 缺点

  1. 有偏估计:初始阶段依赖不准确的估计
  2. 敏感超参:学习率α选择关键
  3. 需要精心设计探索:ε-greedy可能不够
  4. 可能振荡:Q-learning在非平稳环境下可能不稳定

🔧 实战技巧

1. 学习率调度

# 线性衰减
alpha = initial_alpha * (1 - episode / total_episodes)

# 指数衰减
alpha = initial_alpha * (decay_rate ** episode)

2. ε-greedy改进

# GLIE调度:ε = 1/√episode
epsilon = 1.0 / np.sqrt(episode + 1)

# 带最低探索率的衰减
epsilon = max(0.01, epsilon * 0.995)

3. 双重Q-learning(减少乐观偏差)

def double_q_learning(env, num_episodes=10000):
    Q1 = defaultdict(lambda: np.zeros(env.action_space.n))
    Q2 = defaultdict(lambda: np.zeros(env.action_space.n))
    
    for _ in range(num_episodes):
        state = env.reset()
        done = False
        
        while not done:
            action = epsilon_greedy(state, Q1 + Q2)  # 使用Q1+Q2指导探索
            
            next_state, reward, done, _ = env.step(action)
            
            # 随机更新其中一个Q函数
            if random.random() < 0.5:
                best_action = np.argmax(Q1[next_state])
                Q1[state][action] += alpha * (reward + gamma * Q2[next_state][best_action] - Q1[state][action])
            else:
                best_action = np.argmax(Q2[next_state])
                Q2[state][action] += alpha * (reward + gamma * Q1[next_state][best_action] - Q2[state][action])
            
            state = next_state
    
    return {s: (Q1[s] + Q2[s]) / 2 for s in Q1.keys()}

参考文献

  1. Sutton, R. S. (1988). Learning to predict by the methods of temporal differences. Machine Learning.
  2. Watkins, C. J. C. H. (1989). Learning from delayed rewards. PhD thesis.
  3. Sutton, R. S., & Barto, A. G. (2018). Reinforcement Learning: An Introduction (2nd ed.).
  4. van Hasselt, H., Guez, A., & Silver, D. (2016). Deep Reinforcement Learning with Double Q-learning.
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包

打赏作者

litterfinger

你的鼓励将是我创作的最大动力

¥1 ¥2 ¥4 ¥6 ¥10 ¥20
扫码支付:¥1
获取中
扫码支付

您的余额不足,请更换扫码支付或充值

打赏作者

实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值