Hoppa till innehållet
AI-grafen
F· AI engineeringreinforcement-learning· ca 90 min· volatil — kontrolleras ofta· verifierad 2026-09-20

Deep Q-Networks

Kunna träna en DQN med replay buffer och target network.

Förkunskaper

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.

ProblemVarförLösning
Korrelerade datapå varandra följande steg är nästan identiska; SGD antar oberoende stickprovreplay buffer — spara övergångar, slumpa ur den
Rörligt målmålet r+γmax⁡Q(s′,⋅)r + \gamma\max Q(s',\cdot) beräknas med samma nät som tränas; jaga sin egen svanstarget network — en frusen kopia som uppdateras då och då
Överskattningmax⁡\max över brusiga skattningar är systematiskt för högtDouble 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:

  1. Kontrollera (1 - klar).
  2. Skriv ut medel-Q över tid — växer det obegränsat är något fel med målet.
  3. Kontrollera att epsilon faktiskt minskar.
  4. Kontrollera att bufferten hunnit fyllas innan träningen startar.
  5. 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

Alla källor och licenser