sourcelibs/NN/src/NN.xtl

1⍝# NN: neural-network building blocks -- activations, softmax by row, dense layers, one-hot, loss, accuracy 2⍝# 3⍝# Import it with an alias of your choice: "nn:" u_se< "NN". 4⍝# Put libs/NN/src on XETAL_PATH ("just path"); the reference is libs/NN/docs. 5⍝# Names with l: are exported; the others are private to this file. 6⍝# 7⍝# A batch is a matrix, one example per row; a layer's weights are a 8⍝# matrix with one row per input and one column per output, and its 9⍝# bias is one more row at the bottom (so a layer is one array, and 10⍝# x ⁿⁿd̲ense wb takes two arguments). Activations work item by item 11⍝# on any shape. Softmax, log-softmax and argmax work along the last 12⍝# axis of any array: a vector is one example, a matrix one per row, a 13⍝# rank-3 array a batch of sequences. 14 15⍝## Helpers (private) 16 17⍝# The row values r (one per row of m) spread across m's columns. 18ʰs̲pread ← { r m → o̲\ ((-1 t̲ake s̲hape m) c̲at t̲ally r) r̲eshape r } 19 20⍝# m as a matrix of its last-axis rows (a vector becomes one row). 21ʰr̲ows ← { m → (('× r̲/ -1 d̲rop s̲hape m) c̲at -1 t̲ake s̲hape m) r̲eshape m } 22 23⍝## Activations 24 25⍝# r̲elu x: max(x, 0), item by item. 26⍝# >> "nn:" u_se< "NN" 27⍝# >> nn:r_elu -1.5 0.0 2.0 28⍝# 0.0 0.0 2.0 29ˡr̲elu ← { x → x m̲ax 0.0 } 30 31⍝# a l̲eaky x: leaky ReLU, x where positive and a times x elsewhere. 32⍝# >> "nn:" u_se< "NN" 33⍝# >> 0.1 nn:l_eaky -2.0 3.0 34⍝# -0.2 3.0 35ˡl̲eaky ← { a x → (x m̲ax 0.0) + a × x m̲in 0.0 } 36 37⍝# s̲igmoid x: 1 / (1 + e^-x), item by item. 38⍝# >> "nn:" u_se< "NN" 39⍝# >> nn:s_igmoid 0.0 2.0 40⍝# 0.5 0.8807970779778823 41ˡs̲igmoid ← { x → 1.0 ÷ 1.0 + e̲xp n̲eg x } 42 43⍝# t̲anh x: the hyperbolic tangent, as 2 s_igmoid(2x) - 1. 44⍝# >> "nn:" u_se< "NN" 45⍝# >> nn:t_anh 0.0 46⍝# 0.0 47ˡt̲anh ← { x → (2.0 × ˡs̲igmoid 2.0 × x) − 1.0 } 48 49⍝## Softmax and argmax: along the last axis 50 51⍝# s̲oftmax m: each row into probabilities; the row's largest is taken 52⍝# off first (the same result, no overflow), the row's sum divides. 53⍝# >> "nn:" u_se< "NN" 54⍝# >> nn:s_oftmax 1.0 2.0 3.0 55⍝# 0.09003057317038046 0.24472847105479764 0.6652409557748218 56ˡs̲oftmax ← { a → 57 m ← ʰr̲ows a 58 e ← e̲xp m − ('m̲ax r̲/₂ m) ʰs̲pread m 59 (s̲hape a) r̲eshape e ÷ ('+ r̲/₂ e) ʰs̲pread e 60} 61 62⍝# l̲ogSoftmax m: the logarithm of each row's softmax, computed stably. 63ˡl̲ogSoftmax ← { a → 64 m ← ʰr̲ows a 65 z ← m − ('m̲ax r̲/₂ m) ʰs̲pread m 66 (s̲hape a) r̲eshape z − (l̲og '+ r̲/₂ e̲xp z) ʰs̲pread z 67} 68 69⍝## Layers 70 71⍝# x d̲ense wb: the layer x W + b for a batch x (one row each) and wb, 72⍝# the weights W with the bias b as one more row at the bottom. 73⍝# >> "nn:" u_se< "NN" 74⍝# >> (2 2 r_eshape 1.0 2.0 3.0 4.0) nn:d_ense 3 2 r_eshape 1.0 0.0 0.0 2.0 10.0 20.0 75⍝# 11.0 24.0 76⍝# 13.0 28.0 77ˡd̲ense ← { x wb → 78 (x '+ '× i̲nner -1 d̲rop wb) + (o̲ffsets t̲ally x) 'r̲ight t̲able f̲irst -1 t̲ake wb 79} 80 81⍝## Labels 82 83⍝# a̲rgmax a: for each row (along the last axis), the position (from 84⍝# 1) of its largest item (the first, on a tie); a vector gives one Int. 85⍝# >> "nn:" u_se< "NN" 86⍝# >> nn:a_rgmax 2 3 r_eshape 1.0 2.0 3.0 -1.0 0.0 5.0 87⍝# 3 3 88ˡa̲rgmax ← { a → 89 m ← ʰr̲ows a 90 n ← f̲irst -1 t̲ake s̲hape m 91 hit ← f̲loat m = ('m̲ax r̲/₂ m) ʰs̲pread m 92 (-1 d̲rop s̲hape a) r̲eshape f̲loor 'm̲in r̲/₂ (hit × (o̲ffsets t̲ally m) 'r̲ight t̲able f̲loat r̲ange n) + (1.0 − hit) × f̲loat n + 1 93} 94 95⍝# k o̲neHot y: the labels y (1 to k) as rows of k Floats, 1.0 at the label. 96⍝# >> "nn:" u_se< "NN" 97⍝# >> 3 nn:o_neHot 2 1 3 98⍝# 0.0 1.0 0.0 99⍝# 1.0 0.0 0.0 100⍝# 0.0 0.0 1.0 101ˡo̲neHot ← { k y → f̲loat y '= t̲able r̲ange k } 102 103⍝## Loss and accuracy 104 105⍝# y c̲rossEntropy p: the mean over rows of -sum(y log p), for one-hot 106⍝# (or probability) rows y and predicted probabilities p; p is clamped 107⍝# at 1e-12 first, so a probability of 0 costs 27.6, not a log error. 108⍝# >> "nn:" u_se< "NN" 109⍝# >> (2 nn:o_neHot 1 2) nn:c_rossEntropy 2 2 r_eshape 0.9 0.1 0.2 0.8 110⍝# 0.164252033486018 111ˡc̲rossEntropy ← { y p → ('+ r̲/ '+ r̲/₂ n̲eg y × l̲og p m̲ax 1e-12) ÷ f̲loat t̲ally p } 112 113⍝# y m̲se p: the mean squared error over every item. 114⍝# >> "nn:" u_se< "NN" 115⍝# >> 1.0 2.0 nn:m_se 1.0 4.0 116⍝# 2.0 117ˡm̲se ← { y p → ('+ r̲/ r̲avel (y − p) ^ 2) ÷ f̲loat t̲ally r̲avel p } 118 119⍝# y a̲ccuracy p: the fraction of rows of p whose argmax is the label y. 120⍝# >> "nn:" u_se< "NN" 121⍝# >> 3 2 nn:a_ccuracy 2 3 r_eshape 1.0 2.0 3.0 -1.0 0.0 5.0 122⍝# 0.5 123ˡa̲ccuracy ← { y p → ('+ r̲/ f̲loat y = ˡa̲rgmax p) ÷ f̲loat t̲ally y }