train-kit / src / ardegazu / train / forward.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
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
;; ported-from: src/forward.ts @ v1.0.0 (extracted-from: bot/src/rl/forward.ts @ fa686ee)
;;
;; The whole inference engine for live bots: a dense MLP forward pass over the
;; checkpoint's plain-JSON weights. Relu hidden layers, linear output — a
;; trainer (tfjs) exports exactly this shape, so the bots need no framework.
;;
;; The arithmetic order is kept identical to the TS original on purpose: the
;; golden vectors assert bit-for-bit equal doubles.
(ns ardegazu.train.forward)

(defn mlp-forward
  "Forward pass over an MlpJson net; returns the output layer as a JS array."
  [net input]
  (let [layers (.-layers ^js net)
        n-layers (.-length layers)]
    (loop [l 0
           x input]
      (if (< l n-layers)
        (let [layer (aget layers l)
              w (.-w ^js layer)
              b (.-b ^js layer)
              n-out (.-length b)
              n-in (.-length ^js x)
              y (js/Array. n-out)
              last? (== l (dec n-layers))]
          (dotimes [o n-out]
            (let [row (* o n-in)
                  acc (loop [i 0
                             acc (aget b o)]
                        (if (< i n-in)
                          (recur (inc i) (+ acc (* (aget w (+ row i)) (aget x i))))
                          acc))]
            ;; relu on hidden layers, linear output
              (aset y o (if last? acc (js/Math.max 0 acc)))))
          (recur (inc l) y))
        x))))

(defn mlp-init
  "Fresh net with small deterministic-ish random weights (trainer bootstrap)."
  ([sizes] (mlp-init sizes js/Math.random))
  ([sizes rand]
   (let [layers #js []
         n (.-length ^js sizes)]
     (loop [l 1]
       (when (< l n)
         (let [n-in (aget sizes (dec l))
               n-out (aget sizes l)
               scale (js/Math.sqrt (/ 2 n-in))
               w (js/Array. (* n-out n-in))
               b (js/Array. n-out)]
           (dotimes [k (* n-out n-in)]
             (aset w k (* (- (* (rand) 2) 1) scale)))
           (dotimes [k n-out]
             (aset b k 0))
           (.push layers #js {:w w :b b})
           (recur (inc l)))))
     #js {:sizes (.slice ^js sizes) :layers layers})))

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