train-kit / src / ardegazu / train / dqn.cljs
 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
;; ported-from: src/dqn.ts @ v1.0.0 (extracted-from: bot/src/rl/train.ts @ fa686ee,
;;   the fit-prep blocks the neon-grid and valley-blocks trainers shared)
;;
;; The frozen-target DQN pieces every net game shares: candidates in, one
;; Q/value out per candidate, regression targets from the frozen net, and a
;; newest-kept replay cap. The collection loops themselves stay per-game —
;; they are shaped by each gym's step signature.
(ns ardegazu.train.dqn
  (:require [ardegazu.train.forward :as forward]))

(defn net-qs
  "One value per candidate feature vector (forward pass, output[0])."
  [net feats]
  (.map ^js feats (fn [f] (aget (forward/mlp-forward net f) 0))))

(defn dqn-targets
  "Frozen-target regression set: y = done||!f2 ? r : r + γ·max netQs(f2); x = f[a]."
  [net batch gamma]
  (let [xs #js []
        ys #js []
        n (.-length ^js batch)]
    (dotimes [i n]
      (let [t (aget batch i)
            y (if ^boolean (js* "(~{} || !~{})" (.-done ^js t) (.-f2 ^js t))
                (.-r ^js t)
                (+ (.-r ^js t)
                   (* gamma (.apply js/Math.max nil (net-qs net (.-f2 ^js t))))))]
        (.push xs (aget (.-f ^js t) (.-a ^js t)))
        (.push ys y)))
    #js {:xs xs :ys ys}))

(defn cap-replay
  "Keep only the newest `max` transitions, in place."
  [replay max]
  (when (> (.-length ^js replay) max)
    (.splice ^js replay 0 (- (.-length ^js replay) max)))
  js/undefined)

static mirror of HEAD · about · clone: git clone https://git.ardegazu.ro/train-kit.git