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