共计 1635 个字符,预计需要花费 5 分钟才能阅读完成。
BERT 模型的强大能力很大程度上来源于其精心设计的输入表示机制。今天我们就来拆解这个 ” 三合一 ” 的嵌入系统,看看词嵌入 (token embeddings)、位置嵌入(position embeddings) 和句段嵌入 (segment embeddings) 是如何协同工作的。

1. 三大嵌入层各司其职
-
词嵌入:将每个 token 映射到固定维度的向量空间,是语义表示的基础层。例如 ”bank” 在金融和河岸场景会有不同向量。
-
位置嵌入:为每个位置编号分配独特向量,解决 Transformer 无法天然感知顺序的问题。位置 3 和位置 7 的相同单词会获得不同表示。
-
句段嵌入:标识 token 属于句子 A 还是 B(主要用于问答和 NSP 任务),让模型理解句子边界。通常用 0 / 1 区分不同句子。
2. 数学上的优雅相加
最终输入表示是三个嵌入的简单相加:
E = E_{token} + E_{position} + E_{segment}
假设嵌入维度为 d,则三者都是∈ℝ^(d)的向量。这个设计保证了:
- 维度一致性:所有嵌入必须同维度
- 信息平等:相加不会人为赋予某类嵌入更高权重
- 计算高效:比拼接 (concatenation) 节省 2 / 3 的后续计算量
3. PyTorch 实现详解
import torch
import torch.nn as nn
class BERTEmbedding(nn.Module):
def __init__(self, vocab_size, max_len, d_model, n_segments=2):
super().__init__()
self.token_embed = nn.Embedding(vocab_size, d_model) # (V, d)
self.pos_embed = nn.Embedding(max_len, d_model) # (L, d)
self.seg_embed = nn.Embedding(n_segments, d_model) # (2, d)
def forward(self, x, seg_ids):
# x: (batch_size, seq_len)
# seg_ids: (batch_size, seq_len)
batch_size, seq_len = x.shape
# 生成位置编号 [0,1,2,...,seq_len-1]
pos = torch.arange(seq_len).expand(batch_size, seq_len) # (batch, seq)
# 求和并做 LayerNorm(实际 BERT 会额外加)return self.token_embed(x) + \
self.pos_embed(pos) + \
self.seg_embed(seg_ids)
关键点说明:
- 所有嵌入层维度必须相同
- 位置编号从 0 开始且不能超过 max_len
- seg_ids 通常是 0 / 1 组成的张量
4. 为什么选择相加而非拼接?
- 参数效率:拼接会使后续 Transformer 的输入维度膨胀 3 倍
- 实践效果:Google 实验显示相加足以让模型学习到不同嵌入的区分
- 梯度流动:相加操作的反向传播路径更简单直接
5. 实际应用注意事项
- 维度选择:通常与 Transformer 隐藏层一致(如 BERT-base 是 768)
- 初始化策略:位置嵌入常用正弦初始化,词嵌入用截断正态分布
- 长文本处理:当序列超过 max_len 时需截断或分块
- 跨语言场景:共享词嵌入矩阵时需扩展 vocab_size
6. 避坑指南
- ❌ 位置编号越界:比如 max_len=512 但输入了 513 个 token
- ❌ 句段 ID 混乱:三句文本却只用 0 / 1 标识(应扩展为 0 /1/2)
- ❌ 维度不匹配:三种嵌入初始化时设置了不同维度
- ❌ 忘记归一化:实际 BERT 会在相加后做 LayerNorm
7. 思考与延伸
在多语言 BERT 中,如何处理不同语言的词嵌入?可以考虑:
- 为每种语言维护独立的词嵌入矩阵
- 使用共享词典但区分语言标识符
- 在嵌入相加时额外添加语言类型嵌入
这种灵活的嵌入系统正是 BERT 强大的原因之一。下次当你 fine-tune 时,不妨试试调整嵌入层的组合方式,或许会有意外收获。
正文完
发表至: 自然语言处理
近两天内
