Deep Q-Networks
Kunna träna en DQN med replay buffer och target network.
Förkunskaper
- DNeuronnät — forward pass med matriserkrävs
- EQ-learningkrävs
Intuition
Q-learning med en tabell fungerar. Byt tabellen mot ett neuralt nät och allt går sönder — av tre skäl, som alla har varsin lösning.
| Problem | Varför | Lösning |
|---|---|---|
| Korrelerade data | på varandra följande steg är nästan identiska; SGD antar oberoende stickprov | replay buffer — spara övergångar, slumpa ur den |
| Rörligt mål | målet beräknas med samma nät som tränas; jaga sin egen svans | target network — en frusen kopia som uppdateras då och då |
| Överskattning | över brusiga skattningar är systematiskt för högt | Double DQN — välj handling med ett nät, värdera med det andra |
Det var kombinationen av de två första som fick DQN att spela Atari 2013 — inte nätverksarkitekturen, som var helt vanlig.
Överskattningen är subtil och värd en egen mening: även om varje Q-skattning är väntevärdesriktig är väntevärdet av maximum större än maximum av väntevärdena. Bruset går alltid åt samma håll när man tar max.
Kod
import random, torch, torch.nn as nn
from collections import deque
class Replay:
def __init__(self, kapacitet=100_000):
self.buf = deque(maxlen=kapacitet) # ringbuffert: äldst faller ut automatiskt
def lagg_till(self, s, a, r, s2, klar):
self.buf.append((s, a, r, s2, float(klar)))
def stickprov(self, n, rng=random):
batch = rng.sample(self.buf, n)
return [torch.as_tensor(x, dtype=torch.float32) for x in zip(*batch)]
def __len__(self):
return len(self.buf)
q = nn.Sequential(nn.Linear(4, 128), nn.ReLU(), nn.Linear(128, 2))
mal_q = nn.Sequential(nn.Linear(4, 128), nn.ReLU(), nn.Linear(128, 2))
mal_q.load_state_dict(q.state_dict())
opt = torch.optim.Adam(q.parameters(), lr=1e-3)
def trana_steg(replay, batch=64, gamma=0.99, double=True):
s, a, r, s2, klar = replay.stickprov(batch)
qsa = q(s).gather(1, a.long().unsqueeze(1)).squeeze(1)
with torch.no_grad():
if double:
b = q(s2).argmax(1, keepdim=True) # VÄLJ med online-nätet
q2 = mal_q(s2).gather(1, b).squeeze(1) # VÄRDERA med målnätet
else:
q2 = mal_q(s2).max(1).values
mal = r + gamma * q2 * (1 - klar) # (1 - klar) — den klassiska buggen
forlust = nn.functional.smooth_l1_loss(qsa, mal) # Huber tål avvikare bättre än MSE
opt.zero_grad(); forlust.backward()
nn.utils.clip_grad_norm_(q.parameters(), 10.0)
opt.step()
return float(forlust)
# Målnätet uppdateras sällan — var 1 000:e steg, eller mjukt varje steg:
def mjuk_uppdatering(tau=0.005):
for p, pm in zip(q.parameters(), mal_q.parameters()):
pm.data.mul_(1 - tau).add_(tau * p.data)
(1 - klar) är den bugg som kostar flest timmar. Utan den bootstrappar agenten på värdet av ett tillstånd bortom episodens slut. Q-värdena växer, förlusten ser rimlig ut, och agenten lär sig aldrig. Den syns inte i någon felutskrift.
Felsökningsordning när DQN inte lär sig:
- Kontrollera
(1 - klar). - Skriv ut medel-Q över tid — växer det obegränsat är något fel med målet.
- Kontrollera att
epsilonfaktiskt minskar. - Kontrollera att bufferten hunnit fyllas innan träningen startar.
- Först därefter: hyperparametrar.
Behärskning innebär
- Förklarar varför replay buffer och target network behövs
- Implementerar DQN-målet korrekt
- Känner till överskattningsproblemet och Double DQN
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — Playing Atari with Deep Reinforcement Learning — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Deep Reinforcement Learning with Double Q-learning — arXiv (öppen åtkomst; licens per artikel)
- OpenAI Spinning Up in Deep RL (MIT) — MIT