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 }