A.2 PyTorch 实现示例
以下代码展示了 Transformer 核心组件的 PyTorch 实现。更多与正文数学推导配套的可运行代码示例,可参见第 2 章、第 3 章与第 4 章中的内嵌代码。
缩放点积注意力
下面的函数显式区分三类掩码:is_causal 表示因果注意力,key_padding_mask 用布尔值标记有效键位置,attn_mask 支持布尔可见性掩码或加性 bias 掩码。
import torch
import torch.nn.functional as F
import math
def scaled_dot_product_attention(Q, K, V, attn_mask=None, key_padding_mask=None, is_causal=False):
"""缩放点积注意力。
布尔 mask 中 True 表示该位置可见;加性 mask 会直接加到 attention scores 上。
"""
if Q.dim() not in (3, 4) or K.dim() != Q.dim() or V.dim() != Q.dim():
raise ValueError("Q, K, V must all be 3D (B, L, D) or 4D (B, H, L, D)")
d_k = Q.size(-1)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
if is_causal:
q_len, k_len = Q.size(-2), K.size(-2)
causal = torch.ones(q_len, k_len, dtype=torch.bool, device=Q.device).tril(diagonal=k_len - q_len)
scores = scores.masked_fill(~causal, float("-inf"))
if attn_mask is not None:
if attn_mask.dim() == 2:
attn_mask = attn_mask.reshape((1,) * (scores.dim() - 2) + attn_mask.shape)
elif attn_mask.dim() == 3 and scores.dim() == 4:
attn_mask = attn_mask[:, None, :, :]
elif attn_mask.dim() != scores.dim():
raise ValueError("attn_mask must broadcast to attention scores")
attn_mask = attn_mask.to(device=scores.device)
if attn_mask.dtype == torch.bool:
scores = scores.masked_fill(~attn_mask, float("-inf"))
else:
scores = scores + attn_mask.to(dtype=scores.dtype)
if key_padding_mask is not None:
if key_padding_mask.shape != (Q.size(0), K.size(-2)):
raise ValueError("key_padding_mask must be (batch, key_len)")
if scores.dim() == 3:
valid_keys = key_padding_mask[:, None, :]
else:
valid_keys = key_padding_mask[:, None, None, :]
valid_keys = valid_keys.to(device=scores.device, dtype=torch.bool)
scores = scores.masked_fill(~valid_keys, float("-inf"))
fully_masked = torch.isneginf(scores).all(dim=-1, keepdim=True)
safe_scores = scores.masked_fill(fully_masked, 0.0)
attn_weights = F.softmax(safe_scores, dim=-1)
attn_weights = attn_weights.masked_fill(fully_masked, 0.0)
output = torch.matmul(attn_weights, V)
return output, attn_weights多头注意力
多头注意力需要 d_model 能被注意力头数整除,否则无法把投影后的张量均匀拆成多个头。
前馈网络
下面的代码块延续前文导入的 nn 和 F,展示逐位置前馈网络的最小结构。
RMSNorm
与 3.6 节的公式对应:不做去均值,仅以均方根缩放后乘可学习增益。
PyTorch 2.4 起也提供内置的 nn.RMSNorm,行为与上述实现一致,可直接相互对照验证。
旋转位置编码(RoPE)
与 4.3 节对应:按 预计算角度表,对 Q、K 的每对相邻维度做二维旋转。
可以直接验证 4.3 节的核心性质——旋转后的注意力分数只依赖相对位置:
注意:Llama、Hugging Face 等生产实现通常采用“前半-后半”配对(rotate_half)而非这里的相邻维度交错配对。两种布局在数学上等价(相差一个固定的维度置换),但已训练权重与具体布局绑定,移植权重时不能混用。
训练循环示例
下面是训练循环骨架,省略了具体模型类、数据加载器和 epoch 配置,重点展示损失计算、梯度裁剪和优化器更新的位置。
最后更新于
