sourcedemos/ternary-net/ternary-net.xtl
1⍝!/usr/bin/env xetal
2⍝# A 1.58-bit network: a tiny classifier (which of three spiral arms a
3⍝# point is on, 2 -> 16 -> 16 -> 3) run with its weights stored four
4⍝# ways: 32-bit and 16-bit floats (simulated by rounding), 8-bit
5⍝# integers, and ternary -1 / 0 / +1 (log2 3 = 1.58 bits each), where
6⍝# a layer is additions and subtractions of its inputs and one scale.
7⍝# The same points go through all four; the measures compare them.
8
9⍝# The map's side (g by g points) and the ternary threshold (a weight
10⍝# is kept as +1 or -1 when its size passes t times the layer's mean).
11g ← 12
12t ← 0.5
13
14⍝## The weights, read from data/ (just ternary-train)
15⍝# fp: trained in full precision; qa: fine-tuned for ternary
16⍝# weights. Each is 17 x 3 x 16: for each row (16 inputs, padded
17⍝# with 0, then the bias), layer and output (padded with 0).
18fp ← 17 3 16 r̲eshape n̲umbers ⎕N̲GET "data/fp.txt"
19qa ← 17 3 16 r̲eshape n̲umbers ⎕N̲GET "data/qa.txt"
20⍝# The test points, one per line: x, y and the arm (1, 2 or 3).
21tp ← n̲umbers ⎕N̲GET "data/test.txt"
22tp ← o̲\ (((t̲ally tp) d̲iv 3) c̲at 3) r̲eshape tp
23px ← 1 s̲elect tp
24py ← 2 s̲elect tp
25pl ← 3 s̲elect tp
26⍝ -- end of the weights ----------------------------------------------
27
28m ← qa
29⍝## The inputs: points padded to 16 numbers
30⍝# Each point (x, y) becomes x e1 + y e2: a row of 16 with the rest 0,
31⍝# so every layer is the same 16 wide.
32e1 ← 16 t̲ake 1.0
33e2 ← 16 t̲ake 0.0 1.0
34ᵘp̲ad ← { x y → (x '× t̲able e1) + y '× t̲able e2 }
35⍝# The map: g by g points over -1.1 .. 1.1, row by row, y downwards.
36c ← -1.1 + 2.2 × (0.5 + f̲loat o̲ffsets g) ÷ f̲loat g
37gx ← r̲avel (o̲ffsets g) 'r̲ight t̲able c
38gy ← r̲avel (r̲ev c) 'l̲eft t̲able o̲ffsets g
39grid ← gx ᵘp̲ad gy
40test ← px ᵘp̲ad py
41⍝## The network: three layers of one array
42⍝# A layer plane p (17 x 16): 16 rows of weights, then the biases.
43ᵘl̲ayer ← { x p → (x '+ '× i̲nner 16 t̲ake p) + (o̲ffsets t̲ally x) 'r̲ight t̲able f̲irst -1 t̲ake p }
44ᵘr̲elu ← { x → 0.0 m̲ax x }
45ᵘn̲et ← { x k →
46 h ← ᵘr̲elu x ᵘl̲ayer 1 s̲elect₂ k
47 h ← ᵘr̲elu h ᵘl̲ayer 2 s̲elect₂ k
48 3 t̲ake₂ h ᵘl̲ayer 3 s̲elect₂ k
49}
50⍝## The formats: the weights rounded four ways
51⍝# w is 16 x 3 x 16 (input, layer, output); a value per layer is
52⍝# spread back over its layer; the biases stay as they are.
53w ← 16 t̲ake m
54nw ← 32.0 256.0 48.0
55ᵘl̲ayers ← { v → (o̲ffsets 16) 'r̲ight t̲able v 'l̲eft t̲able o̲ffsets 16 }
56ᵘw̲ith ← { q k → q c̲at -1 t̲ake k }
57⍝# Floats: keep b bits after the leading one (23 for FP32, 10 FP16).
58ᵘf̲loats ← { b x →
59 s ← 2.0 ^ (f̲loat f̲loor (l̲og 0.000000001 m̲ax a̲bs x) ÷ l̲og 2.0) − b
60 s × f̲loat f̲loor 0.5 + x ÷ s
61}
62⍝# INT8: each layer's largest weight becomes 127.
63ᵘi̲nt8 ← { x →
64 s ← ᵘl̲ayers ('m̲ax r̲/₁₃ a̲bs x) ÷ 127.0
65 s × f̲loat f̲loor 0.5 + x ÷ s
66}
67⍝# Ternary: +1 or -1 where the size passes t times the layer's mean
68⍝# size, else 0; the scale a is the mean size of the kept weights.
69ᵘt̲ern ← { k x →
70 s ← k × ᵘl̲ayers ('+ r̲/₁₃ a̲bs x) ÷ nw
71 (f̲loat x > s) − f̲loat x < n̲eg s
72}
73q ← t ᵘt̲ern w
74a ← ᵘl̲ayers ('+ r̲/₁₃ (a̲bs w) × a̲bs q) ÷ 1.0 m̲ax '+ r̲/₁₃ a̲bs q
75m32 ← (23.0 ᵘf̲loats w) ᵘw̲ith m
76m16 ← (10.0 ᵘf̲loats w) ᵘw̲ith m
77m8 ← (ᵘi̲nt8 w) ᵘw̲ith m
78m2 ← (a × q) ᵘw̲ith m
79⍝## Additions only: a ternary layer for one input h
80⍝# Sum the inputs where the weight is +1, subtract those where it is
81⍝# -1, then one multiply by the layer's scale.
82ᵘa̲dds ← { h k → ('+ r̲/ (h 'l̲eft t̲able o̲ffsets 16) × f̲loat k = 1.0) − '+ r̲/ (h 'l̲eft t̲able o̲ffsets 16) × f̲loat k = -1.0 }
83⍝## The maps: every map point through each format
84y32 ← grid ᵘn̲et m32
85y16 ← grid ᵘn̲et m16
86y8 ← grid ᵘn̲et m8
87y2 ← grid ᵘn̲et m2
88⍝## The measures
89⍝# The arm each point is given (the largest of its three outputs).
90ᵘa̲rm ← { y → '+ r̲/₂ (f̲loat y = ('m̲ax r̲/₂ y) 'l̲eft t̲able o̲ffsets 3) × (o̲ffsets t̲ally y) 'r̲ight t̲able 1.0 2.0 3.0 }
91ᵘm̲ean ← { x → ('+ r̲/ r̲avel x) ÷ f̲loat t̲ally r̲avel x }
92⍝# For a format's weights k and its map y: accuracy on the test
93⍝# points, agreement with FP32 over the map, mean output error.
94ᵘm̲easures ← { k y →
95 acc ← ᵘm̲ean f̲loat pl = ᵘa̲rm test ᵘn̲et k
96 agree ← ᵘm̲ean f̲loat (ᵘa̲rm y) = ᵘa̲rm y32
97 err ← ᵘm̲ean a̲bs y − y32
98 acc c̲at agree c̲at err
99}
100⍝ -- end of the core ------------------------------------
101
102ᵘr̲ound ← { x → (f̲loat f̲loor 0.5 + x × 1000.0) ÷ 1000.0 }
103⍝ The layer-2 weights as ternary glyphs: - 0 +.
104(1 + f̲loor 1.0 + 2 s̲elect₂ q) s̲elect "-.+"
105⍝ Accuracy on the test points, agreement with FP32 over the map, and
106⍝ the mean output error, for FP32, FP16, INT8 and ternary.
107ᵘr̲ound (m32 ᵘm̲easures y32) c̲at (m16 ᵘm̲easures y16) c̲at (m8 ᵘm̲easures y8) c̲at m2 ᵘm̲easures y2
108⍝ Weight storage in bits, and the adds the ternary layers need (the
109⍝ kept weights) against the 336 multiply-adds of the others.
110(32 × 336) c̲at (16 × 336) c̲at (8 × 336) c̲at ᵘr̲ound 336.0 × (l̲og 3.0) ÷ l̲og 2.0
111'+ r̲/ r̲avel a̲bs q
112⍝ The map's arms, FP32 and then ternary.
113(g c̲at g) r̲eshape (f̲loor ᵘa̲rm y32) s̲elect "abc"
114(g c̲at g) r̲eshape (f̲loor ᵘa̲rm y2) s̲elect "abc"