sourcedemos/moe-router/moe-router.xtl
1⍝!/usr/bin/env xetal
2⍝# A mixture-of-experts router: each token of a sentence is scored
3⍝# against 16 experts by one matrix product, the scores become
4⍝# probabilities (softmax), each token keeps its two most likely
5⍝# experts (top-2), and their probabilities, renormalized, are the
6⍝# gates. The experts' load is how many tokens each one got.
7
8⍝## The vocabulary
9⍝# 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 ?
10⍝# Each word's 8 features (data/features.txt, a row per word, in the
11⍝# order above): animal, number, color, action, place,
12⍝# food, function word, time (fish is an animal and a food; eats an
13⍝# action and food; ? is any word not listed).
14F ← 37 8 r̲eshape n̲umbers ⎕N̲GET "data/features.txt"
15⍝# The embeddings: the features plus a small fixed ripple, so no two
16⍝# words are quite the same.
17wr ← f̲loat (o̲ffsets 37) 'l̲eft t̲able o̲ffsets 8
18wc ← f̲loat (o̲ffsets 37) 'r̲ight t̲able o̲ffsets 8
19E ← (f̲loat F) + 0.15 × s̲in (1.7 × wr) + 2.3 × wc
20⍝## The router weights
21⍝# 16 experts on a 4 x 4 grid: row r likes feature r
22⍝# (animal, number, color, action), column c likes
23⍝# feature 4 + c (place, food, function word, time),
24⍝# plus a small fixed ripple.
25fi ← (r̲ange 8) 'l̲eft t̲able o̲ffsets 16
26ei ← (r̲ange 8) 'r̲ight t̲able o̲ffsets 16
27W ← (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
28⍝## The scores: every token against every expert
29ᵘs̲cores ← { x → 2.0 × x '+ '× i̲nner W }
30⍝## Softmax: each row into probabilities
31⍝# Each row's largest is taken off first (the same
32⍝# result, no overflow); the row's sum divides it.
33ᵘs̲oftmax ← { s →
34 e ← e̲xp s − ('m̲ax r̲/₂ s) 'l̲eft t̲able o̲ffsets 16
35 e ÷ ('+ r̲/₂ e) 'l̲eft t̲able o̲ffsets 16
36}
37⍝## Top-2: the two largest of each row, as masks
38⍝# The largest of each row marks the first expert;
39⍝# taken out, the largest of what is left marks the
40⍝# second. The gates are the two probabilities over
41⍝# their sum, 0 for the other 14 experts.
42ᵘt̲op2 ← { p →
43 m1 ← 'm̲ax r̲/₂ p
44 s1 ← f̲loat p = m1 'l̲eft t̲able o̲ffsets 16
45 p2 ← p − 2.0 × s1
46 m2 ← 'm̲ax r̲/₂ p2
47 s2 ← f̲loat p2 = m2 'l̲eft t̲able o̲ffsets 16
48 (p × s1 + s2) ÷ (m1 + m2) 'l̲eft t̲able o̲ffsets 16
49}
50⍝## The load: how many tokens each expert got
51ᵘl̲oad ← { g → '+ r̲/ f̲loat g > 0.0 }
52⍝ -- end of the core ------------------------------
53
54⍝## The nudge: a token pushed towards other words
55⍝# Word w0 (green) moves towards word wa (fox) by eps
56⍝# from 0 to 1.2 in k steps; the slice spans the
57⍝# directions to wa and to wb (fish), side by side.
58w0 ← 19
59wa ← 9
60wb ← 32
61k ← 61
62side ← 24
63⍝## Epsilon: where the experts change
64⍝# Each token's first and second expert, 1 to 16.
65ᵘ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 }
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 }
67x0 ← w0 s̲elect E
68d1 ← (wa s̲elect E) − x0
69d2 ← (wb s̲elect E) − x0
70⍝# x0 + eps d1 for every eps: one row each.
71eps ← 1.2 × (f̲loat o̲ffsets k) ÷ f̲loat k − 1
72xs ← (eps '× t̲able d1) + (o̲ffsets k) 'r̲ight t̲able x0
73gs ← ᵘt̲op2 ᵘs̲oftmax ᵘs̲cores xs
74⍝# The slice: x0 + a d1 + b d2 over a grid of a and b.
75c ← -0.25 + 1.5 × (f̲loat o̲ffsets side) ÷ f̲loat side − 1
76sa ← r̲avel (o̲ffsets side) 'r̲ight t̲able c
77sb ← r̲avel (r̲ev c) 'l̲eft t̲able o̲ffsets side
78xg ← (sa '× t̲able d1) + (sb '× t̲able d2) + (o̲ffsets side × side) 'r̲ight t̲able x0
79gg ← ᵘt̲op2 ᵘs̲oftmax ᵘs̲cores xg
80⍝ -- end of the nudge -------------------------------
81
82⍝# "the red fox eats fish in the park": word numbers (from 1).
83ids ← 1 17 9 23 32 5 1 26
84x ← ids s̲elect E
85s̲hape x
86gates ← ᵘt̲op2 ᵘs̲oftmax ᵘs̲cores x
87⍝# Each token's first and second expert (numbered 1 to 16, row by
88⍝# row on the grid) and their gates.
89idx ← (o̲ffsets 8) 'r̲ight t̲able f̲loat r̲ange 16
90g1 ← 'm̲ax r̲/₂ gates
91first ← '+ r̲/₂ idx × f̲loat gates = g1 'l̲eft t̲able o̲ffsets 16
92second ← ('+ r̲/₂ idx × f̲loat gates > 0.0) − first
93ᵘr̲ound ← { v → (f̲loat f̲loor 0.5 + v × 1000.0) ÷ 1000.0 }
942 8 r̲eshape first c̲at second
952 8 r̲eshape (ᵘr̲ound g1) c̲at ᵘr̲ound 1.0 − g1
96⍝ The load on the 4 x 4 grid of experts.
974 4 r̲eshape ᵘl̲oad gates
98
99⍝ Nudging "green" towards "fox": the first and second experts at each
100⍝ eps, then the eps values where the pair changes.
101(2 c̲at k) r̲eshape (ᵘf̲irst gs) c̲at ᵘs̲econd gs
102pair ← (16.0 × ᵘf̲irst gs) + ᵘs̲econd gs
103ᵘr̲ound (1 + w̲here (1 d̲rop pair) ≠ -1 d̲rop pair) s̲elect eps
104⍝# The slice, each point's pair as a letter: regions with straight
105⍝# edges, where two experts' scores tie.
106pg ← (16.0 × ᵘf̲irst gg) + ᵘs̲econd gg
107(side c̲at side) r̲eshape ((u̲nique pg) i̲ndexOf pg) s̲elect "abcdefghijklmnopqrstuvwxyz"