Skip to main content

8. Function Approximation

表格型方法为每个状态或状态—动作对单独保存一个数。当状态空间巨大、连续,或者观测本身是图像时,这种方法无法扩展。函数近似用参数向量 w\mathbf w 表示价值:

v^(s,w)≈vπ(s),q^(s,a,w)≈qπ(s,a).\hat v(s,\mathbf w)\approx v_\pi(s), \qquad \hat q(s,a,\mathbf w)\approx q_\pi(s,a).

不同状态共享参数,因此一次更新不仅改变当前样本的估计,也可能泛化到其他相似状态。

1. 价值近似的目标函数​

设 μ(s)\mu(s) 是训练数据中的状态分布。一个自然的均方误差目标为:

J(w)=ES∼μ[(vπ(S)−v^(S,w))2].J(\mathbf w) =\mathbb E_{S\sim\mu} \left[ \bigl(v_\pi(S)-\hat v(S,\mathbf w)\bigr)^2 \right].

梯度为:

∇wJ(w)=E[2(vπ(S)−v^(S,w))(−∇wv^(S,w))].\begin{aligned} \nabla_{\mathbf w}J(\mathbf w) &=\mathbb E\left[ 2\bigl(v_\pi(S)-\hat v(S,\mathbf w)\bigr) \bigl(-\nabla_{\mathbf w}\hat v(S,\mathbf w)\bigr) \right]. \end{aligned}

吸收常数 22 后,理想的随机梯度更新是:

wt+1=wt+αt[vπ(St)−v^(St,wt)]∇wv^(St,wt).\mathbf w_{t+1} =\mathbf w_t +\alpha_t \left[v_\pi(S_t)-\hat v(S_t,\mathbf w_t)\right] \nabla_{\mathbf w}\hat v(S_t,\mathbf w_t).

但 vπ(St)v_\pi(S_t) 正是未知目标,因此还需要用采样 target 近似它。

2. Gradient Monte Carlo​

完整回报满足:

Eπ[Gt∣St=s]=vπ(s).\mathbb E_\pi[G_t\mid S_t=s]=v_\pi(s).

因此可以用 GtG_t 替代未知的 vπ(St)v_\pi(S_t):

wt+1=wt+αt[Gt−v^(St,wt)]∇wv^(St,wt).\boxed{ \mathbf w_{t+1} =\mathbf w_t +\alpha_t \left[G_t-\hat v(S_t,\mathbf w_t)\right] \nabla_{\mathbf w}\hat v(S_t,\mathbf w_t) }.

这是真正意义上对均方价值误差做随机梯度下降,因为 MC target 不依赖当前参数 wt\mathbf w_t。

3. Semi-gradient TD(0)​

TD 用单步自举目标:

Yt=Rt+1+γv^(St+1,wt).Y_t=R_{t+1}+\gamma\hat v(S_{t+1},\mathbf w_t).

更新为:

wt+1=wt+αt[Rt+1+γv^(St+1,wt)−v^(St,wt)]∇wv^(St,wt).\boxed{ \mathbf w_{t+1} =\mathbf w_t +\alpha_t \left[ R_{t+1}+\gamma\hat v(S_{t+1},\mathbf w_t) -\hat v(S_t,\mathbf w_t) \right] \nabla_{\mathbf w}\hat v(S_t,\mathbf w_t) }.

之所以叫 semi-gradient,是因为 YtY_t 也依赖 wt\mathbf w_t,但更新时把 target 当作常量,没有对 target 一侧求导。

如果直接对平方 TD error 的两侧都求导,会得到 residual-gradient 类方法;它优化的是另一种目标,和标准 TD 的不动点更新不是同一回事。

4. 带函数近似的 Sarsa​

令 q^(s,a,w)\hat q(s,a,\mathbf w) 近似动作价值。Sarsa 的 TD error 为:

δt=Rt+1+γq^(St+1,At+1,wt)−q^(St,At,wt).\delta_t =R_{t+1}+\gamma\hat q(S_{t+1},A_{t+1},\mathbf w_t) -\hat q(S_t,A_t,\mathbf w_t).

半梯度更新:

wt+1=wt+αtδt∇wq^(St,At,wt).\boxed{ \mathbf w_{t+1} =\mathbf w_t +\alpha_t\delta_t \nabla_{\mathbf w}\hat q(S_t,A_t,\mathbf w_t) }.

策略仍然可以相对于 q^\hat q 做 ε\varepsilon-greedy 更新。这是 on-policy 控制,因为 At+1A_{t+1} 由当前行为策略实际选出。

5. 带函数近似的 Q-learning​

Q-learning 的目标为:

Yt=Rt+1+γmax⁡a′q^(St+1,a′,wt).Y_t =R_{t+1}+\gamma\max_{a'} \hat q(S_{t+1},a',\mathbf w_t).

半梯度更新为:

wt+1=wt+αt[Yt−q^(St,At,wt)]∇wq^(St,At,wt).\boxed{ \mathbf w_{t+1} =\mathbf w_t +\alpha_t \left[Y_t-\hat q(S_t,A_t,\mathbf w_t)\right] \nabla_{\mathbf w}\hat q(S_t,A_t,\mathbf w_t) }.

这里同样把 YtY_t 当作固定标签。

这里的

max⁡a′q^(St+1,a′,wt)\max_{a'} \hat q(S_{t+1},a',w_t)

就是:把下一个状态 St+1S_{t+1} 下所有可能动作都送进 Q 网络,算出每个动作的 Q 值,然后取最大的那个。

通常不是一个动作一个动作真的循环调用网络 更常见的是网络输入一个状态:

St+1S_{t+1}

然后一次性输出所有动作的 Q:

Qθ(St+1)=[Q(St+1,a1)Q(St+1,a2)⋮Q(St+1,am)].Q_\theta(S_{t+1}) = \begin{bmatrix} Q(S_{t+1},a_1)\\ Q(S_{t+1},a_2)\\ \vdots\\ Q(S_{t+1},a_m) \end{bmatrix}.

所以一般适合离散/离散化后的动作

6. DQN:损失函数、经验回放与目标网络​

从表格型 Q-learning 变成神经网络近似后,原本简单的 TD 更新会遇到几个稳定性问题:

  1. 目标不断移动:TD target 本身包含当前的 QQ 估计。每次更新参数后,预测值变了,训练标签也跟着变,网络等于在追逐一个不断变化的目标。
  2. 样本高度相关:与环境连续交互得到的相邻转移通常相似,不满足 minibatch 梯度方法偏好的近似独立条件。
  3. 数据分布不断变化:随着行为策略和网络参数更新,收集到的状态分布也在变化;如果只使用最新数据,训练容易振荡并且数据利用率很低。

DQN 用两个互补的机制缓解这些问题:

  • 目标网络暂时冻结生成 TD target 所用的参数,减慢目标的变化速度;
  • 经验回放把历史转移混合后随机采样,打散时间相关性、重复利用数据,并让训练分布变化得更平缓。

下面分别说明 DQN 的损失函数、经验回放、目标网络,以及它们组合后的训练流程。

6.1 损失函数​

Deep Q-Network(DQN)用神经网络 Q(s,a;θ)Q(s,a;\boldsymbol\theta) 近似动作价值。若直接使用同一组参数构造 target:

yt=Rt+1+γmax⁡a′Q(St+1,a′;θ),y_t=R_{t+1} +\gamma\max_{a'}Q(S_{t+1},a';\boldsymbol\theta),

则参数每更新一次,预测值和训练标签都会一起变化。DQN 引入在线网络与目标网络:

  • 在线网络 Q(s,a;θ)Q(s,a;\boldsymbol\theta):用于动作选择和梯度更新;
  • 目标网络 Q(s,a;θ−)Q(s,a;\boldsymbol\theta^-):用于生成相对稳定的 target项,作为较为稳定的追逐目标

终止状态需要屏蔽自举项(到达终止态就无需估计未来价值):

yi=ri+γ(1−di)max⁡a′Q(si′,a′;θ−).\boxed{ y_i=r_i+\gamma(1-d_i) \max_{a'}Q(s_i',a';\boldsymbol\theta^-) }.

其中 di=1d_i=1 表示该转移到达终止状态。损失通常写为:

L(θ)=E(s,a,r,s′,d)∼D[(y−Q(s,a;θ))2].L(\boldsymbol\theta) =\mathbb E_{(s,a,r,s',d)\sim D} \left[ \bigl(y-Q(s,a;\boldsymbol\theta)\bigr)^2 \right].

实现时会对 yy (目标网络生成的标签) 停止梯度,只更新在线网络参数。

6.2 Experience Replay​

按时间连续收集的转移高度相关,而且数据分布会随着策略变化。经验回放把转移存入 replay buffer:

(St,At,Rt+1,St+1,dt)⟶D.(S_t,A_t,R_{t+1},S_{t+1},d_t)\longrightarrow D.

训练时从 DD 中均匀随机抽取 minibatch。它的作用包括:

  1. 打散相邻样本的时间相关性;
  2. 让同一条经验被重复使用,提高数据效率;
  3. 混合新旧策略产生的数据,使训练分布变化更平缓;
  4. 便于使用 minibatch 获得更稳定的梯度。

经验回放天然更适合 off-policy 算法,因为 buffer 中的样本可能由旧版本行为策略产生,而 Q-learning 的目标策略由当前贪心操作定义。

6.3 目标网络​

每隔 CC 个优化步骤,把在线网络参数复制给目标网络:

θ−←θ.\boldsymbol\theta^-\leftarrow\boldsymbol\theta.

也可以用软更新:

θ−←τθ+(1−τ)θ−.\boldsymbol\theta^- \leftarrow\tau\boldsymbol\theta +(1-\tau)\boldsymbol\theta^-.

在两次同步之间,目标网络保持不变或缓慢变化,减少“追逐不断移动的目标”带来的不稳定。

目标网络更新后不需要清空 replay buffer。buffer 保存的是环境转移事实,不是旧目标值;每次采样时都会用当前目标网络重新计算 yiy_i。

6.4 DQN 训练流程​

replay_buffer = ReplayBuffer(capacity)
online_q = QNetwork()
target_q = copy.deepcopy(online_q)

for step in range(num_steps):
action = epsilon_greedy(online_q(state), epsilon)
next_state, reward, terminated, truncated, _ = env.step(action)
done = terminated or truncated
replay_buffer.add(state, action, reward, next_state, done)

batch = replay_buffer.sample(batch_size)
with torch.no_grad():
target = batch.reward + gamma * (1 - batch.done) * target_q(batch.next_state).max(dim=1).values
prediction = online_q(batch.state).gather(1, batch.action.unsqueeze(1)).squeeze(1)
loss = ((target - prediction) ** 2).mean()
optimizer.zero_grad()
loss.backward()
optimizer.step()

if step % target_update_interval == 0:
target_q.load_state_dict(online_q.state_dict())

state = env.reset()[0] if done else next_state

7. 为什么函数近似更难​

深度强化学习常同时具有三个因素:

  1. Function approximation:多个状态共享参数;
  2. Bootstrapping:target 包含当前价值估计;
  3. Off-policy learning:生成数据的策略与学习目标不同。

这三者合称 deadly triad。它们并不表示算法一定发散,但会破坏表格型 on-policy 方法中简单的收敛直觉。DQN 的 replay buffer 和 target network 正是在工程上缓和这些问题的关键机制。

例子:

  • 表格型 Q-learning:有自举、离策略,但没有函数近似,通常可以收敛;
  • 神经网络 TD:有函数近似、自举,但如果是 on-policy,风险相对不同;
  • DQN:三个因素同时存在,因此属于 deadly triad 的典型实例。

8. 小结​

从表格型方法到 DQN,更新的共同骨架始终是:

参数更新=旧参数+步长×预测误差×预测梯度.\boxed{ \text{参数更新} =\text{旧参数} +\text{步长}\times\text{预测误差}\times\text{预测梯度} }.

真正变化的是 target 的构造方式:MC 使用完整回报,TD 使用一步自举,Sarsa 使用实际下一动作,Q-learning 使用贪心下一动作,而 DQN 再用目标网络和经验回放让这一过程能在高维非线性模型上稳定运行。