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"