For the complete documentation index, see llms.txt. This page is also available as Markdown.

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 能被注意力头数整除,否则无法把投影后的张量均匀拆成多个头。

前馈网络

下面的代码块延续前文导入的 nnF,展示逐位置前馈网络的最小结构。

RMSNorm

3.6 节的公式对应:不做去均值,仅以均方根缩放后乘可学习增益。

PyTorch 2.4 起也提供内置的 nn.RMSNorm,行为与上述实现一致,可直接相互对照验证。

旋转位置编码(RoPE)

4.3 节对应:按 θi=100002i/d\theta_i = 10000^{-2i/d} 预计算角度表,对 Q、K 的每对相邻维度做二维旋转。

可以直接验证 4.3 节的核心性质——旋转后的注意力分数只依赖相对位置:

注意:Llama、Hugging Face 等生产实现通常采用“前半-后半”配对(rotate_half)而非这里的相邻维度交错配对。两种布局在数学上等价(相差一个固定的维度置换),但已训练权重与具体布局绑定,移植权重时不能混用。

训练循环示例

下面是训练循环骨架,省略了具体模型类、数据加载器和 epoch 配置,重点展示损失计算、梯度裁剪和优化器更新的位置。

最后更新于