从零手写大模型 · 推理篇 04 - Positional Encoding(位置编码)
推理篇 · 查看系列一、上期留下的坑
上一期实现了 Embedding 层,把"猫吃鱼"里每个字都变成了向量。但有个问题没解决:
"猫吃鱼"和"鱼吃猫",经过 Embedding 之后,会变成同一堆向量。
不是 bug,是 Embedding 的本质决定的。
Embedding 说到底就是一张查找表:给一个 token_id(主键),返回对应的一行数据(向量)。这跟写一条 SQL 没有区别:
SELECT * FROM embedding_table WHERE token_id = ?
这条查询只关心"查的是哪个 id",不关心这次查询发生在整个请求的第几步。所以"猫吃鱼"和"鱼吃猫"分别查 {猫, 吃, 鱼} 这三个 id,查回来的三行向量完全一样,只是排列顺序不同。
而下一期要讲的 Self-Attention,本质上是向量两两之间做加权求和——这个操作对"顺序"天然不敏感,把输入打乱重排,算出来的结果只是跟着换个位置,数值本身不会变。
一句话总结: Embedding + Attention 这套组合,天生"位置盲"。但"猫吃鱼"和"鱼吃猫",谁吃谁完全反过来了。所以必须有一种机制,把"位置"信息也塞进向量里——这就是 Positional Encoding。
二、直接加个序号,行不行?
最直觉的想法:像给数据库表加一个自增的 row_number 字段一样,把位置 0, 1, 2... 直接拼上去或者加上去。
这么做有两个问题:
- 数值会失控。 句子长度是 3,位置编号是 0
2;句子长度是 500,编号就跳到 0499。这个数字随句子变长无限增长,而 Embedding 的数值通常在一个较稳定的小范围内波动。把一个可能高达几百的数直接加上去,相当于用巨大的噪音把语义信息淹没。 - 学不到"相对位置"。 语言中重要的往往不是"这是第 37 个字"这种绝对位置,而是"这个字和前一个字挨着""两个词隔了 5 个位置"这种相对关系。纯整数序号很难让模型自然学出这种可泛化的相对关系。
Transformer 论文给出的解法:用一组不同频率的正弦 / 余弦函数来编码位置。
三、论文原文怎么说
这一部分直接对照原论文(Attention Is All You Need, Vaswani et al., 2017,第 3.5 节)来看,人话翻译放在旁边。
| 论文原文核心结论 | 人话翻译 |
|---|---|
| 模型不含任何循环或卷积结构,要让模型利用序列的顺序信息,就必须显式注入位置信息 | 换成大白话:Self-Attention 天生不看顺序,顺序这件事必须有人手动"喂"给模型 |
位置编码的公式为正余弦函数:PE(pos,2i) = sin(pos/10000^(2i/d_model))PE(pos,2i+1) = cos(pos/10000^(2i/d_model)) |
每一对维度对应一个固定频率的波,i 越大波长越长、变化越"慢" |
| 这些波的波长在各维度上构成一个从 2π 到 10000·2π 的等比数列 | 维度从"变化很快"到"变化很慢"均匀铺开,覆盖各种颗粒度的位置信息 |
| 论文作者提到,选择这个函数是因为对任意固定偏移量 k,PE_pos+k 都可以表示成 PE_pos 的一个线性函数 | 只要位置间隔固定,两个位置编码之间总能用一个固定的线性变换互相转换——这让模型更容易学会"相对位置"这种关系,而不只是死记硬背绝对位置 |
| 论文也尝试过用可学习的位置向量代替,效果和固定的正余弦编码几乎一样(见论文 Table 3 第 E 行) | 说明"正余弦"这个具体形式不是唯一解,可学习的位置向量也能达到类似效果 |
| 但论文最终选择了固定的正余弦方案,原因是它可能让模型在推理时处理比训练时更长的序列 | 可学习的位置向量只认识训练时见过的那些位置;固定公式的正余弦不管多长的位置都能直接算出来,泛化性更好 |
这张表其实解释了两件事:为什么需要位置编码,以及为什么偏偏是正余弦这个形式——不是随手拍脑袋定的,而是围绕"能否表达相对位置""能否泛化到更长序列"这两个具体目标设计出来的。
四、直觉理解:像一块机械表
拆开公式看,其实很简单:每一对维度,都是一个固定频率的正弦波在这个位置上的取值。 维度下标 i 越大,波长越长,变化越"慢";i 越小,波长越短,变化越"快"。
这就好比一块机械表:
- 秒针转得最快
- 分针慢一档
- 时针最慢
只看秒针,没法知道现在几点几分——它每 60 秒就转回原点,信息会重复。但秒针 + 分针 + 时针组合在一起,在 12 小时内任意时刻,指针组合状态都是唯一的。
位置编码就是这个思路的连续版本:d_model 维里的每一对维度是一根"指针",转动频率各不相同。单独看某一维会重复,但把所有维度拼在一起看,在实际会用到的句子长度范围内(几千个 token),这个组合几乎不会出现重复的"指纹"。
五、容易搞混的点:Embedding 无界,位置编码有界
上一期特意强调过:Embedding 的数值不是被限制在 [-1, 1] 之间的——它是训练出来的参数,理论上可以是任意实数。
位置编码正相反:
不管句子多长、
pos这个数字有多大,sin和cos的输出永远严格落在 [-1, 1] 区间内。
这是三角函数的天然性质,与 pos 取值无关。哪怕句子长达一万个 token,第 9999 个位置算出来的编码值,依然老老实实待在 -1 到 1 之间。
这正是位置编码解决"数值失控"问题的关键:保证不管序列多长,加到 Embedding 上的这份"位置信息"始终数量级可控,不会喧宾夺主。
六、代码实现
沿用"猫吃鱼"这三个字,d_model = 8(原论文用 512,这里为了方便手写和打印用一个小很多的维度,原理完全等价)。
Step 1:实现位置编码矩阵
import torch
import math
vocab = {"猫": 0, "吃": 1, "鱼": 2}
d_model = 8
max_len = 10 # 预先算好足够长度的位置编码,用多少截多少
def get_positional_encoding(max_len, d_model):
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len).unsqueeze(1).float() # (max_len, 1)
div_term = torch.exp(
torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)
)
pe[:, 0::2] = torch.sin(position * div_term) # 偶数维用 sin
pe[:, 1::2] = torch.cos(position * div_term) # 奇数维用 cos
return pe
pe = get_positional_encoding(max_len, d_model)
print(pe[:3]) # 看看位置 0、1、2 分别编码成了什么
div_term 这行看着复杂,其实就是在按公式批量算每一对维度对应的频率,用 exp(log(...)) 是为了数值计算更稳定,效果和直接算 10000^(2i/d_model) 一样。
Step 2:把位置编码加到 Embedding 上
对比"猫吃鱼"和"鱼吃猫":
torch.manual_seed(42)
embedding = torch.nn.Embedding(len(vocab), d_model)
def encode(sentence_tokens):
ids = torch.tensor([vocab[t] for t in sentence_tokens])
emb = embedding(ids) # (seq_len, d_model) —— 纯语义向量,位置无关
pos_enc = pe[: len(sentence_tokens)] # 取出对应长度的位置编码
return emb + pos_enc # 语义 + 位置,两者相加
cat_eat_fish = encode(["猫", "吃", "鱼"])
fish_eat_cat = encode(["鱼", "吃", "猫"])
print("猫吃鱼:\n", cat_eat_fish)
print("鱼吃猫:\n", fish_eat_cat)
两个句子里"猫""吃""鱼"三个 token 完全一样,embedding 查出来的三行向量的集合也完全一样。但加了位置编码之后:
- "猫吃鱼"里,"猫"排在位置 0,拿到的是
pe[0] - "鱼吃猫"里,"猫"排到了位置 2,拿到的是
pe[2]
pe[0] 和 pe[2] 是两个不同的向量,所以"猫"这个字在两个句子里最终得到的向量不再相同。"鱼"和"吃"同理。整个句子的向量序列因此携带了"谁在前、谁在后"的信息——模型第一次有了区分"猫吃鱼"和"鱼吃猫"的物理基础。
Step 3:验证"有界性"不受句子长度影响
long_pe = get_positional_encoding(5000, d_model)
print("位置编码最小值:", long_pe.min().item())
print("位置编码最大值:", long_pe.max().item())
# 无论 max_len 设成 10 还是 5000,输出恒在 [-1, 1] 之间
哪怕把句子长度拉到 5000,位置编码的取值范围依然被死死摁在 [-1, 1] 之内,不会随长度增长而失控。
七、小结
这一期解决了上期留下的问题:
- 问题: Embedding 层本身是"位置盲"的查表操作,必须额外注入位置信息,语言的顺序语义才能被模型感知
- 方案: 用一组不同频率的正弦、余弦函数组合,生成一份"位置指纹"
- 好处一: 永远有界,不会因为句子变长而数值爆炸
- 好处二: 多维度组合出的模式在实际序列长度内几乎不重复,能唯一标识每一个位置
- 好处三(论文原话的设计初衷): 固定公式对任意长度都能直接算出结果,比"只认识训练时见过的位置"的可学习向量更容易泛化到长序列
现在"猫吃鱼"和"鱼吃猫"终于在向量层面区分开了。但一个新问题马上浮现:模型拿到这些携带位置信息的向量之后,具体是怎么"看"出"猫"和"鱼"之间存在"吃"这个动作关系的?这就要交给下一期的主角——Self-Attention 了。剧透一句:它的核心运算,和写一条数据库的 JOIN 查询,思路出奇地相似。