programdemos/net-macro/train-it.xtl

Training a network from its one-line spec: the Net macro library writes, from "2 16 relu 16 relu 3 softmax", the network's backward pass and an Adam step (net:t_rain<) and the starting state (net:s_tate<); p_ower repeats the step. The task is net-macro's: which of three spiral arms is a point on? Here the 300 training points are made in X_eTaL, and the weights start small and random.

source · imports nn: libs/NN/src/NN.xtl; net: libs/Net/src/Net.xtlm

The spiral: 100 points on each of three arms

n : Int

value · line 13
n ← 300
Used in: i, a, X

i : Int

value · line 14
i ← o̲ffsets n
Used in: arm, t

arm : Int

value · line 15
arm ← i d̲iv 100

t : Float

value · line 16
t ← (0.5 + f̲loat i m̲od 100) ÷ 100.0
Used in: a, X

a : Float

value · line 17
a ← (2.0944 × f̲loat arm) + (5.5 × t) + 0.1 × (f̲loat (r̲oll! n r̲eshape 1001) − 501) ÷ 500.0
Used in: X

X : Float

value · line 18
X ← o̲\ (2 c̲at n) r̲eshape ((0.1 + 0.9 × t) × c̲os a) c̲at (0.1 + 0.9 × t) × s̲in a

Y : Float

value · line 19
Y ← 3 ⁿⁿo̲neHot 1 + arm

lr : Float

value · line 20
lr ← 0.02
Used in: ᵘs̲tep

The network and its training, from one spec

ᵘs̲tep : (Float, Float, Float, Float, Float, Float, Float, Float, Float, Float) -> (Float, Float, Float, Float, Float, Float, Float, Float, Float, Float)

function · line 23
"u:s_tep X Y lr" ⁿᵉᵗt̲rain< "2 16 relu 16 relu 3 softmax"
ⁿᵉᵗt̲rain< expands to
ᵘs̲tep ← { (W1, W2, W3, M1, M2, M3, V1, V2, V3, k) →
  A0 ← X
  A1 ← ⁿⁿr̲elu A0 ⁿⁿd̲ense W1
  A2 ← ⁿⁿr̲elu A1 ⁿⁿd̲ense W2
  A3 ← ⁿⁿs̲oftmax A2 ⁿⁿd̲ense W3
  D3 ← (A3 − Y) ÷ f̲loat t̲ally X
  D2 ← (D3 '+ '× i̲nner o̲\ -1 d̲rop W3) × f̲loat A2 > 0.0
  D1 ← (D2 '+ '× i̲nner o̲\ -1 d̲rop W2) × f̲loat A1 > 0.0
  G1 ← (o̲\ A0 c̲at₂ ((t̲ally A0) c̲at 1) r̲eshape 1.0) '+ '× i̲nner D1
  G2 ← (o̲\ A1 c̲at₂ ((t̲ally A1) c̲at 1) r̲eshape 1.0) '+ '× i̲nner D2
  G3 ← (o̲\ A2 c̲at₂ ((t̲ally A2) c̲at 1) r̲eshape 1.0) '+ '× i̲nner D3
  k ← 1.0 + k
  M1 ← (0.9 × M1) + 0.1 × G1
  V1 ← (0.999 × V1) + 0.001 × G1 × G1
  W1 ← W1 − lr × (M1 ÷ 1.0 − 0.9 ^ k) ÷ 0.00000001 + (V1 ÷ 1.0 − 0.999 ^ k) ^ 0.5
  M2 ← (0.9 × M2) + 0.1 × G2
  V2 ← (0.999 × V2) + 0.001 × G2 × G2
  W2 ← W2 − lr × (M2 ÷ 1.0 − 0.9 ^ k) ÷ 0.00000001 + (V2 ÷ 1.0 − 0.999 ^ k) ^ 0.5
  M3 ← (0.9 × M3) + 0.1 × G3
  V3 ← (0.999 × V3) + 0.001 × G3 × G3
  W3 ← W3 − lr × (M3 ÷ 1.0 − 0.9 ^ k) ÷ 0.00000001 + (V3 ÷ 1.0 − 0.999 ^ k) ^ 0.5
  (W1, W2, W3, M1, M2, M3, V1, V2, V3, k)
}
Used in: %2

ᵘr̲andom : Int -> Float

function · line 25

Small random weights of a shape (inputs + 1 by outputs), then the state.

ᵘr̲andom ← { sh → sh r̲eshape (f̲loat (r̲oll! ('× r̲/ sh) r̲eshape 2001) − 1001) ÷ 2000.0 }
Used in: w1, w2, w3

w1 : Float

value · line 26
w1 ← ᵘr̲andom 3 16

w2 : Float

value · line 27
w2 ← ᵘr̲andom 17 16

w3 : Float

value · line 28
w3 ← ᵘr̲andom 17 3

s0 : (Float, Float, Float, Float, Float, Float, Float, Float, Float, Float)

value · line 29
s0 ← @ ⁿᵉᵗs̲tate< "w1 w2 w3"
ⁿᵉᵗs̲tate< expands to
((w1, w2, w3, 0.0 × w1, 0.0 × w2, 0.0 × w3, 0.0 × w1, 0.0 × w2, 0.0 × w3, 0.0))
Used in: %2

ᵘs̲tart : Float -> Float

function · line 30
ᵘs̲tart ← "2 16 relu 16 relu 3 softmax" ⁿᵉᵗn̲etwork< "w1 w2 w3"
ⁿᵉᵗn̲etwork< expands to
({ x → ⁿⁿs̲oftmax (ⁿⁿr̲elu (ⁿⁿr̲elu x ⁿⁿd̲ense w1) ⁿⁿd̲ense w2) ⁿⁿd̲ense w3 })

%2 : Float

value · line 35

Four hundred steps; then the trained weights, out of the state, and the network on them.

(w1, w2, w3, _, _, _, _, _, _, _) ← 400 'ᵘs̲tep p̲ower s0

w1 : Float

value · line 35

Four hundred steps; then the trained weights, out of the state, and the network on them.

(w1, w2, w3, _, _, _, _, _, _, _) ← 400 'ᵘs̲tep p̲ower s0

w2 : Float

value · line 35

Four hundred steps; then the trained weights, out of the state, and the network on them.

(w1, w2, w3, _, _, _, _, _, _, _) ← 400 'ᵘs̲tep p̲ower s0

w3 : Float -> Float

value · line 35

Four hundred steps; then the trained weights, out of the state, and the network on them.

(w1, w2, w3, _, _, _, _, _, _, _) ← 400 'ᵘs̲tep p̲ower s0

ᵘt̲rained : Float -> Float

function · line 36
ᵘt̲rained ← "2 16 relu 16 relu 3 softmax" ⁿᵉᵗn̲etwork< "w1 w2 w3"
ⁿᵉᵗn̲etwork< expands to
({ x → ⁿⁿs̲oftmax (ⁿⁿr̲elu (ⁿⁿr̲elu x ⁿⁿd̲ense w1) ⁿⁿd̲ense w2) ⁿⁿd̲ense w3 })

ᵘr̲ound

function · line 38

The loss and the share of the 300 points right, before and after.

ᵘr̲ound ← { a → (f̲loat f̲loor 0.5 + 10000.0 × a) ÷ 10000.0 }