programdemos/ternary-net/ternary-net.xtl

A 1.58-bit network: a tiny classifier (which of three spiral arms a point is on, 2 -> 16 -> 16 -> 3) run with its weights stored four ways: 32-bit and 16-bit floats (simulated by rounding), 8-bit integers, and ternary -1 / 0 / +1 (log2 3 = 1.58 bits each), where a layer is additions and subtractions of its inputs and one scale. The same points go through all four; the measures compare them.

source

g : Int

value · line 11

The map's side (g by g points) and the ternary threshold (a weight is kept as +1 or -1 when its size passes t times the layer's mean).

g ← 12

t : Float

value · line 12
t ← 0.5
Used in: q

The weights, read from data/ (just ternary-train)

fp : Float

value · line 18

fp: trained in full precision; qa: fine-tuned for ternary weights. Each is 17 x 3 x 16: for each row (16 inputs, padded with 0, then the bias), layer and output (padded with 0).

fp ← 17 3 16 r̲eshape n̲umbers ⎕N̲GET "data/fp.txt"

qa : Float

value · line 19
qa ← 17 3 16 r̲eshape n̲umbers ⎕N̲GET "data/qa.txt"
Used in: m

tp : Float

value · line 21

The test points, one per line: x, y and the arm (1, 2 or 3).

tp ← n̲umbers ⎕N̲GET "data/test.txt"
Used in: tp, px, py, pl

tp : Float

value · line 22
tp ← o̲\ (((t̲ally tp) d̲iv 3) c̲at 3) r̲eshape tp

px : Float

value · line 23
px ← 1 s̲elect tp
Used in: test

py : Float

value · line 24
py ← 2 s̲elect tp
Used in: test

pl : Float

value · line 25
pl ← 3 s̲elect tp
Used in: ᵘm̲easures

m : Float

value · line 28
m ← qa
Used in: w, m32, m16, m8, m2

The inputs: points padded to 16 numbers

e1 : Float

value · line 32

Each point (x, y) becomes x e1 + y e2: a row of 16 with the rest 0, so every layer is the same 16 wide.

e1 ← 16 t̲ake 1.0
Used in: ᵘp̲ad

e2 : Float

value · line 33
e2 ← 16 t̲ake 0.0 1.0
Used in: ᵘp̲ad

ᵘp̲ad : Float -> Float -> Float

function · line 34
ᵘp̲ad ← { x y → (x '× t̲able e1) + y '× t̲able e2 }
Used in: grid, test

c : Float

value · line 36

The map: g by g points over -1.1 .. 1.1, row by row, y downwards.

c ← -1.1 + 2.2 × (0.5 + f̲loat o̲ffsets g) ÷ f̲loat g
Used in: gx, gy

gx : Float

value · line 37
gx ← r̲avel (o̲ffsets g) 'r̲ight t̲able c
Used in: grid

gy : Float

value · line 38
gy ← r̲avel (r̲ev c) 'l̲eft t̲able o̲ffsets g
Used in: grid

grid : Float

value · line 39
grid ← gx ᵘp̲ad gy
Used in: y32, y16, y8, y2

test : Float

value · line 40
test ← px ᵘp̲ad py
Used in: ᵘm̲easures

The network: three layers of one array

ᵘl̲ayer : Num a => a -> a -> a

function · line 43

A layer plane p (17 x 16): 16 rows of weights, then the biases.

ᵘ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 }
Used in: ᵘn̲et

ᵘr̲elu : Float -> Float

function · line 44
ᵘr̲elu ← { x → 0.0 m̲ax x }
Used in: ᵘn̲et

ᵘn̲et : Float -> Float -> Float

function · line 45
ᵘn̲et ← { x k →
  h ← ᵘr̲elu x ᵘl̲ayer 1 s̲elect₂ k
  h ← ᵘr̲elu h ᵘl̲ayer 2 s̲elect₂ k
  3 t̲ake₂ h ᵘl̲ayer 3 s̲elect₂ k
}
Used in: y32, y16, y8, y2, ᵘm̲easures

The formats: the weights rounded four ways

w : Float

value · line 53

w is 16 x 3 x 16 (input, layer, output); a value per layer is spread back over its layer; the biases stay as they are.

w ← 16 t̲ake m
Used in: q, a, m32, m16, m8

nw : Float

value · line 54
nw ← 32.0 256.0 48.0
Used in: ᵘt̲ern

ᵘl̲ayers : a -> a

function · line 55
ᵘl̲ayers ← { v → (o̲ffsets 16) 'r̲ight t̲able v 'l̲eft t̲able o̲ffsets 16 }
Used in: ᵘi̲nt8, ᵘt̲ern, a

ᵘw̲ith : a -> a -> a

function · line 56
ᵘw̲ith ← { q k → q c̲at -1 t̲ake k }
Used in: m32, m16, m8, m2

ᵘf̲loats : Float -> Float -> Float

function · line 58

Floats: keep b bits after the leading one (23 for FP32, 10 FP16).

ᵘf̲loats ← { b x →
  s ← 2.0 ^ (f̲loat f̲loor (l̲og 0.000000001 m̲ax a̲bs x) ÷ l̲og 2.0) − b
  s × f̲loat f̲loor 0.5 + x ÷ s
}
Used in: m32, m16

ᵘi̲nt8 : Float -> Float

function · line 63

INT8: each layer's largest weight becomes 127.

ᵘi̲nt8 ← { x →
  s ← ᵘl̲ayers ('m̲ax r̲/₁₃ a̲bs x) ÷ 127.0
  s × f̲loat f̲loor 0.5 + x ÷ s
}
Used in: m8

ᵘt̲ern : Float -> Float -> Float

function · line 69

Ternary: +1 or -1 where the size passes t times the layer's mean size, else 0; the scale a is the mean size of the kept weights.

ᵘt̲ern ← { k x →
  s ← k × ᵘl̲ayers ('+ r̲/₁₃ a̲bs x) ÷ nw
  (f̲loat x > s) − f̲loat x < n̲eg s
}
Used in: q

q : Float

value · line 73
q ← t ᵘt̲ern w

a : Float

value · line 74
a ← ᵘl̲ayers ('+ r̲/₁₃ (a̲bs w) × a̲bs q) ÷ 1.0 m̲ax '+ r̲/₁₃ a̲bs q
Used in: m2

m32 : Float

value · line 75
m32 ← (23.0 ᵘf̲loats w) ᵘw̲ith m

m16 : Float

value · line 76
m16 ← (10.0 ᵘf̲loats w) ᵘw̲ith m

m8 : Float

value · line 77
m8 ← (ᵘi̲nt8 w) ᵘw̲ith m

m2 : Float

value · line 78
m2 ← (a × q) ᵘw̲ith m

Additions only: a ternary layer for one input h

ᵘa̲dds : Float -> Float -> Float

function · line 82

Sum the inputs where the weight is +1, subtract those where it is -1, then one multiply by the layer's scale.

ᵘ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 }

The maps: every map point through each format

y32 : Float

value · line 84
y32 ← grid ᵘn̲et m32

y16 : Float

value · line 85
y16 ← grid ᵘn̲et m16

y8 : Float

value · line 86
y8 ← grid ᵘn̲et m8

y2 : Float

value · line 87
y2 ← grid ᵘn̲et m2

The measures

ᵘa̲rm : Num a => a -> Float

function · line 90

The arm each point is given (the largest of its three outputs).

ᵘ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 }

ᵘm̲ean : Float -> Float

function · line 91
ᵘm̲ean ← { x → ('+ r̲/ r̲avel x) ÷ f̲loat t̲ally r̲avel x }
Used in: ᵘm̲easures

ᵘm̲easures : Float -> Float -> Float

function · line 94

For a format's weights k and its map y: accuracy on the test points, agreement with FP32 over the map, mean output error.

ᵘm̲easures ← { k y →
  acc ← ᵘm̲ean f̲loat pl = ᵘa̲rm test ᵘn̲et k
  agree ← ᵘm̲ean f̲loat (ᵘa̲rm y) = ᵘa̲rm y32
  err ← ᵘm̲ean a̲bs y − y32
  acc c̲at agree c̲at err
}

ᵘr̲ound : Float -> Float

function · line 102
ᵘr̲ound ← { x → (f̲loat f̲loor 0.5 + x × 1000.0) ÷ 1000.0 }