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"