从零搭建一个小型语言模型(miniGPT 实战)
现代大模型的架构论文读起来吓人,拆开看其实是一个不太大的 Transformer。这篇文章从零写一个「麻雀虽小五脏俱全」的语言模型——架构对齐 Google Gemma 的现代设计(RMSNorm / RoPE / GQA / SwiGLU),几百行 PyTorch,一张消费级显卡(甚至 CPU)就能训练。
写完你对 GPT 的祛魅程度会显著提升:它真的就是「预测下一个 token」。
一、先说清楚我们在做一件什么事
语言模型唯一的目标:
给定 token 序列 [t1, t2, ..., tn],预测下一个 token t(n+1) 的概率分布
P(t(n+1) | t1, ..., tn)「写作」「问答」「写代码」都只是这个目标训练出来的副产品。我们要做的四步:
- 分词器:把文本切成 token(词表中的编号)
- 模型:Transformer 解码器,输入 token 序列,输出下一个 token 的概率
- 训练:拿语料不断做「预测-对比-修正」
- 生成:模型接龙,逐 token 采样
二、为什么对齐 Gemma 架构
Google 的 Gemma 系列是「小模型」路线的代表:Gemma 3 提供小至 270M(2.7 亿)参数的规格,一张笔记本显卡就能推理。它的架构是现代小模型的教科书配置,也是我们 miniGPT 的蓝图:
| 组件 | 老式 GPT-2 | 现代(Gemma/Qwen 系) | 好处 |
|---|---|---|---|
| 归一化 | LayerNorm | RMSNorm | 少算均值,快 10%+ |
| 位置编码 | 可学习绝对位置 | RoPE 旋转位置 | 相对位置、外推性好 |
| 注意力 | MHA(Q/K/V 头数相同) | GQA 分组查询 | KV 头少,推理显存大降 |
| 前馈层 | GELU 两层 | SwiGLU 三层 | 同参数量下效果更好 |
参考生态现状:Karpathy 的 nanoGPT 社区把「GPT-2 级训练」从 2024 年的 45 分钟卷到 2025 年的 3 分钟以内;国产开源项目 MiniMind 用一张 RTX 3090、约 2 小时就能从零训练一个 64M 参数的完整模型——「自己训一个小模型」已经是业余可负担的事。
三、第一步:分词器
原理是 BPE(Byte Pair Encoding):从字符开始,反复合并语料中最高频的相邻对,直到词表达到目标大小。「人工智能」会被逐渐合并成一个 token,而罕见词拆成子词——词表大小和序列长度的折中艺术。
实战直接用现成的(自己从零写 BPE 对本文目标是弯路):
from tokenizers import Tokenizer, models, trainers, pre_tokenizers
# 训练一个 BPE 分词器(中文场景)
tokenizer = Tokenizer(models.BPE(unk_token="<unk>"))
tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel()
trainer = trainers.BpeTrainer(
vocab_size=4096, # 小词表足够小模型用
special_tokens=["<unk>", "<s>", "</s>"],
)
tokenizer.train(files=["corpus.txt"], trainer=trainer)
enc = tokenizer.encode("从零搭建小型语言模型")
print(enc.ids) # token 编号序列
print(enc.tokens) # 切分结果学习用途也可以用最简单的「字级分词」:每个汉字/字符一个 token。demo 阶段我推荐字级——代码最少,注意力可以直接观察到模型在学什么。
四、第二步:模型(核心,全部代码)
4.1 整体结构
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class MiniGPT(nn.Module):
def __init__(self, vocab_size, dim=256, n_layers=6, n_heads=8, n_kv_heads=4,
ffn_mult=4, max_seq_len=512):
super().__init__()
self.tok_emb = nn.Embedding(vocab_size, dim)
self.blocks = nn.ModuleList([
Block(dim, n_heads, n_kv_heads, ffn_mult) for _ in range(n_layers)
])
self.norm = RMSNorm(dim)
self.head = nn.Linear(dim, vocab_size, bias=False)
# 经典技巧:输出层与嵌入层共享权重(Gemma 同款)
self.head.weight = self.tok_emb.weight
self.ropes = RoPECache(dim // n_heads, max_seq_len)
def forward(self, idx, targets=None):
x = self.tok_emb(idx)
x = x * math.sqrt(x.size(-1)) # Gemma 的嵌入缩放
for block in self.blocks:
x = block(x, self.ropes)
x = self.norm(x)
logits = self.head(x)
loss = None
if targets is not None:
loss = F.cross_entropy(
logits.view(-1, logits.size(-1)), targets.view(-1)
)
return logits, loss4.2 RMSNorm:更简单的归一化
LayerNorm 要减均值、除方差、再仿射;RMSNorm 发现「减均值」那步基本没用,只用均方根缩放:
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x):
rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
return x * rms * self.weight4.3 RoPE:用旋转编码相对位置
绝对位置编码(GPT-2 的可学习向量)的问题:位置 5 和位置 7 的关系模型要单独学。RoPE 把每个位置乘上一个旋转角,Q 和 K 的内积自然只依赖相对距离:
class RoPECache:
def __init__(self, head_dim, max_seq_len, base=10000.0):
inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2).float() / head_dim))
t = torch.arange(max_seq_len).float()
freqs = torch.outer(t, inv_freq) # [seq, head_dim/2]
self.cos = freqs.cos().repeat_interleave(2, dim=-1) # [seq, head_dim]
self.sin = freqs.sin().repeat_interleave(2, dim=-1)
def apply(self, x, layer_past_len=0):
# x: [batch, heads, seq, head_dim]
seq_len = x.size(-2)
cos = self.cos[layer_past_len:layer_past_len + seq_len].to(x.device)
sin = self.sin[layer_past_len:layer_past_len + seq_len].to(x.device)
x1, x2 = x[..., 0::2], x[..., 1::2]
# 相邻两维一组做旋转
rotated = torch.stack(
[x1 * cos[..., 0::2] - x2 * sin[..., 0::2],
x1 * sin[..., 0::2] + x2 * cos[..., 0::2]], dim=-1
).flatten(-2)
return rotated
def rotate_qk(q, k, ropes):
# q: [b, n_heads, s, d] k: [b, n_kv_heads, s, d]
return ropes.apply(q), ropes.apply(k)4.4 GQA:注意力的省钱版本
标准多头注意力(MHA)每个头都要有 Q、K、V。GQA 让多个 Q 头共享一组 K/V 头——比如 8 个 Q 头共享 4 个 KV 头(本文配置):
- 训练时参数量不变(Q 还是 8 组)
- 推理时 KV Cache 体积直接减半——生成速度和显存的主要瓶颈就是 KV Cache
class Attention(nn.Module):
def __init__(self, dim, n_heads, n_kv_heads):
super().__init__()
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads
self.head_dim = dim // n_heads
self.q_proj = nn.Linear(dim, n_heads * self.head_dim, bias=False)
self.k_proj = nn.Linear(dim, n_kv_heads * self.head_dim, bias=False)
self.v_proj = nn.Linear(dim, n_kv_heads * self.head_dim, bias=False)
self.o_proj = nn.Linear(dim, dim, bias=False)
def forward(self, x, ropes, mask):
b, s, d = x.shape
q = self.q_proj(x).view(b, s, self.n_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(b, s, self.n_kv_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(b, s, self.n_kv_heads, self.head_dim).transpose(1, 2)
q, k = rotate_qk(q, k, ropes) # 应用 RoPE
# GQA:把 KV 头扩展到和 Q 头一样多(重复共享)
if self.n_kv_heads != self.n_heads:
rep = self.n_heads // self.n_kv_heads
k = k.repeat_interleave(rep, dim=1)
v = v.repeat_interleave(rep, dim=1)
out = F.scaled_dot_product_attention( # Flash Attention 内核
q, k, v, is_causal=True # 因果掩码:看不见未来
)
return self.o_proj(out.transpose(1, 2).reshape(b, s, d))
F.scaled_dot_product_attention是 PyTorch 2.x 的融合注意力实现,性能接近手写 Flash Attention,因果掩码一个is_causal=True搞定。
4.5 SwiGLU 前馈层
传统 FFN 是两层线性夹一个激活。SwiGLU 用「两条支路相乘」的结构,同参数量下表达力更强(Gemma、Qwen、Llama 全在用):
class SwiGLU(nn.Module):
def __init__(self, dim, hidden):
super().__init__()
self.gate = nn.Linear(dim, hidden, bias=False)
self.up = nn.Linear(dim, hidden, bias=False)
self.down = nn.Linear(hidden, dim, bias=False)
def forward(self, x):
return self.down(F.silu(self.gate(x)) * self.up(x))4.6 组装成 Block
class Block(nn.Module):
def __init__(self, dim, n_heads, n_kv_heads, ffn_mult):
super().__init__()
self.attn_norm = RMSNorm(dim)
self.attn = Attention(dim, n_heads, n_kv_heads)
self.ffn_norm = RMSNorm(dim)
self.ffn = SwiGLU(dim, dim * ffn_mult // 3 * 2) # 保持参数量对齐
def forward(self, x, ropes):
x = x + self.attn(self.attn_norm(x), ropes, None) # Pre-Norm 残差
x = x + self.ffn(self.ffn_norm(x))
return xPre-Norm + 残差 是深层网络能训得动的关键——梯度可以沿着残差高速公路直通底层。
五、第三步:训练
def get_batch(data, block_size, batch_size, device):
ix = torch.randint(len(data) - block_size, (batch_size,))
x = torch.stack([data[i:i + block_size] for i in ix]).to(device)
y = torch.stack([data[i + 1:i + block_size + 1] for i in ix]).to(device)
return x, y
def train():
device = "mps" if torch.backends.mps.is_available() else "cuda" if torch.cuda.is_available() else "cpu"
model = MiniGPT(vocab_size=5000, dim=256, n_layers=6).to(device)
data = torch.tensor(load_tokenized_corpus(), dtype=torch.long)
opt = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.1)
steps, warmup = 3000, 100
for step in range(steps):
x, y = get_batch(data, 256, 32, device)
lr = 3e-4 * min(step / warmup, 1.0) # 简化版学习率预热
for g in opt.param_groups:
g["lr"] = lr
_, loss = model(x, y)
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # 梯度裁剪
opt.step()
if step % 200 == 0:
print(f"step {step} | loss {loss.item():.4f}")训练的正常表现:loss 从 ~8.5(均匀分布的理论值)快速掉到 3 以下,之后缓慢下降。小模型学到语言规律不需要多久,难的是学到「知识」和「推理」——那需要海量数据和参数。
六、第四步:生成(模型接龙)
@torch.no_grad()
def generate(model, idx, max_new_tokens=200, temperature=0.8, top_k=50):
model.eval()
for _ in range(max_new_tokens):
idx_cond = idx[:, -512:] # 截断到最大上下文
logits, _ = model(idx_cond)
logits = logits[:, -1, :] / temperature # 只取最后一个位置
if top_k:
v, _ = torch.topk(logits, top_k)
logits[logits < v[:, [-1]]] = -float("inf")
probs = F.softmax(logits, dim=-1)
idx = torch.cat([idx, torch.multinomial(probs, 1)], dim=1)
return idxtemperature 控制随机性(低=保守,高=放飞),top_k 砍掉长尾低概率 token 防胡言乱语。
七、参数量与训练成本的现实参考
| 规模 | 配置参考 | 硬件成本参考 |
|---|---|---|
| ~3M(本文默认) | dim=256, 6 层 | 笔记本 CPU 可训,看 loss 下降即可 |
| 64M | dim=512, 12 层 | RTX 3090 单卡 2 小时量级(MiniMind 公开数据) |
| 3B | dim=2048+ | 个人可推理(Ollama),从零训练已非个人可负担 |
小模型的正确期望:它会写出统计上像人话的话,但没有「智能」——这恰好让你看清尺度定律(Scaling Law)之前,语言模型的本体是什么。
八、下一步玩什么
- 换真语料:用几十 MB 的中文语料(如维基导出)替换 demo 数据,看模型背课文
- 加 KV Cache:给 Attention 加推理缓存,生成速度提升一个量级(工业推理的必备件)
- 上指令微调:参考我另一篇 AI 模型学习与对比 里的 LoRA 部分,把「接龙模型」变成「对话模型」
- 本地跑真模型:用 Ollama 直接拉 Gemma 3 / Qwen 系列小模型对照感受——RAG 与 Agent 应用实战 有完整路径
参考来源
- Google Gemma 发布记录与模型卡(ai.google.dev/gemma/docs/releases)
- nanoGPT(github.com/karpathy/nanoGPT)及其社区 Speedrun 记录
- MiniMind 开源项目训练成本记录(64M / RTX 3090 / 约 2.3 小时)
- HuggingFace tokenizers 文档(huggingface.co/docs/tokenizers)