programdemos/attention/attention.xtl

Attention as array algebra: every token's query meets every token's key in one matrix product, S = Q K^T / sqrt d; softmax turns each row into weights; the output is the weighted sum of the values, Y = A V. One head, hand-set so that a word looks for what it can describe.

source · imports nn: libs/NN/src/NN.xtl

The vocabulary and the head, read from data/

names : Char

value · line 10
names ← 8 t̲ake₂ 21 9 r̲eshape ⎕N̲GET "data/words.txt"

F : Float

value · line 11
F ← 21 8 r̲eshape n̲umbers ⎕N̲GET "data/features.txt"
Used in: X, B

Wq : Float

value · line 12
Wq ← 8 2 r̲eshape n̲umbers ⎕N̲GET "data/wq.txt"
Used in: ᵘs̲cores

Wk : Float

value · line 13
Wk ← 8 2 r̲eshape n̲umbers ⎕N̲GET "data/wk.txt"
Used in: ᵘs̲cores

The core: attention for tokens X, where mask M allows

ᵘs̲cores : Float -> Float

function · line 16
ᵘs̲cores ← { X → ((X '+ '× i̲nner Wq) '+ '× i̲nner o̲\ X '+ '× i̲nner Wk) ÷ 2.0 ^ 0.5 }
Used in: ᵘa̲ttend

ᵘa̲ttend : Float -> Float -> Float

function · line 17
ᵘa̲ttend ← { X M → ⁿⁿs̲oftmax (ᵘs̲cores X) − 1e9 × 1.0 − M }

Drawing: rows labeled with their words

ᵘs̲hade : Float -> Char

function · line 21
ᵘs̲hade ← { A → (1 + f̲loor 0.5 + 7.0 × A) s̲elect " .:-=+*#" }

ᵘr̲ows : Int -> Char -> Char

function · line 22
ᵘr̲ows ← { ids C → (ids s̲elect names) c̲at₂ (((t̲ally ids) c̲at 2) r̲eshape " ") c̲at₂ C }

ᵘw̲ide : a -> a

function · line 23
ᵘw̲ide ← { C → ((-1 t̲ake s̲hape C) r̲eshape 3) r̲eplicate₂ C }

ᵘh̲eader : Int -> Char

function · line 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 " " }

ᵘt̲op : Int -> Float -> Char

function · line 26

The word each row looks at most, or - when no word gets half its weight.

ᵘt̲op ← { ids A →
  sure ← 0.5 ≤ 'm̲ax r̲/₂ A
  ids ᵘr̲ows ((sure × (ⁿⁿa̲rgmax A) s̲elect ids) + 22 × n̲ot sure) s̲elect names c̲at 1 8 r̲eshape "-       "
}

ids : Int

value · line 32

"the animal did not cross the street because it was tired"

ids ← 1 3 12 15 13 1 7 20 11 14 16

X : Float

value · line 33
X ← ids s̲elect F

n : Int

value · line 34
n ← t̲ally ids
Used in: everyone, causal

everyone : Float

value · line 35
everyone ← (n c̲at n) r̲eshape 1.0
Used in: A, B

A : Float

value · line 36
A ← X ᵘa̲ttend everyone

causal : Float

value · line 50

A causal mask (a word sees only itself and the words before it, as when generating text): "it" still finds the animal; "the" sees itself.

causal ← f̲loat (r̲ange n) '≥ t̲able r̲ange n

w : Int

value · line 57

"... because it was wide": wide finds the street. "it" still finds the animal: one head reads each word alone, so "it" cannot know yet which word comes after it; resolving it takes another layer.

w ← 1 3 12 15 13 1 7 20 11 14 18

B : Float

value · line 58
B ← (w s̲elect F) ᵘa̲ttend everyone