手搓一个小Transformer

1546 字
8 分钟
手搓一个小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 torch
import torch.nn as nn
import 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!#ThereisAlice!\#There is Alice!#how are you? I’m happy!#howareyou?Imok.\#how are you? I'm ok.” * 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 torch
import torch.nn as nn
import torch.nn.functional as F
from model import GLM
from 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 torch
import torch.nn as nn
import torch.nn.functional as F
from model import GLM
from 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()

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

手搓一个小Transformer
https://blog.cplee.cn/posts/make-transformer
作者
cplee
发布于
2026-04-12
许可协议
CC BY-NC-SA 4.0
最后更新于 2026-04-12,距今已过 51 天

部分内容可能已过时

评论区

Profile Image of the Author
cplee
Hello, I'm cplee.
公告
欢迎来到我的博客!点击下方 “了解更多” 查看更多信息
分类
标签
站点统计
文章
6
分类
3
标签
9
总字数
31,269
运行时长
0
最后活动
0 天前

目录