WebSpinner — a spinning top above the WebSpinner wordmark
Academy

Lesson 4 — Attention, Without the Mystery

Three commands. You look at one training example, then at the rule that stops the model reading ahead, then at the weights attention actually produces.

First, open Terminal

  1. Hold ⌘ and press Space. A search box appears in the middle of the screen.
  2. Type Terminal and press Return. A window with plain text opens. That is Terminal.
  3. Click Copy beside a command below, click into the Terminal window, then press ⌘V to paste.
  4. Press Return to run it. Wait until the text stops moving before you do the next one.

1Look at one training example

Thirty-two sequences of one hundred and twenty-eight characters, and for each position, the character that comes next.

cd ~/llm-from-scratch && source .venv/bin/activate && python -c "
import torch, gpt
x, y = gpt.get_batch('train')
print('inputs :', tuple(x.shape))
print('targets:', tuple(y.shape))
print()
ctx = x[0, :12].tolist()
for i in range(1, 6):
    print(repr(gpt.decode(ctx[:i])), '-> next is', repr(gpt.itos[ctx[i]]))
"

Prints the shapes, then a few contexts and the character that followed.

2See the rule

Position three may look at 0, 1, 2 and 3 — and at nothing after it.

cd ~/llm-from-scratch && source .venv/bin/activate && python -c "
import torch
mask = torch.tril(torch.ones(6, 6)).bool()
print('position:   0      1      2      3      4      5')
for i, row in enumerate(mask.tolist()):
    print(f'  sees {i}: ', '  '.join('yes ' if v else ' no ' for v in row))
"

Prints which positions each position is allowed to see.

3See the weights

Each row is one position deciding how much to care about each earlier one. Rows add up to one.

cd ~/llm-from-scratch && source .venv/bin/activate && python -c "
import torch
torch.manual_seed(1337)
T = 6
scores = torch.randn(T, T)
mask = torch.tril(torch.ones(T, T)).bool()
scores = scores.masked_fill(~mask, float('-inf'))
w = torch.softmax(scores, dim=-1)
torch.set_printoptions(precision=2, sci_mode=False)
print(w)
print()
print('every row adds up to:', w.sum(dim=-1).tolist())
"

Builds a small attention matrix and prints it.