sourcedemos/attention/attention.xtl

1⍝!/usr/bin/env xetal 2⍝# Attention as array algebra: every token's query meets every token's 3⍝# key in one matrix product, S = Q K^T / sqrt d; softmax turns each row 4⍝# into weights; the output is the weighted sum of the values, Y = A V. 5⍝# One head, hand-set so that a word looks for what it can describe. 6 7ⁿⁿ⁼u̲se< "NN" 8 9⍝## The vocabulary and the head, read from data/ 10names ← 8 t̲ake₂ 21 9 r̲eshape ⎕N̲GET "data/words.txt" 11F ← 21 8 r̲eshape n̲umbers ⎕N̲GET "data/features.txt" 12Wq ← 8 2 r̲eshape n̲umbers ⎕N̲GET "data/wq.txt" 13Wk ← 8 2 r̲eshape n̲umbers ⎕N̲GET "data/wk.txt" 14 15⍝## The core: attention for tokens X, where mask M allows 16ᵘs̲cores ← { X → ((X '+ '× i̲nner Wq) '+ '× i̲nner o̲\ X '+ '× i̲nner Wk) ÷ 2.0 ^ 0.5 } 17ᵘa̲ttend ← { X M → ⁿⁿs̲oftmax (ᵘs̲cores X) − 1e9 × 1.0 − M } 18⍝ -- end of the core --------------------------------------- 19 20⍝## Drawing: rows labeled with their words 21ᵘs̲hade ← { A → (1 + f̲loor 0.5 + 7.0 × A) s̲elect " .:-=+*#" } 22ᵘr̲ows ← { ids C → (ids s̲elect names) c̲at₂ (((t̲ally ids) c̲at 2) r̲eshape " ") c̲at₂ C } 23ᵘw̲ide ← { C → ((-1 t̲ake s̲hape C) r̲eshape 3) r̲eplicate₂ C } 24ᵘh̲eader ← { ids → " " c̲at r̲avel (1 t̲ake₂ ids s̲elect names) c̲at₂ ((t̲ally ids) c̲at 2) r̲eshape " " } 25⍝# The word each row looks at most, or - when no word gets half its weight. 26ᵘt̲op ← { ids A → 27 sure ← 0.5 ≤ 'm̲ax r̲/₂ A 28 ids ᵘr̲ows ((sure × (ⁿⁿa̲rgmax A) s̲elect ids) + 22 × n̲ot sure) s̲elect names c̲at 1 8 r̲eshape "- " 29} 30 31⍝# "the animal did not cross the street because it was tired" 32ids ← 1 3 12 15 13 1 7 20 11 14 16 33X ← ids s̲elect F 34n ← t̲ally ids 35everyone ← (n c̲at n) r̲eshape 1.0 36A ← X ᵘa̲ttend everyone 37⍝ The weights: row i is how much word i looks at each word (the 38⍝ columns are the same words, left to right). 39ᵘh̲eader ids 40ids ᵘr̲ows ᵘw̲ide ᵘs̲hade A 41⍝ What each word looks at most: tired finds the animal, it finds the 42⍝ animal; words with nothing to look for spread evenly (-). 43ids ᵘt̲op A 44⍝ The output for "tired" (Y = A V, V = X): it has taken on the 45⍝ animal's feature (animate, the first). 46(f̲loat f̲loor 0.5 + 100.0 × 11 s̲elect A '+ '× i̲nner X) ÷ 100.0 47 48⍝# A causal mask (a word sees only itself and the words before it, as 49⍝# when generating text): "it" still finds the animal; "the" sees itself. 50causal ← f̲loat (r̲ange n) '≥ t̲able r̲ange n 51ᵘh̲eader ids 52ids ᵘr̲ows ᵘw̲ide ᵘs̲hade X ᵘa̲ttend causal 53 54⍝# "... because it was wide": wide finds the street. "it" still finds 55⍝# the animal: one head reads each word alone, so "it" cannot know yet 56⍝# which word comes after it; resolving it takes another layer. 57w ← 1 3 12 15 13 1 7 20 11 14 18 58B ← (w s̲elect F) ᵘa̲ttend everyone 59w ᵘt̲op B