;; ported-from: src/tf.ts @ v1.0.0 (extracted-from: bot/src/rl/train.ts @ fa686ee,
;; netToTf / tfToNet / the dynamic tfjs-node import)
;;
;; The tfjs half of the trainer: MlpJson ↔ tf.Sequential conversion and one
;; gradient-descent fit. @tensorflow/tfjs-node is an OPTIONAL peer dependency
;; reachable only through load-tf's dynamic import — everything else in this
;; kit runs without it. Nothing here may grow a static require of tfjs.
(ns ardegazu.train.tf
(:require [shadow.esm :refer (dynamic-import)]))
(defonce ^:private cached (volatile! nil))
(defn load-tf
"Load tfjs-node once; a clear error when the optional peer dep is absent."
[]
(if-some [tf @cached]
(js/Promise.resolve tf)
(-> (dynamic-import "@tensorflow/tfjs-node")
(.then (fn [m]
(vreset! cached m)
m))
(.catch (fn [err]
(throw (js/Error.
(str "@tensorflow/tfjs-node is not installed — it is an optional peer dependency needed only for training (fitNet/loadTf): "
(.-message ^js err)))))))))
(defn net-to-tf
"Build a compiled tf.Sequential carrying the net's weights (kernels transposed)."
[^js tf ^js net lr]
(let [m (.sequential tf)
sizes (.-sizes net)
n-sizes (.-length sizes)]
(loop [l 1]
(when (< l n-sizes)
(.add m (.dense (.-layers tf)
#js {:units (aget sizes l)
:inputShape (if (== l 1) #js [(aget sizes 0)] js/undefined)
:activation (if (< l (dec n-sizes)) "relu" "linear")}))
(recur (inc l))))
(.compile m #js {:optimizer (.adam (.-train tf) lr) :loss "meanSquaredError"})
;; load weights: our w is row-major out×in; tf kernels are [in, out]
(let [tensors #js []
layers (.-layers net)]
(dotimes [l (.-length layers)]
(let [n-in (aget sizes l)
n-out (aget sizes (inc l))
layer (aget layers l)
w (.-w ^js layer)
kernel (js/Float32Array. (* n-in n-out))]
(dotimes [o n-out]
(dotimes [i n-in]
(aset kernel (+ (* i n-out) o) (aget w (+ (* o n-in) i)))))
(.push tensors (.tensor2d tf kernel #js [n-in n-out]))
(.push tensors (.tensor1d tf (js/Float32Array.from (.-b ^js layer))))))
(.setWeights m tensors))
m))
(defn tf-to-net
"Read a Sequential's weights back into the plain-JSON MlpJson shape."
[_tf ^js m sizes]
(let [weights (.getWeights m)
layers #js []
n (dec (.-length ^js sizes))]
(dotimes [l n]
(let [^js kernel-t (aget weights (* l 2))
^js bias-t (aget weights (inc (* l 2)))
kernel (.dataSync kernel-t)
bias (.dataSync bias-t)
n-in (aget sizes l)
n-out (aget sizes (inc l))
w (js/Array. (* n-out n-in))]
(dotimes [o n-out]
(dotimes [i n-in]
(aset w (+ (* o n-in) i) (aget kernel (+ (* i n-out) o)))))
(.push layers #js {:w w :b (js/Array.from bias)})))
#js {:sizes (.slice ^js sizes) :layers layers}))
(defn fit-net
"netToTf → fit → tfToNet, with every tensor disposed. Resolves to the new net."
[^js net xs ys ^js opts]
(-> (load-tf)
(.then
(fn [^js tf]
(let [m (net-to-tf tf net (.-lr opts))
x-t (.tensor2d tf xs #js [(.-length ^js xs) (aget (.-sizes net) 0)])
y-t (.tensor2d tf ys #js [(.-length ^js ys) 1])
batch-size (let [b (.-batchSize opts)] (if (nil? b) 256 b))]
(-> (.fit ^js m x-t y-t #js {:epochs (.-epochs opts)
:batchSize batch-size
:shuffle true
:verbose 0})
(.then (fn [_] (tf-to-net tf m (.-sizes net))))
(.finally (fn []
(.dispose x-t)
(.dispose y-t)
(.dispose ^js m)))))))))