class GPT(nn.Module):
"""Decoder-only transformer language model built from configurable
blocks."""
class CausalAttention(nn.Module):
"""Multi-head causal self-attention, optionally rotary."""
def __init__(self, num_hiddens, num_heads, bias=False, rope=False):
super().__init__()
self.num_heads, self.rope = num_heads, rope
self.W_qkv = nn.Linear(num_hiddens, 3 * num_hiddens, bias=bias)
self.W_o = nn.Linear(num_hiddens, num_hiddens, bias=bias)
def _rope(self, x):
d = x.shape[-1]
pos = torch.arange(x.shape[-2], dtype=torch.float32,
device=x.device)
inv_freq = 10000.0 ** (
-torch.arange(0, d, 2, device=x.device) / d)
theta = pos[:, None] * inv_freq[None, :]
cos, sin = torch.cos(theta), torch.sin(theta)
x1, x2 = x[..., 0::2], x[..., 1::2]
return torch.stack([x1 * cos - x2 * sin,
x1 * sin + x2 * cos], -1).flatten(-2)
def forward(self, X, *_):
B, T, D = X.shape
q, k, v = self.W_qkv(X).chunk(3, -1)
q, k, v = (u.reshape(B, T, self.num_heads, -1).transpose(1, 2)
for u in (q, k, v))
if self.rope:
q, k = self._rope(q), self._rope(k)
Y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
return self.W_o(Y.transpose(1, 2).reshape(B, T, D))
def __init__(self, vocab_size, num_hiddens=256, num_heads=8, num_blks=6,
max_len=1024, pos='rope', norm='rms', act='swiglu',
pre_norm=True, bias=False, dropout=0):
super().__init__()
self.pos, self.max_len = pos, max_len
self.token_emb = nn.Embedding(vocab_size, num_hiddens)
nn.init.normal_(self.token_emb.weight, std=0.02)
if pos == 'learned':
self.pos_emb = nn.Embedding(max_len, num_hiddens)
nn.init.normal_(self.pos_emb.weight, std=0.02)
attn = lambda: self.CausalAttention(num_hiddens, num_heads, bias,
rope=(pos == 'rope'))
self.blks = nn.ModuleList([
d2l.TransformerBlock(num_hiddens, num_heads, dropout, norm, act,
pre_norm, bias, attn_factory=attn)
for _ in range(num_blks)])
self.norm = (nn.RMSNorm if norm == 'rms'
else nn.LayerNorm)(num_hiddens)
def forward(self, X):
H = self.token_emb(X)
if self.pos == 'learned':
H = H + self.pos_emb(torch.arange(X.shape[1], device=X.device))
for blk in self.blks:
H = blk(H)
return F.linear(self.norm(H), self.token_emb.weight)