从零实现小型 GPT(进阶·小白超详细版)

升级说明:本文对 每一步 都做了「这一段要干嘛」+「逐行解释」+「数据形状变化」+「常见坑」四件套讲解;Step 6(Transformer Block, Pre-LN)做了特别扩展,既讲整体设计,又逐行拆解,适合完全新手。文末仍提供一键可运行脚本


目录


总览:我们要造一台“会接话”的机器

目标:给它一段开头文字,它能自己续写
方案:用解码器型 Transformer(GPT 家族的骨架),按块堆起来:

(每层)
LayerNorm → 多头自注意力(带因果遮罩) → 残差
LayerNorm → 前馈网络(FFN + GELU) → 残差

训练任务是预测下一个 token(本文用“字符”当 token,最直观)。


图片取自网络

Step 0:环境与整体思路

要干嘛?

  • 准备 PyTorch;固定随机数种子(便于复现实验);明确“自回归”训练套路。

代码 + 逐行解释

import math                       # 基本数学运算,例如开方 √d
import random                     # Python 随机(配合设种子,结果可复现)
import torch                      # PyTorch 核心库
import torch.nn as nn             # 神经网络层(Linear、Embedding、LayerNorm 等)
import torch.nn.functional as F   # 常用函数(softmax、cross_entropy 等)

def set_seed(seed: int = 42):
    random.seed(seed)                 # 固定 Python 层面的随机
    torch.manual_seed(seed)           # 固定 CPU 上的 torch 随机
    torch.cuda.manual_seed_all(seed)  # 固定 GPU 上的 torch 随机(如可用)

set_seed(42)

形状:此步不涉及张量形状。
:不设种子可能导致每次运行结果不同,初学者以为“出 bug 了”。


Step 1:字符级 Tokenizer(最直观)

要干嘛?

  • 字符映射成整数 id(模型只吃数字);反向解码做可视化输出。

代码 + 逐行解释

class CharTokenizer:
    def __init__(self, text: str):
        chars = sorted(list(set(text)))               # 去重+排序:所有出现过的字符
        self.stoi = {ch: i for i, ch in enumerate(chars)}  # 字符→id(string to id)
        self.itos = {i: ch for ch, i in self.stoi.items()} # id→字符(id to string)
        self.vocab_size = len(self.stoi)              # 词表大小

    def encode(self, s: str):
        return [self.stoi[ch] for ch in s]            # 逐字符查表

    def decode(self, ids):
        return "".join(self.itos[i] for i in ids)     # 逐 id 查表 + 拼接回字符串

形状:编码后是一个 Python 列表 [int, int, ...]
:字符集太小/太大都会影响效果;本文只是演示最直观路径。


Step 2:模型配置(Config)

要干嘛?

  • 把“可调旋钮”(层数、头数、维度、dropout 等)集中管理,初始化更清晰。

代码 + 逐行解释

class GPTConfig:
    vocab_size: int = None                 # 由 Tokenizer 决定,初始化后再填
    context_len: int = 128                 # 上下文窗口长度(最大序列长度)
    d_model: int = 256                     # 通道维/嵌入维(越大模型越“宽”)
    n_heads: int = 4                       # 注意力头数(d_model 必须被整除)
    n_layers: int = 6                      # Transformer 层数(越多越“深”)
    dropout: float = 0.1                   # 丢弃比例,训练时随机失活部分单元
    qkv_bias: bool = False                 # QKV 线性层是否带偏置(GPT-2 常 False)
    device: str = "cuda" if torch.cuda.is_available() else "cpu"  # 自动选 GPU/CPU

形状:此步不涉及张量形状。
d_model % n_heads != 0 会在多头拆分时报错。


Step 3:数据切块与小批采样

要干嘛?

  • 把长文本切成很多个定长窗口 (x, y),其中 yx 右移一位。模型任务就是“用前面预测后面”。

代码 + 逐行解释

def make_dataset(encoded_ids, block_size):
    inputs, targets = [], []
    for i in range(0, len(encoded_ids) - block_size):
        x = torch.tensor(encoded_ids[i:i+block_size], dtype=torch.long)     # 输入序列
        y = torch.tensor(encoded_ids[i+1:i+block_size+1], dtype=torch.long) # 目标序列=右移一位
        inputs.append(x); targets.append(y)
    X = torch.stack(inputs)   # (N, T)
    Y = torch.stack(targets)  # (N, T)
    return X, Y

def get_batch(X, Y, batch_size):
    idx = torch.randint(0, X.size(0), (batch_size,))  # 随机抽 batch 个样本
    return X[idx], Y[idx]                             # (B, T), (B, T)

形状

  • X/Y(N, T);batch 后变 (B, T)
    :目标一定要右移1 位;否则等于“看答案抄答案”。

Step 4:因果多头自注意力(含遮罩)

要干嘛?

  • 让序列中每个位置,按照内容挑选它想关注的历史 token;不能偷看未来(因果遮罩)。

核心逻辑

  1. 线性层一次性得到 Q, K, V
  2. 注意力分数 QKᵀ / √d,加上上三角为 -inf 的遮罩
  3. softmax 得到权重;与 V 相乘聚合;
  4. 多头并行,最后 concat + 线性投影回 d_model

代码(带关键注释)

class CausalSelfAttention(nn.Module):
    def __init__(self, d_model, n_heads, dropout=0.0, qkv_bias=False, context_len=128):
        super().__init__()
        assert d_model % n_heads == 0
        self.d_model = d_model
        self.n_heads = n_heads
        self.head_dim = d_model // n_heads
        self.context_len = context_len

        self.qkv = nn.Linear(d_model, 3 * d_model, bias=qkv_bias)   # 合并 QKV
        self.proj = nn.Linear(d_model, d_model, bias=True)          # 输出投影

        self.attn_drop = nn.Dropout(dropout)
        self.resid_drop = nn.Dropout(dropout)

        mask = torch.full((context_len, context_len), float("-inf"))# 初始化为 -inf
        mask = torch.triu(mask, diagonal=1)                         # 上三角(不含对角)保留 -inf,其他为 0
        self.register_buffer("mask", mask)                          # 注册为 buffer(不训练)

    def forward(self, x):                    # x: (B, T, C)
        B, T, C = x.size()
        qkv = self.qkv(x)                    # (B, T, 3C)
        q, k, v = qkv.chunk(3, dim=-1)       # 各自 (B, T, C)

        # 拆多头: (B, nH, T, head_dim)
        q = q.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
        k = k.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
        v = v.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)

        # 注意力: (B, nH, T, T)
        att = (q @ k.transpose(-2, -1)) / math.sqrt(self.head_dim)
        att = att + self.mask[:T, :T]        # 因果遮罩:禁止看未来
        att = F.softmax(att, dim=-1)         # 对每一行做概率归一化
        att = self.attn_drop(att)

        y = att @ v                           # 聚合 Value → (B, nH, T, head_dim)
        y = y.transpose(1, 2).contiguous().view(B, T, C)  # 拼回 (B, T, C)
        y = self.resid_drop(self.proj(y))
        return y

形状关键点

  • q/k/v(B, nH, T, d_head)att(B, nH, T, T)y 回到 (B, T, C)
    :遮罩必须用浮点 -inf;整型/布尔型会导致 softmax 数值错误。

Step 5:前馈网络 FFN(GELU)

要干嘛?

  • 在注意力之后再加一层非线性转换,单位置逐点计算,增强表达力。

代码 + 逐行解释

class FeedForward(nn.Module):
    def __init__(self, d_model, dropout=0.0):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(d_model, 4 * d_model),  # 扩维(经验上 4 倍效果好)
            nn.GELU(),                        # 比 ReLU 更平滑的激活
            nn.Linear(4 * d_model, d_model),  # 压回原维度
            nn.Dropout(dropout),              # 正则防过拟合
        )
    def forward(self, x):
        return self.net(x)                    # 逐位置独立处理(不看邻居)

形状:输入输出均 (B, T, C)
:把 dropout 放在中间/后面都有人用;本文采用常见的“末尾”做法。


Step 6:Transformer Block(Pre-LN)——超详细

这一段要干嘛?

  • 注意力前馈网络两大子层,用残差连接(skip-connection)串起来;
  • 采用 Pre-LayerNorm先 LN 再算子层,训练更稳定;
  • 一个 Block 就是“信息先交流(注意力),再深化(FFN)”。多层堆叠,能力递进。

整体结构(单层)

x_in
 ├─ LN ──> Self-Attention ──> Dropout ──┐
 └─────────────── Residual Add ◄────────┘  → x_mid
      │
      ├─ LN ──> FFN ──> Dropout ─────────┐
      └─────────────── Residual Add ◄────┘  → x_out

为什么要 LayerNorm(LN)?

  • 规范化每个样本在最后一维的分布(均值≈0、方差≈1),让不同层之间数值尺度稳定
  • Pre-LN(先 LN 再算子层)能让梯度从残差支路直接回传,深层训练不易崩。

代码 + 逐行解释

class TransformerBlock(nn.Module):
    def __init__(self, d_model, n_heads, dropout, qkv_bias, context_len):
        super().__init__()
        # 子层1:注意力前的规范化
        self.ln1 = nn.LayerNorm(d_model, eps=1e-5)
        # 注意力模块(含因果遮罩、多头并行、投影)
        self.attn = CausalSelfAttention(d_model, n_heads, dropout, qkv_bias, context_len)
        # 子层2:前馈前的规范化
        self.ln2 = nn.LayerNorm(d_model, eps=1e-5)
        # 前馈网络模块(两层 MLP + GELU + Dropout)
        self.ffn = FeedForward(d_model, dropout)

    def forward(self, x):
        # --- 子层1:注意力 ---
        # 1) 先对输入 x 做 LayerNorm,得到规范化表示
        # 2) 喂给自注意力,得到“与上下文交互后”的特征
        # 3) 再与原输入 x 做 “残差相加”(skip-connection)
        x = x + self.attn(self.ln1(x))

        # --- 子层2:前馈网络 ---
        # 1) 对上一步输出再做 LayerNorm
        # 2) 喂给前馈网络(逐位置非线性变换)
        # 3) 与输入做残差相加
        x = x + self.ffn(self.ln2(x))
        return x

形状跟踪(假设输入 x 形状 (B, T, C))

  • ln1(x):仍是 (B, T, C)
  • attn(ln1(x)):输出 (B, T, C)x + (...)(B, T, C)
  • ln2(x):仍 (B, T, C)ffn(ln2(x)):仍 (B, T, C);最终输出 (B, T, C)

为什么残差要“+ 原输入”?

  • 像修路:新建一条“高速路”(子层),但保留一条“老路”(原输入)并并行,防止新路修坏了就走不通(梯度也能从老路返回去)。

Pre-LN vs Post-LN(感性理解)

  • Post-LN:先过子层,再 LN,等于“先跑后洗澡”。深堆叠时,梯度可能过小/不稳
  • Pre-LN:先 LN,再跑子层,等于“先整理状态再跑”,梯度直通残差,更稳。

常见坑

  • 忘了 eps 可能导致数值不稳定(除以极小方差时发散),1e-5 是常用安全值;
  • 残差相加要求形状一致,注意力和 FFN 都必须输出 (B, T, C)

Step 7:组装 MiniGPT(权重共享)

要干嘛?

  • 词嵌入 + 位置嵌入;堆 n_layers 个 Block;末端 LN;线性输出头;
  • 权重共享:输出层权重与词嵌入矩阵共享,省参数、效果也常更好。

代码 + 逐行解释(节选)

class MiniGPT(nn.Module):
    def __init__(self, cfg: GPTConfig):
        super().__init__()
        self.cfg = cfg
        self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)     # 词嵌入
        self.pos_emb = nn.Embedding(cfg.context_len, cfg.d_model)    # 位置嵌入(可学习)
        self.drop = nn.Dropout(cfg.dropout)

        self.blocks = nn.ModuleList([                                # 堆叠若干 Block
            TransformerBlock(cfg.d_model, cfg.n_heads, cfg.dropout, cfg.qkv_bias, cfg.context_len)
            for _ in range(cfg.n_layers)
        ])
        self.ln_f = nn.LayerNorm(cfg.d_model, eps=1e-5)              # 末端 LayerNorm
        self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)

        self.lm_head.weight = self.tok_emb.weight                    # **权重共享**

        self.apply(self._init_weights)                               # 统一初始化

    def forward(self, idx):                 # idx: (B, T)
        B, T = idx.shape
        assert T <= self.cfg.context_len

        pos = torch.arange(0, T, device=idx.device).unsqueeze(0)     # (1, T)
        x = self.tok_emb(idx) + self.pos_emb(pos)                    # (B, T, C)
        x = self.drop(x)
        for blk in self.blocks:                                      # 依次通过各层
            x = blk(x)
        x = self.ln_f(x)
        logits = self.lm_head(x)                                     # (B, T, vocab_size)
        return logits

形状(B, T, C) → … → (B, T, C)(B, T, V)
:权重共享时要保证 lm_head 的权重尺寸与 tok_emb 完全一致。


Step 8:生成函数(贪心/Top-k/温度)

要干嘛?

  • 逐步自回归地生成下一个 token;控制“创造力/多样性”。

代码 + 逐行解释(节选)

@torch.no_grad()
def generate(self, idx, max_new_tokens=100, temperature=1.0, top_k=None):
    self.eval()                                      # 关闭 dropout
    for _ in range(max_new_tokens):
        idx_cond = idx[:, -self.cfg.context_len:]    # 只看最近窗口
        logits = self.forward(idx_cond)              # 所有位置的预测
        logits = logits[:, -1, :] / max(temperature, 1e-6)  # 最后一个位置 + 温度缩放

        if top_k is not None:                        # Top-k 过滤
            v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
            thresh = v[:, [-1]]
            logits = torch.where(logits < thresh, torch.full_like(logits, float("-inf")), logits)

        probs = F.softmax(logits, dim=-1)            # 转概率
        next_id = torch.multinomial(probs, num_samples=1)     # 按概率采样
        idx = torch.cat([idx, next_id], dim=1)       # 拼接在后面
    return idx

直觉temperature 越小越保守;top_k 越小越不容易“胡写”。


Step 9:最小训练循环(可直接运行)

要干嘛?

  • 构造 (x, y) 批次;前向 → 交叉熵损失 → 反向 → AdamW 更新;打印 loss;训练后采样一段文本。

关键代码(节选)

xb, yb = get_batch(X, Y, batch_size)                  # (B, T)
logits = model(xb)                                    # (B, T, V)
loss = F.cross_entropy(logits.view(-1, cfg.vocab_size), yb.view(-1))  # 展平后计算 CE
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)               # 防梯度爆炸
optimizer.step()

形状(B, T, V) 展平为 (B*T, V),与目标 (B*T) 对齐。


常见坑与答疑速查表

  • 为什么 loss 有时上上下下?
    随机采样 batch + dropout 带来噪声,整体下降就好。
  • 为什么小数据会“胡说八道”?
    训练语料太少,字符级粒度又细,泛化能力弱。
  • Pre-LN 一定比 Post-LN 好吗?
    深层/大模型里,Pre-LN 更稳更好训;但有些改进技巧可让 Post-LN 也可用。
  • 权重共享有什么用?
    省参、对齐输入输出空间,实践中常带来轻微收益。

完整可运行脚本(复制即用)

与基础版一致,仅讲解更细。保存为 mini_gpt.pypython mini_gpt.py 运行即可。

import math, random, torch
import torch.nn as nn
import torch.nn.functional as F

def set_seed(seed: int = 42):
    random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)

class CharTokenizer:
    def __init__(self, text: str):
        chars = sorted(list(set(text)))
        self.stoi = {ch: i for i, ch in enumerate(chars)}
        self.itos = {i: ch for ch, i in self.stoi.items()}
        self.vocab_size = len(self.stoi)
    def encode(self, s: str): return [self.stoi[ch] for ch in s]
    def decode(self, ids): return "".join(self.itos[i] for i in ids)

class GPTConfig:
    vocab_size: int = None
    context_len: int = 128
    d_model: int = 256
    n_heads: int = 4
    n_layers: int = 6
    dropout: float = 0.1
    qkv_bias: bool = False
    device: str = "cuda" if torch.cuda.is_available() else "cpu"

class CausalSelfAttention(nn.Module):
    def __init__(self, d_model, n_heads, dropout=0.0, qkv_bias=False, context_len=128):
        super().__init__()
        assert d_model % n_heads == 0
        self.d_model = d_model
        self.n_heads = n_heads
        self.head_dim = d_model // n_heads
        self.context_len = context_len
        self.qkv = nn.Linear(d_model, 3 * d_model, bias=qkv_bias)
        self.proj = nn.Linear(d_model, d_model, bias=True)
        self.attn_drop = nn.Dropout(dropout)
        self.resid_drop = nn.Dropout(dropout)
        mask = torch.full((context_len, context_len), float("-inf"))
        mask = torch.triu(mask, diagonal=1)
        self.register_buffer("mask", mask)
    def forward(self, x):
        B, T, C = x.size()
        qkv = self.qkv(x)
        q, k, v = qkv.chunk(3, dim=-1)
        q = q.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
        k = k.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
        v = v.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
        att = (q @ k.transpose(-2, -1)) / math.sqrt(self.head_dim)
        att = att + self.mask[:T, :T]
        att = F.softmax(att, dim=-1)
        att = self.attn_drop(att)
        y = att @ v
        y = y.transpose(1, 2).contiguous().view(B, T, C)
        y = self.resid_drop(self.proj(y))
        return y

class FeedForward(nn.Module):
    def __init__(self, d_model, dropout=0.0):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(d_model, 4 * d_model),
            nn.GELU(),
            nn.Linear(4 * d_model, d_model),
            nn.Dropout(dropout),
        )
    def forward(self, x): return self.net(x)

class TransformerBlock(nn.Module):
    def __init__(self, d_model, n_heads, dropout, qkv_bias, context_len):
        super().__init__()
        self.ln1 = nn.LayerNorm(d_model, eps=1e-5)
        self.attn = CausalSelfAttention(d_model, n_heads, dropout, qkv_bias, context_len)
        self.ln2 = nn.LayerNorm(d_model, eps=1e-5)
        self.ffn = FeedForward(d_model, dropout)
    def forward(self, x):
        x = x + self.attn(self.ln1(x))
        x = x + self.ffn(self.ln2(x))
        return x

class MiniGPT(nn.Module):
    def __init__(self, cfg: GPTConfig):
        super().__init__()
        self.cfg = cfg
        self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)
        self.pos_emb = nn.Embedding(cfg.context_len, cfg.d_model)
        self.drop = nn.Dropout(cfg.dropout)
        self.blocks = nn.ModuleList([
            TransformerBlock(cfg.d_model, cfg.n_heads, cfg.dropout, cfg.qkv_bias, cfg.context_len)
            for _ in range(cfg.n_layers)
        ])
        self.ln_f = nn.LayerNorm(cfg.d_model, eps=1e-5)
        self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
        self.lm_head.weight = self.tok_emb.weight
        self.apply(self._init_weights)
    def _init_weights(self, m):
        if isinstance(m, (nn.Linear, nn.Embedding)):
            nn.init.normal_(m.weight, mean=0.0, std=0.02)
            if isinstance(m, nn.Linear) and m.bias is not None:
                nn.init.zeros_(m.bias)
        elif isinstance(m, nn.LayerNorm):
            nn.init.ones_(m.weight); nn.init.zeros_(m.bias)
    def forward(self, idx):
        B, T = idx.shape
        assert T <= self.cfg.context_len
        pos = torch.arange(0, T, device=idx.device).unsqueeze(0)
        x = self.tok_emb(idx) + self.pos_emb(pos)
        x = self.drop(x)
        for blk in self.blocks: x = blk(x)
        x = self.ln_f(x)
        return self.lm_head(x)
    @torch.no_grad()
    def generate(self, idx, max_new_tokens=100, temperature=1.0, top_k=None):
        self.eval()
        for _ in range(max_new_tokens):
            idx_cond = idx[:, -self.cfg.context_len:]
            logits = self.forward(idx_cond)
            logits = logits[:, -1, :] / max(temperature, 1e-6)
            if top_k is not None:
                v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
                thresh = v[:, [-1]]
                logits = torch.where(logits < thresh, torch.full_like(logits, float("-inf")), logits)
            probs = F.softmax(logits, dim=-1)
            next_id = torch.multinomial(probs, num_samples=1)
            idx = torch.cat([idx, next_id], dim=1)
        return idx

def make_dataset(encoded_ids, block_size):
    inputs, targets = [], []
    for i in range(0, len(encoded_ids) - block_size):
        x = torch.tensor(encoded_ids[i:i+block_size], dtype=torch.long)
        y = torch.tensor(encoded_ids[i+1:i+block_size+1], dtype=torch.long)
        inputs.append(x); targets.append(y)
    return torch.stack(inputs), torch.stack(targets)

def get_batch(X, Y, batch_size):
    idx = torch.randint(0, X.size(0), (batch_size,))
    return X[idx], Y[idx]

def train_tiny():
    set_seed(42)
    raw_text = (
        "To code or not to code, that is the question.\n"
        "All models are wrong, but some are useful.\n"
        "Hello GPT!\n"
    )
    tokenizer = CharTokenizer(raw_text)
    cfg = GPTConfig(); cfg.vocab_size = tokenizer.vocab_size
    cfg.device = "cuda" if torch.cuda.is_available() else "cpu"
    model = MiniGPT(cfg).to(cfg.device)
    ids = tokenizer.encode(raw_text)
    block_size = min(cfg.context_len, 64)
    X, Y = make_dataset(ids, block_size)
    X = X.to(cfg.device); Y = Y.to(cfg.device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.01)
    model.train(); max_steps = 600; batch_size = 32
    for step in range(1, max_steps + 1):
        xb, yb = get_batch(X, Y, batch_size)
        logits = model(xb)
        loss = F.cross_entropy(logits.view(-1, cfg.vocab_size), yb.view(-1))
        optimizer.zero_grad(set_to_none=True); loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        optimizer.step()
        if step % 50 == 0: print(f"[step {step:04d}] loss={loss.item():.4f}")
    model.eval()
    start_ids = torch.tensor([tokenizer.encode("To ")], dtype=torch.long, device=cfg.device)
    out_ids = model.generate(start_ids, max_new_tokens=100, temperature=0.8, top_k=50)
    print("\n=== SAMPLE ===")
    print(tokenizer.decode(out_ids[0].tolist()))

if __name__ == "__main__":
    train_tiny()

这段代码整体都在干嘛(串起来说)

1. 构造训练样本:用滑动窗口把文本切成 (input, target),其中 target 就是 input 右移一位(下一个字符)。

2. 前向 + 反向

  1. logits = model(xb) 得到每个位置对下一个字符的预测;

  2. F.cross_entropy(logits.view(-1, V), yb.view(-1)) 计算平均交叉熵损失;

  3. 反向传播 + AdamW 更新参数。

3. 每 50 步打印一次 loss:让你确认模型是否在变好。

4. 训练结束做采样:用 generate() 从提示 "To " 开始,按设定的温度/Top-k 逐字符生成,得到 SAMPLE。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐