class TinyLM(nn.Module):
"""A small decoder-only transformer language model."""
def __init__(self, vocab_size, d_model=128, num_heads=2, num_blks=2,
max_len=64):
super().__init__()
self.num_heads = num_heads
self.token_emb = nn.Embedding(vocab_size, d_model)
self.pos_emb = nn.Embedding(max_len, d_model)
self.blks = nn.ModuleList([nn.ModuleDict(dict(
norm1=nn.LayerNorm(d_model),
qkv=nn.Linear(d_model, 3 * d_model),
proj=nn.Linear(d_model, d_model),
norm2=nn.LayerNorm(d_model),
mlp=nn.Sequential(nn.Linear(d_model, 4 * d_model), nn.GELU(),
nn.Linear(4 * d_model, d_model))))
for _ in range(num_blks)])
self.norm = nn.LayerNorm(d_model)
self.head = nn.Linear(d_model, vocab_size)
def attention(self, blk, X):
B, T, D = X.shape
q, k, v = blk['qkv'](X).chunk(3, dim=-1)
q, k, v = (u.reshape(B, T, self.num_heads, -1).transpose(1, 2)
for u in (q, k, v))
Y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
return blk['proj'](Y.transpose(1, 2).reshape(B, T, D))
def forward(self, X):
H = self.token_emb(X) + self.pos_emb(torch.arange(X.shape[1],
device=X.device))
for blk in self.blks:
H = H + self.attention(blk, blk['norm1'](H))
H = H + blk['mlp'](blk['norm2'](H))
return self.head(self.norm(H))