sourcedemos/backprop/backprop.xtl
1⍝!/usr/bin/env xetal
2⍝# One step of training, every array shown: a small network (2 inputs,
3⍝# 4 tanh units, softmax over 3 classes) on six points of a spiral. The
4⍝# forward pass, the loss, each layer's gradient as one expression, the
5⍝# step; then every gradient checked against a finite difference.
6
7ⁿⁿ⁼u̲se< "NN"
8
9⍝## The batch and the starting weights, read from data/
10b ← 6 3 r̲eshape n̲umbers ⎕N̲GET "data/batch.txt"
11X ← 2 t̲ake₂ b ⍝ 6 points, x and y
12Y ← 3 ⁿⁿo̲neHot f̲loor r̲avel -1 t̲ake₂ b ⍝ their arms, one-hot
13W1 ← 3 4 r̲eshape n̲umbers ⎕N̲GET "data/w1.txt" ⍝ 2 inputs, then the bias row
14W2 ← 5 3 r̲eshape n̲umbers ⎕N̲GET "data/w2.txt" ⍝ 4 inputs, then the bias row
15lr ← 0.5
16
17⍝## Forward
18H ← ⁿⁿt̲anh X ⁿⁿd̲ense W1 ⍝ 6 x 4
19P ← ⁿⁿs̲oftmax H ⁿⁿd̲ense W2 ⍝ 6 x 3
20L ← Y ⁿⁿc̲rossEntropy P
21
22⍝## Backward: each layer's gradient, one expression each
23ᵘo̲nes ← { x → x c̲at₂ ((t̲ally x) c̲at 1) r̲eshape 1.0 } ⍝ a column of 1s: the bias input
24D2 ← (P − Y) ÷ f̲loat t̲ally X ⍝ loss by the output scores (softmax with cross-entropy)
25G2 ← (o̲\ ᵘo̲nes H) '+ '× i̲nner D2 ⍝ W2's gradient, 5 x 3
26D1 ← (D2 '+ '× i̲nner o̲\ -1 d̲rop W2) × 1.0 − H × H ⍝ back through W2 and tanh
27G1 ← (o̲\ ᵘo̲nes X) '+ '× i̲nner D1 ⍝ W1's gradient, 3 x 4
28
29⍝## The step
30V1 ← W1 − lr × G1
31V2 ← W2 − lr × G2
32ᵘl̲oss ← { w1 w2 → Y ⁿⁿc̲rossEntropy ⁿⁿs̲oftmax (ⁿⁿt̲anh X ⁿⁿd̲ense w1) ⁿⁿd̲ense w2 }
33⍝ -- end of the core ----------------------------------------------
34
35ᵘr̲ound ← { a → (f̲loat f̲loor 0.5 + 10000.0 × a) ÷ 10000.0 }
36⍝ Each point's probabilities for the three arms (its arm is 1, 1, 2, 2, 3, 3).
37ᵘr̲ound P
38⍝ The loss before and after one step.
39L c̲at V1 ᵘl̲oss V2
40⍝ The gradients (rounded): W1's, then W2's.
41ᵘr̲ound G1
42ᵘr̲ound G2
43⍝# Each analytic gradient against a central finite difference, (L(w + e)
44⍝# - L(w - e)) / 2e for every weight: the largest difference of all 27.
45e ← 0.00001
46ᵘb̲ump ← { w i → e × f̲loat (s̲hape w) r̲eshape i = r̲ange t̲ally r̲avel w }
47F1 ← (s̲hape W1) r̲eshape '{ i → (((W1 + W1 ᵘb̲ump i) ᵘl̲oss W2) − (W1 − W1 ᵘb̲ump i) ᵘl̲oss W2) ÷ 2.0 × e } e̲ach r̲ange 12
48F2 ← (s̲hape W2) r̲eshape '{ i → ((W1 ᵘl̲oss W2 + W2 ᵘb̲ump i) − W1 ᵘl̲oss W2 − W2 ᵘb̲ump i) ÷ 2.0 × e } e̲ach r̲ange 15
49gap ← ('m̲ax r̲/ r̲avel a̲bs G1 − F1) m̲ax 'm̲ax r̲/ r̲avel a̲bs G2 − F2
50gap < 1e-8