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.
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.
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.
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.