MAX。头像
关注

手搓gpt_model

gpt_model.py 逐行讲解

本文件定义了一个完整的 GPT-2 模型(纯 PyTorch 手写,不依赖 HuggingFace)。它是后面所有训练脚本的"地基"。

阅读前请先看 00-整体概览.md 里的维度约定(B、T、E、V、H、d_h)。本笔记把每一步的张量形状都标注出来。

贯穿全文的运行示例:B=2(2 条样本),T=4(每条 4 个 token),E=768,V=21128,H=12,d_h=64。

阅读前必备:如果对"embedding、nn.Linear、softmax、交叉熵、广播、梯度下降"这些基础概念还不熟,强烈建议先花 20 分钟读 00-背景知识.md。本笔记会默认你已经懂了这些,只在个别关键处再补"背景知识"小框。


目录

  1. 导入
  2. LayerNorm(层归一化)
  3. GELU(激活函数)
  4. FeedForward(前馈层)
  5. MultiHeadAttention(多头因果自注意力)★
  6. TransformerBlock(Transformer 块)
  7. GPTModel(整机)
  8. generate(自回归生成)★
  9. text_to_token_ids / token_ids_to_text(文本↔张量)

1. 导入

import torch
import torch.nn as nn
  • torch:PyTorch 核心,张量运算、随机数、自动求导都在这里。
  • torch.nn as nn:神经网络模块库。nn.Module(所有网络层的基类)、nn.Linearnn.Embeddingnn.Dropoutnn.Sequential 都来自这里。

2. LayerNorm(层归一化)—— 第 5~16 行

class LayerNorm(nn.Module):
    def __init__(self, emb_dim):
        super().__init__()
        self.eps = 1e-5
        self.scale = nn.Parameter(torch.ones(emb_dim))
        self.shift = nn.Parameter(torch.zeros(emb_dim))

逐行:

  • class LayerNorm(nn.Module):定义一个网络层,必须继承 nn.Module。它免费提供:
    • model.parameters() → 拿到所有可学习参数(优化器要用);
    • model.to(device) → 把参数搬到 GPU/CPU;
    • model.train() / model.eval() → 切换训练/推理模式(影响 Dropout 等)。
  • def __init__(self, emb_dim):构造函数。emb_dim 是要归一化的维度,本项目 = 768。
  • super().__init__():必须先调用父类的初始化,否则 nn.Module 内部状态没建立,后面的参数注册会报错。
  • self.eps = 1e-5:一个极小的数。归一化公式 (x-μ)/√(σ²+eps) 里加它,防止方差 σ²=0 时除以 0。
  • self.scale = nn.Parameter(torch.ones(emb_dim))可学习参数 γ,形状 (768,),初值全 1。
  • self.shift = nn.Parameter(torch.zeros(emb_dim))可学习参数 β,形状 (768,),初值全 0。

nn.Parameter 的作用:告诉 PyTorch"这个张量是模型权重",它会被自动加入 model.parameters() 并被优化器更新。LayerNorm 里归一化后不直接输出,而是再做一次 γ·x + β——把"该缩放多少、平移多少"交给网络自己学。

    def forward(self, x):
        mean = x.mean(dim=-1, keepdim=True)
        var = x.var(dim=-1, keepdim=True, unbiased=False)
        norm_x = (x - mean) / torch.sqrt(var + self.eps)
        return self.scale * norm_x + self.shift

逐行(假设输入 x 形状 (2, 4, 768)):

  • x.mean(dim=-1, keepdim=True)
    • 沿最后一维(768 那维)求均值 → (2, 4, 1)
    • keepdim=True 让结果保留最后一维的 1,形状 (2,4,1)。这样减的时候靠广播自动对齐:(2,4,768) - (2,4,1) 会把 1 扩充成 768。如果 keepdim=False,结果 (2,4),维数都不匹配,减法会报错或语义出错。
  • x.var(dim=-1, keepdim=True, unbiased=False):沿最后一维求方差 → (2, 4, 1)
    • unbiased=False:分母用 N(总体方差)。Transformer 里标准做法就是这样(HF、原始 GPT 代码都如此)。
    • 对比:unbiased=True(默认)分母用 N-1,那是统计里估计样本方差用的,这里不需要。
  • norm_x = (x - mean) / torch.sqrt(var + self.eps)
    • (2,4,768) - (2,4,1)(2,4,768)(广播)。
    • sqrt 是逐元素开方,(2,4,1)。除法再广播 → (2,4,768)。现在每个 token 的 768 维向量都是"均值 0、方差 1"。
  • return self.scale * norm_x + self.shift(768,)(2,4,768) 广播,等价于每个维度乘以各自的 γ、加上各自的 β → 输出 (2,4,768)

要点:LayerNorm 是对最后一个维度做归一化。在 GPT 里就是"对每个 token 自己的 768 维向量归一化",跟其他 token、其他样本无关。

背景知识 · 手算一遍归一化(用 3 维向量演示,假装 emb_dim=3):
x = [1, 2, 3]

  • mean = 2,var = ((1−2)² + (2−2)² + (3−2)²)/3 = 2/3 ≈ 0.667,std ≈ 0.816;
  • norm_x = (x − mean)/std = [−1.225, 0, 1.225](均值 0、方差 1);
  • 初始 scale=[1,1,1]、shift=[0,0,0],所以输出就是 norm_x 本身;
  • 训练后 scale/shift 学到别的值,输出就变成"归一化后再缩放、平移"的样子——这正是 LayerNorm 比直接归一化更灵活的原因。

比方:把全班(768 个维度)成绩标准化,然后老师给每个科目配一个缩放系数(γ)和一个加分项(β),系数可学。


3. GELU(激活函数)—— 第 19~27 行

class GELU(nn.Module):
    def __init__(self):
        super().__init__()

    def forward(self, x):
        return 0.5 * x * (1 + torch.tanh(
            torch.sqrt(torch.tensor(2.0 / torch.pi)) *
            (x + 0.044715 * torch.pow(x, 3))
        ))

逐行:

  • __init__:没有参数,所以什么都不做。定义空构造函数是为了继承 nn.Module 的框架。
  • forward:GELU 的 tanh 近似公式
    • 真正的 GELU = x · Φ(x)(Φ 是标准正态分布的累积分布函数),里面含误差函数 erf,计算贵。
    • 这里用 0.5x(1+tanh(√(2/π)·(x+0.044715·x³))) 近似,误差只有 1e-3 量级。
    • torch.tensor(2.0 / torch.pi):构造常量 2/π 的张量。
    • torch.pow(x, 3):逐元素三次方。0.044715 是凑出来的最优系数。
  • 和 ReLU 的区别:ReLU 在 0 处硬截断(负全变 0,梯度在负区恒 0);GELU 是平滑的 S 形过渡,负数也有(很小的)非零梯度,训练更稳。GPT-2、BERT 都用它。
  • 形状:GELU 是逐元素运算,输入什么形状输出就是什么形状,不改变维度。

4. FeedForward(前馈层)—— 第 30~40 行

class FeedForward(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.layers = nn.Sequential(
            nn.Linear(cfg["emb_dim"], 4 * cfg["emb_dim"]),
            GELU(),
            nn.Linear(4 * cfg["emb_dim"], cfg["emb_dim"]),
        )

    def forward(self, x):
        return self.layers(x)

逐行:

  • cfg:配置字典,里面有 emb_dim 等字段。
  • nn.Sequential(...):把三个子模块按顺序打包,输入依次流过。
  • 三个子模块:
    1. nn.Linear(768, 3072):升维 4 倍。权重形状 (3072, 768),偏置 (3072,)
    2. GELU():非线性。
    3. nn.Linear(3072, 768):降回 768。权重形状 (768, 3072),偏置 (768,)
  • forward:直接返回 self.layers(x)

形状变化(输入 (2,4,768)):

(2,4,768) --Linear(768→3072)--> (2,4,3072) --GELU--> (2,4,3072) --Linear(3072→768)--> (2,4,768)

为什么中间要放大 4 倍再缩回来? Transformer 块里,注意力负责让 token 之间"互通消息",前馈层负责让每个 token"自己深度加工"。先摊到高维空间做非线性变换(表达能力更强),再压缩回原维度给下一层。4 倍是 GPT-2 的标准做法。


5. MultiHeadAttention(多头因果自注意力)—— 第 43~106 行 ★

这是整个模型最核心的部分,也是最容易绕晕维度的地方。别急,一步步来。

5.1 构造函数(第 43~61 行)

class MultiHeadAttention(nn.Module):
    def __init__(self, d_in, d_out, context_length, dropout, num_heads, qkv_bias=False):
        super().__init__()
        assert d_out % num_heads == 0, "d_out must be divisible by n_heads"

        self.d_out = d_out
        self.num_heads = num_heads
        self.head_dim = d_out // num_heads

逐行:

  • d_in:输入维度(=768),d_out:输出维度(=768)。
  • assert d_out % num_heads == 0:断言 768 能被 12 整除。如果不满足直接抛异常,防止后面 view 拆头失败。这是防御性检查。
  • self.head_dim = d_out // num_heads每个注意力头分到的维度 = 768 ÷ 12 = 64
        self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_key = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.out_proj = nn.Linear(d_out, d_out)
        self.dropout = nn.Dropout(dropout)
        self.register_buffer("mask", torch.triu(
            torch.ones(context_length, context_length), diagonal=1))

逐行:

  • 三个投影层:W_queryW_keyW_value,每个权重形状 (768, 768)
    • 输入 token 经过它们得到"查询 Q / 键 K / 值 V"三个身份。
    • bias=qkv_bias:是否带偏置。预训练(从头训)用 False;加载网上下载的权重用 True(见总览里的配置对比)。
  • out_proj = nn.Linear(768, 768):把多头拼回的向量再做一次线性变换,融合各头信息。
  • self.dropout = nn.Dropout(dropout):随机把一部分值置 0,防过拟合。
  • self.register_buffer("mask", torch.triu(torch.ones(context_length, context_length), diagonal=1)):这行是因果掩码。
    • torch.ones(context_length, context_length):全 1 方阵,如 256×256。
    • torch.triu(..., diagonal=1):只保留主对角线以上(diagonal=1 表示从"再往上一格"开始算)的三角形。得到的 mask:第 i 行第 j 列 = 1 当且仅当 j > i,即"后面的位置"是 1。
    • register_buffer 而非普通 self.mask = ...:buffer 会自动跟着模型 .to(device) 搬到 GPU,且不会被当成可学习参数去更新。掩码是固定的,不需要梯度。

为什么叫 buffer? nn.Parameter 要更新,buffer 不更新但跟着设备走。mask 正好属于后者。

5.2 forward 前段:投影与拆头(第 63~79 行)

    def forward(self, x):
        b, num_tokens, d_in = x.shape # (B, T, E)
        keys = self.W_key(x)  # 形状: (B, T, E)
        queries = self.W_query(x) # (B, T, E)
        values = self.W_value(x) # (B, T, E)

逐行(输入 x = (2, 4, 768)):

  • b, num_tokens, d_in = x.shape:解包得到 b=2,num_tokens=4,d_in=768。
  • 三个投影:Q/K/V 都还是 (2, 4, 768)。现在每个 token 有了三个"分身"。
        keys = keys.view(b, num_tokens, self.num_heads, self.head_dim)
        values = values.view(b, num_tokens, self.num_heads, self.head_dim)
        queries = queries.view(b, num_tokens, self.num_heads, self.head_dim)

        keys = keys.transpose(1, 2)
        queries = queries.transpose(1, 2)
        values = values.transpose(1, 2)

逐行(关键维度操作):

  • view(b, num_tokens, num_heads, head_dim):把最后一维 768 拆成 (12, 64):
    • (2, 4, 768)(2, 4, 12, 64)
    • view不复制数据的重排,要求总元素数一致:2×4×768 = 2×4×12×64。✓
  • transpose(1, 2):把"头"的维度从第 2 位换到第 1 位:
    • (2, 4, 12, 64)(2, 12, 4, 64)
    • 为什么?因为 torch 的矩阵乘法 a @ b 是把最后两维当矩阵乘。我们想让每个头内部做 T×d_h 的矩阵乘,所以要把 (T, H, d_h) 变成 (H, T, d_h),让头维度靠前、T 和 d_h 在最后两维。

经过拆头后:数据在内存里被看成 12 个"独立的小注意力",每个小注意力处理 (2, 4, 64),互不干扰。这就是"多头"。

5.3 forward 中段:计算注意力分数与掩码(第 81~95 行)

        attn_scores = queries @ keys.transpose(2, 3)
  • keys.transpose(2, 3):把 K 从 (2,12,4,64) 变成 (2,12,64,4)(交换 T 和 d_h)。
  • 矩阵乘:(2,12,4,64) @ (2,12,64,4) = (2,12,4,4)
  • 语义:attn_scores[b, h, i, j] = 第 b 条样本、第 h 个头里,第 i 个 token 的 Q 与第 j 个 token 的 K 的相似度。值越大表示"i 越该关注 j"。
        mask_bool = self.mask.bool()[:num_tokens, :num_tokens]
        attn_scores.masked_fill_(mask_bool, -torch.inf)

逐行:

  • self.mask.bool():把 0/1 变成 False/True(.bool() 转换)。[:num_tokens, :num_tokens]:掩码是 (context_length, context_length),这里按当前实际序列长度裁剪成 (4, 4)
  • masked_fill_(mask_bool, -torch.inf):凡是掩码为 True 的位置(右上三角 = 未来位置)填成 -infattn_scores(2,12,4,4) 就地修改。

为什么填 -inf 而不是 0 因为后面要过 softmax。-inf 经过 softmax 变成 0 概率(e^-inf = 0),0 却会变成正概率。只有 -inf 才能彻底禁止关注未来 token。

        attn_weights = torch.softmax(attn_scores / keys.shape[-1]**0.5, dim=-1)
        attn_weights = self.dropout(attn_weights)

逐行:

  • keys.shape[-1]head_dim = 64,**0.5 即 √64 = 8。除以 8 是缩放点积注意力的关键:当维度变大,点积结果的方差变大,softmax 输入过大会把概率推向极端(梯度消失)。除以 √d_h 把方差拉回 1。
  • softmax(..., dim=-1):沿最后一行(每个 token 的 4 个相似度)归一化成概率,和为 1。(2,12,4,4)
  • self.dropout(attn_weights):随机把部分权重置 0(同时放大剩余权重),防止模型过度依赖某个 token。形状不变。

背景知识 · softmax 手算(假设某一行分数是 [1.0, 2.0, 3.0]):

  • exp(1)≈2.72, exp(2)≈7.39, exp(3)≈20.09,和 ≈ 30.20;
  • 概率 ≈ [0.09, 0.24, 0.67]
    看出规律:分数 1→3,概率从 0.09 跳到 0.67,指数把差距放大了。这就是"注意力聚焦"的原理——相似度稍高的位置,权重会被显著放大。

背景知识 · 为什么除以 √d_k? 点积是 d_k=64 个数相乘再相加,维度越大结果方差越大;方差大了 softmax 输入"过饱和"(概率全挤向 0 或 1),梯度接近消失。除以 √64=8 把方差拉回 1,让 softmax 落在"敏感区",梯度才好传。

5.4 forward 后段:加权求和、拼头、输出(第 97~106 行)

        context_vec = (attn_weights @ values).transpose(1, 2)
  • attn_weights (2,12,4,4) @ values (2,12,4,64) = (2,12,4,64)
  • 语义:第 i 个 token 的输出 = 所有 j 的 attn_weights[i,j] × values[j] 的加权和。注意力越高的 token,贡献越大。
  • .transpose(1, 2)(2,12,4,64)(2,4,12,64),把 T 挪回来。
        context_vec = context_vec.reshape(b, num_tokens, self.d_out)
        context_vec = self.out_proj(context_vec)
        return context_vec

逐行:

  • reshape(b, num_tokens, self.d_out)(2,4,12,64)(2,4,768),把 12 个头的 64 维拼回一整个 768 维向量。
    • 注意这里 reshape 和前面 view 的区别:view 要求内存连续,reshape 更宽容(必要时会复制),二者在连续张量上等价。
  • out_proj(2,4,768)(2,4,768),把 12 个头的信息线性融合一遍。
  • 返回 (2,4,768),形状与输入一致,Transformer 块的残差连接就能用了。

5.5 注意力维度完整流程图

x                    (2, 4, 768)
W_query/W_key/W_value
Q/K/V                (2, 4, 768)          ← 三个投影
view                 (2, 4, 12, 64)       ← 拆成 12 头
transpose(1,2)       (2, 12, 4, 64)       ← 头放前面,方便矩阵乘
Q @ Kᵀ               (2, 12, 4, 4)        ← 每头内 T×T 相似度
masked_fill(-inf)    (2, 12, 4, 4)        ← 右上三角禁止
softmax(÷√64)        (2, 12, 4, 4)        ← 归一化成权重,和为 1
dropout              (2, 12, 4, 4)
× V                  (2, 12, 4, 64)       ← 加权求和
transpose(1,2)       (2, 4, 12, 64)
reshape              (2, 4, 768)          ← 拼回
out_proj             (2, 4, 768)

一句话:每个 token 用 Q 问"我该看谁",用 K 回答"我是谁",softmax 算好关注权重,再把 V(“我有什么信息”)按权重加权取回。


6. TransformerBlock(Transformer 块)—— 第 109~140 行

class TransformerBlock(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.att = MultiHeadAttention(
            d_in=cfg["emb_dim"],
            d_out=cfg["emb_dim"],
            context_length=cfg["context_length"],
            num_heads=cfg["n_heads"],
            dropout=cfg["drop_rate"],
            qkv_bias=cfg["qkv_bias"])
        self.ff = FeedForward(cfg)
        self.norm1 = LayerNorm(cfg["emb_dim"])
        self.norm2 = LayerNorm(cfg["emb_dim"])
        self.drop_resid = nn.Dropout(cfg["drop_rate"])

逐行:

  • 一个 Transformer 块 = 多头注意力 + 前馈层 + 两个 LayerNorm
  • self.att:实例化注意力,输入输出都是 emb_dim=768。
  • self.ff:实例化前馈层。
  • self.norm1 / self.norm2:两块子层各自的归一化。
  • self.drop_resid:残差连接上再加个 Dropout,进一步防过拟合。
  • 所有超参数从 cfg 字典读,保证各处一致。
    def forward(self, x):
        shortcut = x
        x = self.norm1(x) # (B, T, E)
        x = self.att(x)   # 形状 [B, T, E]
        x = self.drop_resid(x) # [B, T, E]
        x = x + shortcut  # 加回原始输入

逐行(输入 x = (2,4,768)):

  • shortcut = x:先保存输入,做残差连接
  • self.norm1(x):先归一化 → (2,4,768)
  • self.att(x):注意力 → (2,4,768)
  • self.drop_resid(x):dropout → (2,4,768)
  • x = x + shortcut残差加法(2,4,768)
    • 顺序是 norm → 子层 → 加回原始输入,这叫 Pre-LayerNorm(归一化在子层前)。这是 GPT-2 的排列方式(区别于原始 Transformer 的 Post-LN)。
    • Pre-LN 的梯度经过残差捷径可以"直通"上层,训练更稳定,所以现代大模型都用它。
        shortcut = x
        x = self.norm2(x)
        x = self.ff(x) # (B, T, E) --> (B, T, 4E) --> (B, T, E)
        x = self.drop_resid(x)
        x = x + shortcut  # 加回原始输入
        return x # shape: (B, T, E)
  • 第二个残差块,结构完全一样:norm2 → 前馈(内部升 4 倍再降回)→ dropout → 加回。
  • 输出 (2,4,768)

残差连接为什么重要? 梯度在深层网络里会随层数指数级衰减/爆炸(消失梯度)。有了 x + shortcut 这条"捷径",梯度可以原样穿透每一层,让 12 层也能稳定训练。


7. GPTModel(整机)—— 第 143~170 行

class GPTModel(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.tok_emb = nn.Embedding(cfg["vocab_size"], cfg["emb_dim"])
        self.pos_emb = nn.Embedding(cfg["context_length"], cfg["emb_dim"])
        self.drop_emb = nn.Dropout(cfg["drop_rate"])

        self.trf_blocks = nn.Sequential(
            *[TransformerBlock(cfg) for _ in range(cfg["n_layers"])])

        self.final_norm = LayerNorm(cfg["emb_dim"])
        self.out_head = nn.Linear(
            cfg["emb_dim"], cfg["vocab_size"], bias=False)

逐行:

  • self.tok_emb = nn.Embedding(vocab_size, emb_dim)词嵌入查找表,形状 (21128, 768)
    • 输入 token id(整数 0~21127),输出第 id 行对应的 768 维向量。等价于"查字典:第几个词对应哪个向量"。
  • self.pos_emb = nn.Embedding(context_length, emb_dim)位置嵌入表,形状 (256, 768)
    • 位置 id = token 在序列中的位置,第 0 位、第 1 位……各自学一个向量。
    • GPT-2 用可学习位置嵌入(不是三角函数),训练时一起更新。
  • self.drop_emb = nn.Dropout(drop_rate):对"词向量 + 位置向量"整体做 dropout。
  • self.trf_blocks = nn.Sequential(*[...])
    • 列表推导式 [TransformerBlock(cfg) for _ in range(cfg["n_layers"])] 生成 12 个块。
    • *(星号解包):把列表拆成 12 个位置参数喂给 Sequential。于是 12 个 Transformer 块首尾串联。
  • self.final_norm = LayerNorm(emb_dim):整机最后一层归一化(输出之前再来一次,规范表示)。
  • self.out_head = nn.Linear(emb_dim, vocab_size, bias=False):输出头,形状 (768, 21128)
    • 把每个 token 的 768 维表示映射成对全词表 21128 个词的打分。
    • bias=False:GPT-2 经典设计(并且常与词嵌入共享权重,见 load_weight.py)。
    def forward(self, in_idx):
        batch_size, seq_len = in_idx.shape # (B, T)
        tok_embeds = self.tok_emb(in_idx) # (B, T, E) 词嵌入
        pos_embeds = self.pos_emb(torch.arange(seq_len, device=in_idx.device)) # (B, T, E) 位置嵌入
        x = tok_embeds + pos_embeds  # 形状 [B, T, E]
        x = self.drop_emb(x) # (B, T, E)
        x = self.trf_blocks(x) # (B, T, E)
        x = self.final_norm(x) # (B, T, E)
        last_hidden_state = x
        logits = self.out_head(x)
        return {
            "logits": logits, # shape: (B, T, V)
            "last_hidden_state": last_hidden_state, # shape: (B, T, E)
        }

逐行(输入 (2,4)):

  • batch_size, seq_len = in_idx.shape:解包得到 B=2,T=4。
  • self.tok_emb(in_idx):查表 → (2, 4, 768)
  • torch.arange(seq_len, device=in_idx.device):生成 [0,1,2,3](位置 id)。
    • device=in_idx.device:位置 id 必须和输入在同一个设备(GPU/CPU),否则加法报错。这是常见 bug 点。
  • self.pos_emb(...):查表 → (4, 768)
  • x = tok_embeds + pos_embeds广播相加(2,4,768)。每个 token 的向量 = 词向量 + 它在序列中的位置向量。

背景知识 · 广播(broadcasting)(2,4,768) + (4,768) 为什么能加?PyTorch 自动把 (4,768) 沿 batch 维"复制"成 2 份再相加。规则是从右往左对齐维度,只要两边维度相等或有一边是 1 就合法。所以位置嵌入不用手动扩成 (2,4,768)。如果形状是 (2,4,768)(3,768),第 2 维 4≠3,直接报错——这是新手最常遇到的报错。

  • 为什么位置要加进去?Transformer 没有"先后顺序"的概念,注意力是全连接的;不加位置,把句子打乱顺序输出完全一样,模型就不知道"第一个字 vs 最后一个字"的区别。位置嵌入就是给 token 打上"第几个"的标签。
  • self.drop_emb(x)(2,4,768)
  • self.trf_blocks(x):12 层 Transformer 块依次处理 → (2,4,768)。到这里每个 token 的向量已经"融合了上下文信息"。
  • self.final_norm(x)(2,4,768)
  • last_hidden_state = x保存归一化后的隐藏状态。奖励模型(2-RM.py)会用到它——在词嵌入后面接一个线性层打分数。
  • self.out_head(x)(2,4,768)(2,4,21128)。每个 token 对全词表打分。
  • return {"logits": ..., "last_hidden_state": ...}:返回字典,两个结果都带上。
    • 语言模型训练用 logits(算下一个词预测误差)。
    • 奖励模型用 last_hidden_state(接打分头)。

为什么返回字典而不返回元组? 不同下游任务要的东西不同(LM 要 logits,RM 要 hidden state),字典让调用方按名字取,代码更可读、不怕位置顺序搞错。

整机形状总览:

token ids         (2, 4)
  ↓ tok_emb       查表
tok_embeds        (2, 4, 768)
  + pos_embeds    (4, 768) ——广播
x                 (2, 4, 768)
  ↓ dropout + 12×TransformerBlock + final_norm
last_hidden_state (2, 4, 768)   → RM 用
  ↓ out_head
logits            (2, 4, 21128) → LM 用

8. generate(自回归生成)—— 第 173~231 行 ★

def generate(model, idx, max_new_tokens, context_size, temperature=0.0, top_k=None, eos_id=None):

参数说明:

  • model:GPTModel。
  • idx提示词的 token id,形状 (B, T)
  • max_new_tokens:最多新生成多少个 token。
  • context_size:模型能看的最长上下文(= pos_emb 的行数)。
  • temperature:温度。>0 用采样生成;默认 0.0 走贪心。
  • top_k:只从分数最高的 k 个里挑。
  • eos_id:结束符 id,遇到就提前停止。

核心思想:自回归。模型每步只预测"下一个 token",然后把预测结果拼回去,再预测再拼,一个词一个词地"蹦"出整段话。

    for _ in range(max_new_tokens):
        idx_cond = idx[:, -context_size:]
        with torch.no_grad():
            outputs = model(idx_cond)
            logits = outputs["logits"]
        logits = logits[:, -1, :]

逐行(假设 idx = (1, 5),5 个 token):

  • 循环 max_new_tokens 次,每次蹦一个新 token。
  • idx[:, -context_size:]:取最后 context_size 个 token。如果提示词超长就裁掉前面,防止超过 pos_emb 上限(查表越界)。形状 (1, ≤context_size)
  • with torch.no_grad():推理不需要梯度。省显存、提速。外面调用的代码包了 torch.no_grad() 也可以,这里再包一层是双保险。
  • outputs["logits"](1, T', 21128)
  • logits[:, -1, :]:只取最后一个位置的 logits → (1, 21128)。因为要预测的是"下一个"token,只需要最后一个位置输出。
        if top_k is not None:
            top_logits, _ = torch.topk(logits, top_k)
            min_val = top_logits[:, -1]
            logits = torch.where(logits < min_val, torch.tensor(
                float("-inf")).to(logits.device), logits)

逐行(Top-k 过滤):

  • torch.topk(logits, top_k):取分数最高的前 k 个,返回 (值, 索引)。我们只用值 top_logits,所以索引用 _ 扔掉。
  • top_logits[:, -1]:第 k 大的值(也就是"门槛"),形状 (1,)
  • torch.where(logits < min_val, -inf, logits):比门槛低的全部置 -inf,留下分数最高的 k 个候选。-inf.to(logits.device) 确保设备一致。

Top-k 的作用:避免模型每次都选最高分(会重复、单调),只从"最有把握的 k 个"里随机抽一个,多样性更好。

        if temperature > 0.0:
            logits = logits / temperature
            logits = logits - logits.max(dim=-1, keepdim=True).values
            probs = torch.softmax(logits, dim=-1)
            idx_next = torch.multinomial(probs, num_samples=1)
        else:
            idx_next = torch.argmax(logits, dim=-1, keepdim=True)

逐行(temperature > 0 走采样分支):

  • logits / temperature温度缩放
    • temperature < 1:logits 变大,softmax 更"尖锐",输出更确定。
    • temperature > 1:logits 变小,分布更"平",输出更随机。
    • = 1:不变。
  • logits - logits.max(dim=-1, keepdim=True).values:每行减去本行最大值 → 最大变 0。数值稳定性技巧:防止 e^x 溢出(x 很大时 e^x 会变 inf 或 nan)。keepdim=True 保证减的时候形状 (1,1) 能广播。
  • torch.softmax(logits, dim=-1):归一化成概率分布 → (1, 21128),和为 1。
  • torch.multinomial(probs, num_samples=1)按概率分布抽样。像转盘抽奖,概率高的更容易中,但不是必中 → (1, 1)

逐行(temperature=0 走贪心分支):

  • torch.argmax(logits, dim=-1, keepdim=True):直接选分数最高的 token id → (1, 1)。完全确定,每次结果一样。
        if idx_next == eos_id:  # 如果设置了 eos_id,并且生成到结束 token,就提前停止
            break

        idx = torch.cat((idx, idx_next), dim=1)  # (batch_size, num_tokens+1)
    return idx

逐行:

  • if idx_next == eos_id::生成到了结束符就停。
    • 注意:idx_next(B,1) 张量,eos_id 是标量整数。单样本(B=1)时比较正确;若 B>1,这行会有歧义(PyTorch 会报错或做逐元素比较后整体判断),这是本代码的一个简化,实际生产会改成 torch.any(idx_next == eos_id)。单条生成没问题。
  • torch.cat((idx, idx_next), dim=1):把新 token 拼到序列末尾,(1,T)(1,T+1)
  • 循环结束后 return idx:返回 prompt + 新生成的全部 token

generate 与 0-PRETRAIN.py 里 generate_text_simple 的区别:后者是纯贪心(没有 top_k / temperature / eos 参数),本质一样,是教学简化版。


9. 文本 ↔ 张量工具(第 234~242 行)

def text_to_token_ids(text, tokenizer):
    encoded = tokenizer.encode(text, add_special_tokens=False)
    encoded_tensor = torch.tensor(encoded).unsqueeze(0)  # 增加 batch 维度
    return encoded_tensor

逐行:

  • tokenizer.encode(text, add_special_tokens=False):把中文文本切成 token id 列表。add_special_tokens=False 表示不加 <BOS> 等特殊标记。
    • 例:"这本书真是"[6821, 1726, 1210, 2682, 1878, 2123](6 个 id)。
  • torch.tensor(encoded):列表 → 张量,形状 (6,)
  • .unsqueeze(0):在 0 维加一个"1" → (1, 6),补上 batch 维度。模型要 (B, T),所以必须加。
def token_ids_to_text(token_ids, tokenizer):
    flat = token_ids.squeeze(0)  # 去掉 batch 维度
    return tokenizer.decode(flat.tolist(), skip_special_tokens=True)

逐行:

  • token_ids.squeeze(0):去掉 batch 维度:(1, 6)(6,)
  • .tolist():张量 → Python 列表。
  • tokenizer.decode(..., skip_special_tokens=True):把 id 还原成文本,跳过特殊 token。

这两个函数是一对:一个文本进去(加 batch),一个张量出来(去 batch),生成/推理时来回用。


本文件小结

  1. LayerNorm:对最后一维做"标准化 + 可学习缩放平移",输入输出形状相同。
  2. GELU:平滑激活函数,逐元素运算,不改形状。
  3. FeedForward:768 → 3072 → 768 的"膨胀-收缩"结构,让每个 token 深度加工。
  4. MultiHeadAttention:把 768 拆成 12 个 64 维头分别做缩放点积注意力,用因果掩码禁止偷看未来,最后拼回 768。
  5. TransformerBlock:norm → 注意力 → 加残差 → norm → 前馈 → 加残差(Pre-LN 结构)。
  6. GPTModel:词嵌入 + 位置嵌入 + 12 层块 + 归一化 + 输出头,输入 (B,T),输出 (B,T,V) 的 logits 和 (B,T,E) 的隐藏状态。
  7. generate:自回归逐 token 生成,支持 top-k 和温度采样。
  8. 两个工具函数负责文本与张量互转。

一句话记忆:模型把每个 token 变成 768 维向量,让它在 12 层里不断和上下文"交流",最后预测下一个词对全词表 21128 个词的分数。

下一篇:0-PRETRAIN.md——怎么用普通中文文本把模型从零训出来。

转载自 CSDN-专业IT技术社区

原文链接:https://blog.csdn.net/2302_80130040/article/details/163374832

文章来源crawl

评论

赞0

评论列表

微信小程序
QQ小程序

关于作者

点赞数:0
关注数:0
粉丝:0
文章:0
关注标签:0
加入于:--