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.
The vocabulary
F : Float
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"
wr : Float
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
The router weights
fi : Int
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
W : Float
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
The scores: every token against every expert
Softmax: each row into probabilities
ᵘs̲oftmax : Num a => a -> Float
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 }
Top-2: the two largest of each row, as masks
ᵘt̲op2 : Float -> Float
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 }
The load: how many tokens each expert got
ᵘl̲oad : Float -> Float
ᵘl̲oad ← { g → '+ r̲/ f̲loat g > 0.0 }
The nudge: a token pushed towards other words
w0 : Int
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
Epsilon: where the experts change
ᵘf̲irst : Num a => a -> Float
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
ᵘ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 }
eps : Float
x0 + eps d1 for every eps: one row each.
eps ← 1.2 × (f̲loat o̲ffsets k) ÷ f̲loat k − 1
c : Float
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
xg : Float
xg ← (sa '× t̲able d1) + (sb '× t̲able d2) + (o̲ffsets side × side) 'r̲ight t̲able x0
ids : Int
"the red fox eats fish in the park": word numbers (from 1).
ids ← 1 17 9 23 32 5 1 26
idx : Float
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
first : Float
first ← '+ r̲/₂ idx × f̲loat gates = g1 'l̲eft t̲able o̲ffsets 16
second : Float
second ← ('+ r̲/₂ idx × f̲loat gates > 0.0) − first