train-kit / src / ardegazu / train / tf.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
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
;; 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)))))))))

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