手搓一个小Transformer
最近在做大模型相关的项目,每次做类似项目时,都会再次拜读一下《Attention Is All You Need》。
虽然平时接触过不少基于Transformer的模型,例如BERT、GPT、Qwen等,也知道Self-Attention、多头注意力、位置编码等概念,但总感觉自己对Transformer的理解停留在“看过公式”的阶段。
于是决定尝试不依赖HuggingFace,从零实现一个最小可运行的Transformer文本生成模型,即 GNPT(这个名字是我乱起的,意为Generative Non-Pretrained Transformer)。
搭模型
由于模型只用于生成任务,因此本人搭建了一个decoder-only的模型,后续在关于encoder、decoder设计部分不再单独强调decoder-only。
Self Attetion
Transformer最核心的部分其实就是Self-Attention。输入序列经过Embedding后,可以表示为:
x.shape = [B, T, d_model]
其中:
B:Batch Size T:序列长度 d_model:Embedding维度
首先需要生成 Q、K、V:
self.q_proj = nn.Linear(d_model, d_model)self.k_proj = nn.Linear(d_model, d_model)self.v_proj = nn.Linear(d_model, d_model)
q = self.q_proj(x)k = self.k_proj(x)v = self.v_proj(x)刚开始我有个疑问:为什么要有 Q、K、V 三个矩阵?
我的理解是:
Q(Query):我想关注什么 K(Key):我有哪些特征可供匹配 V(Value):真正携带的信息
Attention 其实就是:
score = q @ k.transpose(-1, -2)
计算每个 Token 与其他 Token 的相关性。
Masked Attention
Transformer Decoder 有一个特点:
当前 Token 不能看到未来 Token。
例如:
hello
预测:
h -> e he -> l hel -> l hell -> o
在预测 e 时,模型不应该提前看到后面的 l。
因此需要构造下三角 Mask:
mask = torch.tril(torch.ones(T, T, device=x.device)).bool()得到:
1 0 0 0 1 1 0 0 1 1 1 0 1 1 1 1
然后:
score = score.masked_fill(~mask, float("-inf"))这样 Softmax 后未来位置概率自动变成 0。
多头注意力
原论文中提到:
Project the queries, keys and values h times
我刚开始理解为:
for i in range(h): q_i = Linear(x)即真的做 h 次投影。
后来阅读 PyTorch 实现后发现:
工程上通常采用:
nn.Linear(d_model, d_model)
一次性投影。
例如:
512 -> 512
然后 reshape:
[B,T,512]
变成:
[B,H,T,64]
例如:
8 Heads 512 = 8 × 64
这样实现效率更高。
Transformer Block
实现完 Attention 后,需要构建完整 Block。
主要包括:
Multi-Head Attention self.attn LayerNorm self.ln1 self.ln2
Feed Forward Network Linear ReLU Linear
Residual Connection x = x + …
这里还有一个有趣的问题。
原论文采用:
Attention ↓ Add ↓ LayerNorm
即 Post-LN。
而现代 GPT、LLaMA 等模型大多采用:
LayerNorm ↓ Attention ↓ Add
即 Pre-LN。
为了实现简单,我最终采用了 Pre-LN 的写法。
Token Embedding 与位置编码
Self-Attention 本身无法感知顺序。
例如:
I love you
和:
you love I
在纯 Attention 看来只是同一组 Token。
因此需要加入位置编码。
首先是 Token Embedding:
self.token_emb = nn.Embedding( vocab_size, d_model)然后构造位置:
pos = torch.arange( T, device=x.device)查位置表:
pos_vec = self.pos_emb(pos)
最后:
x = token_vec + pos_vec
这里我刚开始有个疑问:
为什么不是:
self.pos_emb(x)
后来才意识到:
Token ID 和 Position ID 完全不是一个东西。
完整代码
import math
import torchimport torch.nn as nnimport torch.nn.functional as F
class MultiHeadAttention(nn.Module): def __init__(self, in_dim, out_dim, num_heads, causal=False): super().__init__() assert out_dim % num_heads == 0
self.q_proj = nn.Linear(in_dim, out_dim) self.k_proj = nn.Linear(in_dim, out_dim) self.v_proj = nn.Linear(in_dim, out_dim) self.o_proj = nn.Linear(out_dim, out_dim) self.num_heads = num_heads self.head_dim = out_dim // num_heads self.causal = causal
def forward(self, x): B, T, dim = x.shape q = self.q_proj(x) k = self.k_proj(x) v = self.v_proj(x)
q = q.view(B, T, self.num_heads, self.head_dim).transpose(1, 2) k = k.view(B, T, self.num_heads, self.head_dim).transpose(1, 2) v = v.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
score = q @ k.transpose(-1, -2) / math.sqrt(k.size(-1))
if self.causal: mask = torch.tril(torch.ones(T, T)).to(x.device).bool() score = score.masked_fill(~mask, float('-inf'))
out = torch.softmax(score, -1) @ v out = out.transpose(1, 2).contiguous().view(B, T, -1) out = self.o_proj(out)
return out
class FFN(nn.Module): def __init__(self, in_dim, hidden_dim, out_dim): super().__init__() self.net = nn.Sequential( nn.Linear(in_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, out_dim) ) # self.layer1 = nn.Linear(in_dim, hidden_dim) # self.layer2 = nn.Linear(hidden_dim, out_dim)
def forward(self, x): # x = self.layer1(x).clamp(0) # x = self.layer2(x) x = self.net(x) return x
class DecoderBlock(nn.Module): def __init__(self, model_dim, num_heads, ffn_hidden): super().__init__() self.att = MultiHeadAttention(model_dim, model_dim, num_heads=num_heads, causal=True) self.ffn = FFN(model_dim, ffn_hidden, model_dim) self.ln = nn.LayerNorm(model_dim)
def forward(self, x): x = x + self.att(self.ln(x)) x = x + self.ffn(self.ln(x)) return x
class GLM(nn.Module): ''' Generational Language Model (这个名字也是乱起的,并非专业名词) ''' def __init__(self, vocab_size, max_len, model_dim, num_heads, num_blocks=2): super().__init__() self.emb = nn.Embedding(vocab_size, model_dim) self.pe = nn.Embedding(max_len, model_dim)
self.blocks = nn.Sequential(*[ DecoderBlock(model_dim, num_heads, 4 * model_dim) for _ in range(num_blocks) ])
self.ln = nn.LayerNorm(model_dim) self.head = nn.Linear(model_dim, vocab_size)
self.max_len = max_len
def forward(self, ids): B, T = ids.shape assert T <= self.max_len
token_embd = self.emb(ids) pos = torch.arange(T).to(ids.device) pe = self.pe(pos)
x = token_embd + pe
x = self.blocks(x) x = self.ln(x)
logits = self.head(x) return logits数据集
以往训练一个大模型是非常耗时耗力耗财的事情,但由于目标只是验证 Transformer 是否能够完成文本生成,因此没有引入复杂的数据集,而是直接构造了几组简单文本:
text = “#hello this is mike!#how are you? I’m happy!” * 100
这里有两个特殊字符:
# -> 开始符(BOS) $ -> 结束符(EOS)
例如:
#hello this is mike!$
表示一条完整样本。
Token划分
这里我采用的是最简单的字符级 Tokenizer。
例如:
hello
被切分为:
h e l l o
而不是:
hello
作为一个 Token。
对应代码:
vocab = sorted(list(set(text)))
stoi = {ch: tid for tid, ch in enumerate(vocab)}itos = {tid: ch for ch, tid in stoi.items()}其中:
stoi:字符 -> Token Id itos:Token Id -> 字符
这样构造出来的词表非常小,很省成。虽然真实大模型通常采用 BPE 或 SentencePiece,但对于学习 Transformer 来说,字符级 Tokenizer 已经足够。
训练
import math
import torchimport torch.nn as nnimport torch.nn.functional as Ffrom model import GLMfrom config import Config
def train(): text = Config.text vocab = Config.vocab
stoi = Config.stoi itos = Config.itos
data = torch.tensor([stoi[ch] for ch in text], dtype=torch.long)
vocab_size = Config.vocab_size max_context = Config.max_context batch_size = 16
def get_batch(): ix = torch.randint(0, len(text) - max_context + 1 - 1, (batch_size,)) x = torch.stack([data[i: i+max_context] for i in ix]) y = torch.stack([data[i+1: i+1+max_context] for i in ix]) return x, y
device = 'cuda' if torch.cuda.is_available() else 'cpu'
model = GLM( vocab_size=Config.vocab_size, max_len=Config.max_context, model_dim=Config.model_dim, num_heads=Config.num_heads, num_blocks=Config.num_blocks ).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
steps = 300 for step in range(steps): x, y = get_batch() x, y = x.to(device), y.to(device)
logits = model(x)
loss = F.cross_entropy( logits.view(-1, vocab_size), y.view(-1) )
optimizer.zero_grad() loss.backward() optimizer.step()
if step % 50 == 0: print(f'[step {step}] loss: {loss.item()}')
torch.save(model.state_dict(), 'model_state_dict.pth')
if __name__ == '__main__': train()使用
import math
import torchimport torch.nn as nnimport torch.nn.functional as Ffrom model import GLMfrom config import Config
@torch.no_grad()def generate(start_text, max_new_tokens=20): device = 'cuda' if torch.cuda.is_available() else 'cpu'
model = GLM( vocab_size=Config.vocab_size, max_len=Config.max_context, model_dim=Config.model_dim, num_heads=Config.num_heads, num_blocks=Config.num_blocks ).to(device)
model.load_state_dict(torch.load('model_state_dict.pth', map_location=device))
model.eval()
data = torch.tensor([Config.stoi[ch] for ch in start_text], dtype=torch.long).to(device).unsqueeze(0)
for _ in range(max_new_tokens): x = data[:, -Config.max_context:] logits = model(x).view(-1, Config.vocab_size) logits = logits[-1].view(Config.vocab_size)
probs = torch.softmax(logits, -1) nxt_id = torch.multinomial(probs, num_samples=1)
nxt = Config.itos[nxt_id.item()] if nxt == '$': break print(nxt, end='')
data = torch.cat([data, nxt_id.unsqueeze(0)], dim=1)
if __name__ == '__main__': while True: text = input() for ch in text: if ch not in Config.vocab: print(f'out of vocab') exit(0) generate(start_text=text) print()文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!
部分内容可能已过时