programdemos/moe-router/moe-router.xtl

A mixture-of-experts router: each token of a sentence is scored against 16 experts by one matrix product, the scores become probabilities (softmax), each token keeps its two most likely experts (top-2), and their probabilities, renormalized, are the gates. The experts' load is how many tokens each one got.

source

The vocabulary

F : Float

value · line 14

words: the a of and in to cat dog fox bird horse one two three seven ten red blue green brown runs jumps eats sees sleeps park river city home apple bread fish cheese today night morning ? Each word's 8 features (data/features.txt, a row per word, in the order above): animal, number, color, action, place, food, function word, time (fish is an animal and a food; eats an action and food; ? is any word not listed).

F ← 37 8 r̲eshape n̲umbers ⎕N̲GET "data/features.txt"
Used in: E

wr : Float

value · line 17

The embeddings: the features plus a small fixed ripple, so no two words are quite the same.

wr ← f̲loat (o̲ffsets 37) 'l̲eft t̲able o̲ffsets 8
Used in: E

wc : Float

value · line 18
wc ← f̲loat (o̲ffsets 37) 'r̲ight t̲able o̲ffsets 8
Used in: E

E : Float

value · line 19
E ← (f̲loat F) + 0.15 × s̲in (1.7 × wr) + 2.3 × wc
Used in: x0, d1, d2, x

The router weights

fi : Int

value · line 25

16 experts on a 4 x 4 grid: row r likes feature r (animal, number, color, action), column c likes feature 4 + c (place, food, function word, time), plus a small fixed ripple.

fi ← (r̲ange 8) 'l̲eft t̲able o̲ffsets 16
Used in: W

ei : Int

value · line 26
ei ← (r̲ange 8) 'r̲ight t̲able o̲ffsets 16
Used in: W

W : Float

value · line 27
W ← (2.0 × f̲loat (fi = 1 + ei d̲iv 4) ∨ fi = 5 + ei m̲od 4) + 0.4 × s̲in (0.9 × f̲loat fi) + 1.3 × f̲loat ei
Used in: ᵘs̲cores

The scores: every token against every expert

ᵘs̲cores : Float -> Float

function · line 29
ᵘs̲cores ← { x → 2.0 × x '+ '× i̲nner W }
Used in: gs, gg, gates

Softmax: each row into probabilities

ᵘs̲oftmax : Num a => a -> Float

function · line 33

Each row's largest is taken off first (the same result, no overflow); the row's sum divides it.

ᵘs̲oftmax ← { s →
  e ← e̲xp s − ('m̲ax r̲/₂ s) 'l̲eft t̲able o̲ffsets 16
  e ÷ ('+ r̲/₂ e) 'l̲eft t̲able o̲ffsets 16
}
Used in: gs, gg, gates

Top-2: the two largest of each row, as masks

ᵘt̲op2 : Float -> Float

function · line 42

The largest of each row marks the first expert; taken out, the largest of what is left marks the second. The gates are the two probabilities over their sum, 0 for the other 14 experts.

ᵘt̲op2 ← { p →
  m1 ← 'm̲ax r̲/₂ p
  s1 ← f̲loat p = m1 'l̲eft t̲able o̲ffsets 16
  p2 ← p − 2.0 × s1
  m2 ← 'm̲ax r̲/₂ p2
  s2 ← f̲loat p2 = m2 'l̲eft t̲able o̲ffsets 16
  (p × s1 + s2) ÷ (m1 + m2) 'l̲eft t̲able o̲ffsets 16
}
Used in: gs, gg, gates

The load: how many tokens each expert got

ᵘl̲oad : Float -> Float

function · line 51
ᵘl̲oad ← { g → '+ r̲/ f̲loat g > 0.0 }

The nudge: a token pushed towards other words

w0 : Int

value · line 58

Word w0 (green) moves towards word wa (fox) by eps from 0 to 1.2 in k steps; the slice spans the directions to wa and to wb (fish), side by side.

w0 ← 19
Used in: x0

wa : Int

value · line 59
wa ← 9
Used in: d1

wb : Int

value · line 60
wb ← 32
Used in: d2

k : Int

value · line 61
k ← 61

side : Int

value · line 62
side ← 24

Epsilon: where the experts change

ᵘf̲irst : Num a => a -> Float

function · line 65

Each token's first and second expert, 1 to 16.

ᵘf̲irst ← { g → '+ r̲/₂ (f̲loat g = ('m̲ax r̲/₂ g) 'l̲eft t̲able o̲ffsets 16) × (o̲ffsets t̲ally g) 'r̲ight t̲able f̲loat r̲ange 16 }

ᵘs̲econd : Float -> Float

function · line 66
ᵘs̲econd ← { g → ('+ r̲/₂ (f̲loat g > 0.0) × (o̲ffsets t̲ally g) 'r̲ight t̲able f̲loat r̲ange 16) − ᵘf̲irst g }

x0 : Float

value · line 67
x0 ← w0 s̲elect E
Used in: d1, d2, xs, xg

d1 : Float

value · line 68
d1 ← (wa s̲elect E) − x0
Used in: xs, xg

d2 : Float

value · line 69
d2 ← (wb s̲elect E) − x0
Used in: xg

eps : Float

value · line 71

x0 + eps d1 for every eps: one row each.

eps ← 1.2 × (f̲loat o̲ffsets k) ÷ f̲loat k − 1

xs : Float

value · line 72
xs ← (eps '× t̲able d1) + (o̲ffsets k) 'r̲ight t̲able x0
Used in: gs

gs : Float

value · line 73
gs ← ᵘt̲op2 ᵘs̲oftmax ᵘs̲cores xs

c : Float

value · line 75

The slice: x0 + a d1 + b d2 over a grid of a and b.

c ← -0.25 + 1.5 × (f̲loat o̲ffsets side) ÷ f̲loat side − 1
Used in: sa, sb

sa : Float

value · line 76
sa ← r̲avel (o̲ffsets side) 'r̲ight t̲able c
Used in: xg

sb : Float

value · line 77
sb ← r̲avel (r̲ev c) 'l̲eft t̲able o̲ffsets side
Used in: xg

xg : Float

value · line 78
xg ← (sa '× t̲able d1) + (sb '× t̲able d2) + (o̲ffsets side × side) 'r̲ight t̲able x0
Used in: gg

gg : Float

value · line 79
gg ← ᵘt̲op2 ᵘs̲oftmax ᵘs̲cores xg
Used in: pg

ids : Int

value · line 83

"the red fox eats fish in the park": word numbers (from 1).

ids ← 1 17 9 23 32 5 1 26
Used in: x

x : Float

value · line 84
x ← ids s̲elect E

gates : Float

value · line 86
gates ← ᵘt̲op2 ᵘs̲oftmax ᵘs̲cores x

idx : Float

value · line 89

Each token's first and second expert (numbered 1 to 16, row by row on the grid) and their gates.

idx ← (o̲ffsets 8) 'r̲ight t̲able f̲loat r̲ange 16
Used in: first, second

g1 : Float

value · line 90
g1 ← 'm̲ax r̲/₂ gates

first : Float

value · line 91
first ← '+ r̲/₂ idx × f̲loat gates = g1 'l̲eft t̲able o̲ffsets 16

second : Float

value · line 92
second ← ('+ r̲/₂ idx × f̲loat gates > 0.0) − first

ᵘr̲ound : Float -> Float

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

pair : Float

value · line 102
pair ← (16.0 × ᵘf̲irst gs) + ᵘs̲econd gs

pg : Float

value · line 106

The slice, each point's pair as a letter: regions with straight edges, where two experts' scores tie.

pg ← (16.0 × ᵘf̲irst gg) + ᵘs̲econd gg