强化学习中的时序差分方法:从TD(0)到 eligibility traces 详解
引言
时序差分(Temporal Difference, TD)学习是强化学习的"黄金中庸之道"——它像蒙特卡洛一样无需环境模型,又像动态规划一样通过自举(Bootstrapping)实现高效学习。TD方法结合了二者的精华,成为现代强化学习(如DQN、A3C等)的核心引擎。
本文将深入剖析TD的理论基础、经典算法及其实现细节。
目录
- 一、核心思想:TD误差的诞生
- 二、TD预测:估计价值函数
- 三、TD控制:SARSA与Q-learning
- 四、n步TD与TD(λ)
- 五、完整代码实现:悬崖漫步
- 六、收敛性分析
- 七、MC vs DP vs 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=0T−t−1γkRt+1+k | Episode结束 | 无偏差,高方差 |
| 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)=∑k∑tI(Stk=s)∑k∑tI(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对比
| 维度 | SARSA | Q-learning |
|---|---|---|
| 策略类型 | On-policy | Off-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+⋯+γn−1Rt+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∑∞λn−1Gt(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)=γλEt−1(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算法收敛需满足:
- 学习率条件: ∑ 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<∞
- 充分探索:所有状态-动作对无限次访问
- 策略性质:对于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(首选)
八、优缺点与实战技巧
✅ 优点
- 无需模型:适用于真实世界问题
- 在线学习:可增量更新,边学边用
- 低方差:比MC更稳定
- 实现简单:代码量小,易于调试
- 收敛快速:通常比MC快得多
❌ 缺点
- 有偏估计:初始阶段依赖不准确的估计
- 敏感超参:学习率α选择关键
- 需要精心设计探索:ε-greedy可能不够
- 可能振荡: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()}
参考文献
- Sutton, R. S. (1988). Learning to predict by the methods of temporal differences. Machine Learning.
- Watkins, C. J. C. H. (1989). Learning from delayed rewards. PhD thesis.
- Sutton, R. S., & Barto, A. G. (2018). Reinforcement Learning: An Introduction (2nd ed.).
- van Hasselt, H., Guez, A., & Silver, D. (2016). Deep Reinforcement Learning with Double Q-learning.



1660

被折叠的 条评论
为什么被折叠?



