programdemos/train-live/train-live.xtl

Training in X_eTaL: a network (2 inputs, 16 tanh units, softmax over 3) learns which of three spiral arms a point is on. The points are made here; the backward pass is the backprop microscope's; Adam, written out, moves the weights; p_ower repeats the step.

source · imports nn: libs/NN/src/NN.xtl

The spiral: 100 points on each of three arms

n : Int

value · line 10
n ← 300
Used in: i, a, X, ᵘg̲rad

i : Int

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

arm : Int

value · line 12
arm ← i d̲iv 100
Used in: a, Y, ᵘr̲ight

t : Float

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

a : Float

value · line 14
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 15
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 16
Y ← 3 ⁿⁿo̲neHot 1 + arm

The network: W1 (3 x 16), W2 (17 x 3), each bias its last row

ᵘo̲nes : Float -> Float

function · line 19
ᵘo̲nes ← { x → x c̲at₂ ((t̲ally x) c̲at 1) r̲eshape 1.0 }
Used in: ᵘg̲rad

ᵘf̲orward : Float -> Float -> Float

function · line 20
ᵘf̲orward ← { W1 W2 → ⁿⁿs̲oftmax (ⁿⁿt̲anh X ⁿⁿd̲ense W1) ⁿⁿd̲ense W2 }

ᵘg̲rad : Float -> Float -> (Float, Float)

function · line 23

The gradients of the loss by W1 and by W2, a pair (the backprop microscope's four lines).

ᵘg̲rad ← { W1 W2 →
  H ← ⁿⁿt̲anh X ⁿⁿd̲ense W1
  D2 ← ((ⁿⁿs̲oftmax H ⁿⁿd̲ense W2) − Y) ÷ f̲loat n
  D1 ← (D2 '+ '× i̲nner o̲\ -1 d̲rop W2) × 1.0 − H × H
  ((o̲\ ᵘo̲nes X) '+ '× i̲nner D1, (o̲\ ᵘo̲nes H) '+ '× i̲nner D2)
}
Used in: ᵘa̲dam

Adam

lr : Float

value · line 34

The learning rate.

lr ← 0.02
Used in: ᵘa̲dam

ᵘm̲ove : (Float, Float) -> Float -> Float

function · line 36

How far Adam moves a weight array, from its averages (m, v) at step k.

ᵘm̲ove ← { (m, v) k → (m ÷ 1.0 − 0.9 ^ k) ÷ 0.00000001 + (v ÷ 1.0 − 0.999 ^ k) ^ 0.5 }
Used in: ᵘa̲dam

ᵘa̲dam : (Float, Float, Float, Float, Float, Float, Float) -> (Float, Float, Float, Float, Float, Float, Float)

function · line 38

One step of Adam.

ᵘa̲dam ← { (W1, W2, M1, M2, V1, V2, k) →
  k ← 1.0 + k
  (G1, G2) ← W1 ᵘg̲rad W2
  M1 ← (0.9 × M1) + 0.1 × G1
  M2 ← (0.9 × M2) + 0.1 × G2
  V1 ← (0.999 × V1) + 0.001 × G1 × G1
  V2 ← (0.999 × V2) + 0.001 × G2 × G2
  (W1 − lr × (M1, V1) ᵘm̲ove k, W2 − lr × (M2, V2) ᵘm̲ove k, M1, M2, V1, V2, k)
}
Used in: s1, s2

w0 : Float

value · line 48

The starting state: small random weights, zero averages, step 0.

w0 ← (f̲loat (r̲oll! 99 r̲eshape 2001) − 1001) ÷ 1000.0
Used in: W1, W2

W1 : Float

value · line 49
W1 ← 3 16 r̲eshape 48 t̲ake w0

W2 : Float

value · line 50
W2 ← 17 3 r̲eshape 48 d̲rop w0

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

value · line 51
s0 ← (W1, W2, 0.0 × W1, 0.0 × W2, 0.0 × W1, 0.0 × W2, 0.0)

ᵘl̲oss : (Any a, Any b, Any c, Any d, Any e) => (Float, Float, a, b, c, d, e) -> Float

function · line 55

The loss, and the share of points right, of a state.

ᵘl̲oss ← { (W1, W2, _, _, _, _, _) → Y ⁿⁿc̲rossEntropy W1 ᵘf̲orward W2 }

ᵘr̲ight : (Any a, Any b, Any c, Any d, Any e) => (Float, Float, a, b, c, d, e) -> Float

function · line 56
ᵘr̲ight ← { (W1, W2, _, _, _, _, _) → (1 + arm) ⁿⁿa̲ccuracy W1 ᵘf̲orward W2 }

s1 : (Float, Float, Float, Float, Float, Float, Float)

value · line 59

The loss and the share of points right: before, after 100 steps, after 400.

s1 ← 100 'ᵘa̲dam p̲ower s0

s2 : (Float, Float, Float, Float, Float, Float, Float)

value · line 60
s2 ← 300 'ᵘa̲dam p̲ower s1

ᵘr̲ound : Float -> Float

function · line 61
ᵘr̲ound ← { a → (f̲loat f̲loor 0.5 + 1000.0 × a) ÷ 1000.0 }