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