;; 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})))