JOVANA
Explore Library Glossary Getting Started Three Levels Fields How it works Mission
Join the mission
All guides

DQN 配方:經驗回放與目標網路

兩個簡單的點子,把會發散的深度 Q-learning 變成征服 Atari 的代理。我們把每一個都講清楚。

DQN 的損失函數

一個 深度 Q 網路(Deep Q-Network, DQN)會為給定狀態的每個動作輸出一個值:Q(s, ·)。我們希望這些值滿足貝爾曼最適方程(Bellman optimality equation),所以把每個 Q(s, a) 回歸到一個目標 y = r + γ·maxₐ′ Q(s′, a′)。其中的差距 y − Q(s, a) 正是 TD 誤差(TD error)

如果你直接這樣訓練,兩個問題會咬人。第一,連續的影格高度相關,每個小批次看起來幾乎一模一樣,網路會對最近幾秒的玩法過擬合。第二,目標 y 是用同一個正在更新的網路算出來的,所以目標每一步都在動——你在追一個會移動的球門。下面兩個穩定器正好分別解決這兩個問題。

L(\theta)=\mathbb{E}_{(s,a,r,s')\sim D}\left[\left(r+\gamma\max_{a'}Q(s',a';\theta^{-})-Q(s,a;\theta)\right)^{2}\right]

DQN 損失同時承載兩項修復:目標網路 θ⁻ 與回放緩衝區 D 就嵌在平方 TD 誤差之中。

經驗回放與回放緩衝區

經驗回放(experience replay)會把每一筆轉移 (s, a, r, s′, done) 存進一個大型的回放緩衝區(replay buffer)(通常約一百萬筆),再從中抽取隨機小批次來訓練。隨機抽樣打碎了時間上的相關性,而且每筆經驗都能被重複使用很多次——這在樣本效率上是一大勝利。

目標網路

為了不再追移動的球門,DQN 保留一份網路的凍結副本——目標網路(target network)——並用它來計算 y。每隔 C 步(例如 10,000 步)就把線上網路的權重複製到目標網路。在兩次複製之間,目標是靜止的,於是這段期間回歸問題保持良好定義,學習也不再震盪。

回放對付的是相關性;目標網路對付的是移動的目標。兩者合起來,把不穩定的三要素轉化成可以可靠訓練的東西。任一個單獨都不夠——把其中之一移除,Atari 的分數就崩潰。

獎勵裁剪與 Huber 損失

一套網路、57 款遊戲、分數尺度天差地遠(乒乓給 ±1,彈珠台給上千分)。獎勵裁剪(reward clipping)把每個獎勵壓到 {−1, 0, +1},這樣單一學習率就能到處通用。代價是:代理再也分不出小勝和大勝——這是已知的限制,不是免費的午餐。

至於損失本身,DQN 採用 Huber 損失(Huber loss)而非單純的平方誤差。對小的 TD 誤差,它的行為像均方誤差(MSE);對大的誤差則切換成絕對(線性)誤差,這樣少數離群的轉移就無法產生破壞訓練的巨大梯度。

L_{\delta}(e)=\begin{cases}\tfrac{1}{2}e^{2} & |e|\le\delta\\[4pt]\delta\left(|e|-\tfrac{1}{2}\delta\right) & |e|>\delta\end{cases}

作用於 TD 誤差 e 的 Huber 損失:小誤差時為二次函數,超過 δ 後變為線性——對裁剪後仍存在的大誤差更穩健。

把它組合起來

下面是完整迴圈的虛擬碼。注意前處理(灰階化 + 影格堆疊)、緩衝區、凍結目標、裁剪與 Huber 損失,是如何整齊地嵌進同一個訓練步驟。

虛擬碼所實現的智能體—環境迴圈:執行動作、觀察獎勵與下一狀態、儲存轉移,然後學習。

一張圖示:智能體在環境中執行動作,並獲得獎勵與下一狀態作為回報。

buffer = ReplayBuffer(capacity=1_000_000)
Q_target = copy(Q)

for step in range(total_steps):
    a = epsilon_greedy(Q(s))                 # act
    s2, r, done = env.step(a)                # interact
    buffer.add(s, a, clip(r, -1, 1), s2, done)
    s = env.reset() if done else s2

    batch = buffer.sample(32)               # learn from the past
    y = r + gamma * max_a2 Q_target(s2, a2)  # frozen target net
    y = r if done else y                     # no bootstrap past terminal
    loss = huber(Q(s, a) - y)
    Q.update(loss)

    if step % C == 0:                        # periodic sync
        Q_target = copy(Q)
DQN 訓練迴圈:回放抽樣、凍結目標、裁剪獎勵、Huber 損失。