GOCLAWLLM ENGINEERING
GoClaw 首页

5. 从神经网络到语言模型

核心原理90~120 分钟
学习目标
  1. 把文本转换成错位的输入与标签
  2. 追踪 embedding、hidden state 与 logits 的形状
  3. 估算 embedding 和 LM Head 参数量
前置知识
  • 张量、softmax 与交叉熵
  • Tokenizer 基础

本章产物构造一个 batch,并逐位置解释 next-token loss。

5.1 Embedding 是可学习查表

若词表大小为 V、隐藏维度为 C,embedding 权重形状:

[V, C]

输入 token ID 是整数索引,embedding 取出对应行。它不是固定“词义表”,而是在训练中根据下一个 token 目标逐渐形成有用表示。

5.2 语言模型的训练样本

文本 token:

[12, 9, 31, 7, 5]

输入和标签:

input:  [12, 9, 31, 7]
target: [ 9,31,  7, 5]

模型在每个位置预测右移一位的 token。一次长度为 T 的序列能同时产生大约 T 个训练目标,这也是预训练可以充分利用普通文本的原因之一。

5.3 Logits 不是概率

最终隐藏状态 [B,T,C] 经 LM Head 映射为 [B,T,V]。每个位置有 V 个 logits。训练时交叉熵内部完成 log-softmax,不要先手动 softmax 再交叉熵,否则更慢且数值稳定性更差。

5.4 参数量粗估

一个标准 decoder-only Transformer Block 的主要矩阵:

Q/K/V/O 投影:约 4 × C²
MLP(扩展到约 4C 再投回):约 8 × C²
每层合计:约 12 × C²

L 层时主体约:

12 × L × C²

再加 embedding/LM head:

V × C

现代模型使用 GQA、SwiGLU 和不同扩展倍数,实际系数会变化,但该公式非常适合做量级判断。

5.5 逐位置算一次 Next-token Loss

设词表只有四个 token:[A, B, C, D]。某个位置的正确标签是 C,模型 logits 为:

[2.0, 1.0, 0.0, -1.0]

先减去最大值避免指数溢出,再做 softmax:

shifted = [0, -1, -2, -3]
exp     ≈ [1.000, 0.368, 0.135, 0.050]
sum     ≈ 1.553
P(C)    ≈ 0.135 / 1.553 ≈ 0.087
NLL     = -log(0.087) ≈ 2.44

虽然 A 的概率最高,但正确标签是 C,因此 loss 较大。反向传播会提高正确类别相对其他类别的 logit。这里优化的是条件分布,不是直接修改某条文本规则。

一个 batch 的交叉熵通常对所有有效位置取平均。若 labels 中某些位置为 -100,这些位置不进入平均;因此比较两个 loss 前必须确认有效 token 数和 mask 规则一致。

5.6 从形状检查完整前向路径

B=2T=4、隐藏维度 C=8、词表 V=32

阶段输入形状输出形状容易出错的位置
Token Embedding[2,4] 整数 ID[2,4,8]dtype 必须是整数索引
Decoder Blocks[2,4,8][2,4,8]residual 两侧形状必须一致
LM Head[2,4,8][2,4,32]最后一维对应词表
Shift Labels[2,4][2,4]输入与标签错开一位
Cross Entropylogits + labels标量展平后位置顺序必须一致

可执行的最小核对:

import torch

torch.manual_seed(42)
token_ids = torch.tensor([[1, 4, 2, 8], [3, 2, 9, 1]])
embedding = torch.nn.Embedding(32, 8)
lm_head = torch.nn.Linear(8, 32, bias=False)
hidden = embedding(token_ids)
logits = lm_head(hidden)
labels = torch.tensor([[4, 2, 8, -100], [2, 9, 1, -100]])
loss = torch.nn.functional.cross_entropy(
    logits.reshape(-1, 32), labels.reshape(-1), ignore_index=-100
)
assert hidden.shape == (2, 4, 8)
assert logits.shape == (2, 4, 32)
assert torch.isfinite(loss)

这段代码没有 attention,因此不能学习上下文关系;它只用于确认“ID → 表示 → 词表 logits → loss”的接口。

配套实验:从 Token ID 到 Next-token Loss Notebook

5.7 练习与验收

  1. [12,9,31,7,5] 切成长度 3 的两个训练窗口,写出每个窗口的 input 和 target。
  2. 手算上面四分类例题中正确标签改为 A 后的 NLL,并解释为什么变小。
  3. V=50,000C=4,096,不共享 embedding 与 LM Head 时二者合计多少参数?FP16 权重约占多少字节?
  4. 解释为什么训练时能并行计算 T 个位置,而生成时仍需逐 token 解码。
  5. 找到 TinyGPT 中 labels 错位、logits 展平和交叉熵的位置,用实际形状回答上述问题。

验收标准:不看代码也能画出 [B,T] → [B,T,C] → [B,T,V] → scalar loss,并解释每个维度及 mask 对平均 loss 的影响。


本章依据

原理性结论以原始论文、官方文档或公开教材为依据。论文中的实验结果只适用于其声明的模型、数据、硬件和评估设置。

  1. 神经语言模型、分布式词表示与条件概率建模。

  2. 自回归语言建模在大规模 decoder-only 模型中的实现。

  3. 交叉熵、困惑度与语言模型评估的定义。