Skip to content
当前页大纲

从零搭建一个小型语言模型(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)

「写作」「问答」「写代码」都只是这个目标训练出来的副产品。我们要做的四步:

  1. 分词器:把文本切成 token(词表中的编号)
  2. 模型:Transformer 解码器,输入 token 序列,输出下一个 token 的概率
  3. 训练:拿语料不断做「预测-对比-修正」
  4. 生成:模型接龙,逐 token 采样

二、为什么对齐 Gemma 架构

Google 的 Gemma 系列是「小模型」路线的代表:Gemma 3 提供小至 270M(2.7 亿)参数的规格,一张笔记本显卡就能推理。它的架构是现代小模型的教科书配置,也是我们 miniGPT 的蓝图:

组件老式 GPT-2现代(Gemma/Qwen 系)好处
归一化LayerNormRMSNorm少算均值,快 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 对本文目标是弯路):

python
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 整体结构

python
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, loss

4.2 RMSNorm:更简单的归一化

LayerNorm 要减均值、除方差、再仿射;RMSNorm 发现「减均值」那步基本没用,只用均方根缩放:

python
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.weight

4.3 RoPE:用旋转编码相对位置

绝对位置编码(GPT-2 的可学习向量)的问题:位置 5 和位置 7 的关系模型要单独学。RoPE 把每个位置乘上一个旋转角,Q 和 K 的内积自然只依赖相对距离

python
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
python
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 全在用):

python
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

python
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 x

Pre-Norm + 残差 是深层网络能训得动的关键——梯度可以沿着残差高速公路直通底层。

五、第三步:训练

python
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 以下,之后缓慢下降。小模型学到语言规律不需要多久,难的是学到「知识」和「推理」——那需要海量数据和参数

六、第四步:生成(模型接龙)

python
@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 idx

temperature 控制随机性(低=保守,高=放飞),top_k 砍掉长尾低概率 token 防胡言乱语。

七、参数量与训练成本的现实参考

规模配置参考硬件成本参考
~3M(本文默认)dim=256, 6 层笔记本 CPU 可训,看 loss 下降即可
64Mdim=512, 12 层RTX 3090 单卡 2 小时量级(MiniMind 公开数据)
3Bdim=2048+个人可推理(Ollama),从零训练已非个人可负担

小模型的正确期望:它会写出统计上像人话的话,但没有「智能」——这恰好让你看清尺度定律(Scaling Law)之前,语言模型的本体是什么。

八、下一步玩什么

  1. 换真语料:用几十 MB 的中文语料(如维基导出)替换 demo 数据,看模型背课文
  2. 加 KV Cache:给 Attention 加推理缓存,生成速度提升一个量级(工业推理的必备件)
  3. 上指令微调:参考我另一篇 AI 模型学习与对比 里的 LoRA 部分,把「接龙模型」变成「对话模型」
  4. 本地跑真模型:用 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)

MIT License.