Transformer - Transformer 的输入嵌入与位置编码
解释 Token Embedding、正弦余弦位置编码及其 PyTorch 实现与张量形状。
Transformer 的输入表示通常由 Token Embedding 与 Positional Encoding 相加得到。前者表示 token 的语义,后者为不具备循环结构的注意力网络提供顺序信息。
本文整理自
hengproject/ML_recall↗ 中的transformers/input_embedding.ipynb↗,公式来自 Attention Is All You Need ↗。
Token Embedding#
输入 token ID 的形状为 ,嵌入层维护一个形状为 的可学习矩阵,其中 是词表大小。查表后,每个 token 都被映射为一个 维向量:
import torch
import torch.nn as nn
class TokenEmbedding(nn.Module):
def __init__(self, vocab_size, d_model):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
def forward(self, token_ids):
return self.embedding(token_ids)python如果只使用 Token Embedding,模型无法区分相同 token 出现在不同位置的情况。Self-attention 本身对输入排列是等变的,因此还需要显式注入位置信息。
Sinusoidal Positional Encoding#
原始 Transformer 使用固定的正弦余弦位置编码。对位置 和通道索引 :
相邻的偶数、奇数通道使用相同频率的一对正弦和余弦函数,不同通道对覆盖不同尺度。低索引通道变化较快,高索引通道变化较慢,因此每个位置都获得一个多尺度表示。
最终输入为 token 向量与位置向量之和:
二者维度相同,所以相加不会改变张量形状。
为什么同时使用正弦与余弦#
对同一频率,在完整周期上的正弦和余弦正交。例如:
更关键的是,正弦与余弦的相位关系使位置偏移可表示为线性变换。对固定偏移 , 和 都能由位置 上同频率的正弦、余弦线性组合得到,这为模型学习相对位置关系提供了便利。
PyTorch 实现#
import math
import torch
import torch.nn as nn
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000):
super().__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(max_len).unsqueeze(1).float()
frequencies = torch.exp(
torch.arange(0, d_model, 2).float()
* (-math.log(10000.0) / d_model)
)
pe[:, 0::2] = torch.sin(position * frequencies) # 偶数索引
pe[:, 1::2] = torch.cos(position * frequencies) # 奇数索引
pe = pe.unsqueeze(0) # (1, max_len, d_model)
self.register_buffer("pe", pe)
def forward(self, x):
# x: (batch, seq_len, d_model)
return x + self.pe[:, : x.size(1)]pythonregister_buffer 表示 pe 是模块状态的一部分,但不是需要优化器更新的参数。它会随模型保存、加载和迁移设备。
位置编码预先生成到 max_len。前向传播时通过 self.pe[:, :x.size(1)] 取得当前序列长度对应的部分,并沿 batch 维广播。
组合输入层#
class InputEmbedding(nn.Module):
def __init__(self, vocab_size, d_model, max_len=5000):
super().__init__()
self.token_embedding = TokenEmbedding(vocab_size, d_model)
self.positional_encoding = PositionalEncoding(d_model, max_len)
def forward(self, token_ids):
x = self.token_embedding(token_ids)
return self.positional_encoding(x)python可以用一个简单测试确认形状:
batch_size = 2
seq_len = 10
d_model = 512
vocab_size = 1000
token_ids = torch.randint(0, vocab_size, (batch_size, seq_len))
embedding = InputEmbedding(vocab_size, d_model, max_len=seq_len)
output = embedding(token_ids)
print(token_ids.shape) # torch.Size([2, 10])
print(output.shape) # torch.Size([2, 10, 512])python实现注意点#
- 原论文会将 Token Embedding 乘以 后再与位置编码相加;是否保留这一缩放应与整体实现保持一致。
- 上面的切片赋值默认
d_model为偶数。若允许奇数维度,需要单独处理最后一个偶数索引通道。 - 固定位置编码没有可学习参数,但
max_len是硬上限;超长序列需要扩展 buffer 或动态生成。 - 现代 Transformer 也常使用可学习位置嵌入、RoPE 或相对位置偏置。它们改变位置建模方式,但不会改变 Token Embedding 的基本职责。