跳转至

大模型手撕

手写多头自注意力前向传播:实现 MultiHeadSelfAttention 类,完成 Q/K/V 线性投影、分头、缩放点积注意力、合并多头和输出投影。

1.1 原理概述

多头自注意力是 Transformer 的核心组件。它通过将输入投影到多个不同的子空间,并行计算注意力,再从不同角度捕捉序列中的依赖关系。最终将所有头的输出拼接起来,通过一个线性投影层融合信息。

1.2 实现步骤

  • 线性投影:使用三个独立的线性层(或一个合并的大线性层)将输入 X 投影为 Query、Key、Value。

  • 分头:将投影后的张量从形状 (batch, seq_len, hidden_dim) 转换为 (batch, num_heads, seq_len, head_dim)。这一步涉及 view 和 transpose 操作,需要谨慎处理维度顺序以避免数据错位。

  • 缩放点积注意力:计算注意力分数 scores = (Q * K^T) / sqrt(d_k),应用 softmax 得到权重,再与 V 相乘。整个过程需要正确处理掩码(如 padding mask、causal mask)。

  • 合并多头:将所有头的输出转置回 (batch, seq_len, hidden_dim) 形状,再通过最终的输出投影层。

1.3 完整代码实现

import torch
import torch.nn as nn
import torch.nn.functional as F
import math

class MultiHeadSelfAttention(nn.Module):
    def __init__(self, hidden_dim, num_heads, dropout=0.1):
        super().__init__()
        assert hidden_dim % num_heads == 0, "hidden_dim must be divisible by num_heads"
        self.hidden_dim = hidden_dim
        self.num_heads = num_heads
        self.head_dim = hidden_dim // num_heads

        # 使用一个合并的线性层进行 Q、K、V 投影,提高计算效率
        self.qkv_proj = nn.Linear(hidden_dim, 3 * hidden_dim, bias=False)
        self.out_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x, key_padding_mask=None, attn_mask=None):
        """
        x: (batch, seq_len, hidden_dim)
        key_padding_mask: (batch, seq_len) 布尔张量,True 表示需要忽略的位置(如 padding)
        attn_mask: 可选的注意力掩码,形状可为 (seq_len, seq_len) 或 (batch, seq_len, seq_len)
        """
        B, seq_len, _ = x.shape

        # 1. 线性投影并分割 Q、K、V
        qkv = self.qkv_proj(x)  # (B, seq_len, 3 * hidden_dim)
        q, k, v = qkv.chunk(3, dim=-1)  # 每个 (B, seq_len, hidden_dim)

        # 2. 分头:reshape + transpose
        # 错误做法:q.view(B, self.num_heads, seq_len, self.head_dim)
        # 正确做法:先 reshape 成 (B, seq_len, num_heads, head_dim) 再 transpose
        q = q.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        k = k.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        v = v.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        # 现在形状均为 (B, num_heads, seq_len, head_dim)

        # 3. 缩放点积注意力
        scale = math.sqrt(self.head_dim)
        attn_scores = torch.matmul(q, k.transpose(-2, -1)) / scale  # (B, n_heads, L_q, L_k)

        # 4. 掩码处理
        # key_padding_mask: (B, L_k) -> (B, 1, 1, L_k)
        if key_padding_mask is not None:
            attn_scores = attn_scores.masked_fill(
                key_padding_mask.unsqueeze(1).unsqueeze(2), float('-inf')
            )
        # attn_mask 通常形状兼容,可直接应用
        if attn_mask is not None:
            attn_scores = attn_scores.masked_fill(attn_mask == 0, float('-inf'))

        # 5. Softmax 与 Dropout
        attn_weights = F.softmax(attn_scores, dim=-1)
        attn_weights = self.dropout(attn_weights)

        # 6. 加权求和
        attn_output = torch.matmul(attn_weights, v)  # (B, n_heads, L_q, head_dim)

        # 7. 合并多头
        # 转置为 (B, L_q, n_heads, head_dim) 再合并最后两个维度
        attn_output = attn_output.transpose(1, 2).contiguous()
        attn_output = attn_output.view(B, seq_len, self.hidden_dim)

        # 8. 输出投影
        output = self.out_proj(attn_output)
        return output

1.4 设计要点

  • 合并 QKV 投影:将三个线性层合并为一个可以提升并行度,是现代实现的标准做法。

  • 掩码的顺序:先应用 key_padding_mask,再应用 attn_mask,两者都是通过将对应位置设为 -inf 来实现屏蔽。

  • 头合并时的 contiguous():transpose 后张量在内存中不连续,需调用 .contiguous() 后才能使用 .view(),否则会报错。


详细写出分头时的 reshape 与 permute/transpose 操作,解释为何不能直接 reshape 成 (B, num_heads, seq_len, head_dim)。

2.1 分头操作的维度变换过程

假设输入经过线性投影后形状为 (B, seq_len, num_heads * head_dim)。我们需要将其转换为 (B, num_heads, seq_len, head_dim) 以便在注意力计算中每个头独立处理。

分两步:

  1. q = q.view(B, seq_len, self.num_heads, self.head_dim) → 形状 (B, seq_len, num_heads, head_dim)

  2. q = q.transpose(1, 2) → 形状 (B, num_heads, seq_len, head_dim)

2.2 为什么不能直接 reshape 成 (B, num_heads, seq_len, head_dim)?

关键在于 PyTorch 张量的内存布局和 view 的工作原理。

线性层输出的最后一维是 num_heads * head_dim 个元素。这些元素在内存中按如下顺序排列:

[seq_pos_0: head0_dim0, head0_dim1, ..., head0_dim_last, head1_dim0, head1_dim1, ..., head1_dim_last, ...]

即对于每个序列位置,各个头的元素是交错存储的。view 操作要求新形状在内存中是连续的,并且按照行优先顺序依次读取。

如果直接 view(B, num_heads, seq_len, head_dim),则新形状下的内存读取顺序为:先填充 batch 维度,然后 num_heads 维度,接着 seq_len 维度,最后 head_dim 维度。这会将原本属于同一个序列位置、不同头的元素错误地分配给不同的头,导致数据错乱。

因此,必须先将 head_dimnum_heads 维度分离开来,即先 view 成 (B, seq_len, num_heads, head_dim)。此时,由于我们在原始形状中最后一维就是按 (head_dim) 为一组连续存放的,然后才轮到下一个头,所以这个 view 是合法的——它只是将原本隐式的 num_heads * head_dim 拆分成了显式的两个维度。然后通过 transpose(1, 2) 交换 seq_lennum_heads 的位置,得到最终的标准形状。

2.3 验证正确性

可通过一个小例子来验证两种方式的差异。如果直接 view 成错误形状,会导致同一序列位置的 Q 向量中包含不同头的混合信息,使得注意力计算完全错误。


实现带因果遮罩的自回归自注意力,要求用 torch.ones 和 torch.tril 生成下三角 mask,并在 softmax 前将非法位置置为 -inf。

3.1 因果掩码原理

在自回归解码中,位置 i 只能关注到位置 j (j ≤ i),以防止信息从未来位置泄露。这通过一个下三角矩阵实现:对 attn_scores 中 j > i 的位置设为负无穷,使得 softmax 后权重为 0。

3.2 实现代码

def generate_causal_mask(seq_len, device):
    # 生成 (seq_len, seq_len) 的下三角矩阵,1 表示保留,0 表示屏蔽
    mask = torch.tril(torch.ones(seq_len, seq_len, device=device))
    return mask

# 在 MultiHeadSelfAttention 的 forward 中应用
causal_mask = generate_causal_mask(seq_len, x.device)
attn_scores = attn_scores.masked_fill(causal_mask == 0, float('-inf'))

如果需要支持 batch 维度,可以将因果掩码扩展为 (1, 1, seq_len, seq_len),自动广播到 (B, num_heads, seq_len, seq_len)

3.3 与 key_padding_mask 的组合

通常先应用 key_padding_mask,再应用 causal_mask。两者互不影响,因为 causal_mask 只限制位置关系,而 key_padding_mask 屏蔽特定位置的任何交互。


实现同时支持 key_padding_mask 和 attn_mask 的注意力函数,正确处理布尔填充掩码与附加掩码的广播。

4.1 掩码的类型与作用

  • key_padding_mask:形状 (B, L_k),布尔张量,True 表示该位置是填充 token,需要被完全忽略(在注意力计算中对所有 query 位置均屏蔽)。这通常用于变长序列。

  • attn_mask:形状灵活,可以是 (L_q, L_k)(B, L_q, L_k) 等,用于实现因果掩码或其他自定义的注意力模式(如局部注意力、特定 token 间的交互控制)。

4.2 广播与处理

  • key_padding_mask 需要从 (B, L_k) 广播到 (B, 1, 1, L_k),这样它可以应用到所有头和所有查询位置。

  • attn_mask 需要能广播到 (B, num_heads, L_q, L_k)。常见的做法是传入形状为 (L_q, L_k) 的因果掩码,利用 PyTorch 的广播机制自动匹配。

4.3 实现代码

def apply_masks(attn_scores, key_padding_mask=None, attn_mask=None):
    """
    attn_scores: (B, num_heads, L_q, L_k)
    key_padding_mask: (B, L_k) bool, True = pad
    attn_mask: 可以是 (L_q, L_k) 或 (B, L_q, L_k),0 表示屏蔽
    """
    if key_padding_mask is not None:
        # 扩展维度以便广播
        attn_scores = attn_scores.masked_fill(
            key_padding_mask.unsqueeze(1).unsqueeze(2), float('-inf')
        )
    if attn_mask is not None:
        # attn_mask 可能含有 0 或 False 表示屏蔽,也支持 True/1 保留
        attn_scores = attn_scores.masked_fill(attn_mask == 0, float('-inf'))
    return attn_scores

4.4 注意事项

  • masked_fill 要求掩码张量与被填充张量形状兼容(可广播)。在实际调用时,可以将 attn_mask 扩展为 (1, 1, L_q, L_k) 再传入,以避免形状不匹配。

  • 当同时使用两种掩码时,先应用 key_padding_mask 再应用 attn_mask 是合理的顺序,因为前者是“必须屏蔽”,后者是“额外模式”。


手写交叉注意力:query 来自 decoder,key/value 来自 encoder,写出完整 CrossAttention 类,并说明与 self-attention 的输入和掩码区别。

5.1 交叉注意力机制

在 Transformer 的 Decoder 中,交叉注意力层允许 Decoder 在生成每个目标词时,关注 Encoder 输出的所有源序列位置。其 Q 来自 Decoder 的上一自注意力输出,K 和 V 来自 Encoder 的最终输出。

5.2 与 Self-Attention 的主要区别

  • 输入来源:Self-Attention 的 Q、K、V 来自同一输入;Cross-Attention 的 Q 来自 Decoder,K、V 来自 Encoder。

  • 掩码:Cross-Attention 不需要因果掩码,因为 Decoder 在推理时可以并行地看到所有 Encoder 位置。只需要对 Encoder 输出进行 key_padding_mask。

  • 维度可能不同:Decoder 和 Encoder 的隐藏维度可能不同,需要通过投影层将 Encoder 的维度映射到 Decoder 的维度(通常使 K、V 投影的输出维度与 Q 相同)。

5.3 完整实现

class CrossAttention(nn.Module):
    def __init__(self, decoder_dim, encoder_dim, num_heads, dropout=0.1):
        super().__init__()
        self.num_heads = num_heads
        self.head_dim = decoder_dim // num_heads
        self.scale = math.sqrt(self.head_dim)

        # Q 投影:从 decoder 隐藏状态
        self.q_proj = nn.Linear(decoder_dim, decoder_dim, bias=False)
        # K, V 投影:从 encoder 隐藏状态,但输出维度与 decoder_dim 对齐
        self.k_proj = nn.Linear(encoder_dim, decoder_dim, bias=False)
        self.v_proj = nn.Linear(encoder_dim, decoder_dim, bias=False)
        self.out_proj = nn.Linear(decoder_dim, decoder_dim, bias=False)
        self.dropout = nn.Dropout(dropout)

    def forward(self, decoder_hidden, encoder_hidden, key_padding_mask=None):
        """
        decoder_hidden: (B, tgt_len, decoder_dim)
        encoder_hidden: (B, src_len, encoder_dim)
        key_padding_mask: (B, src_len) bool, True = pad
        """
        B, tgt_len, _ = decoder_hidden.shape
        src_len = encoder_hidden.size(1)

        # 线性投影
        q = self.q_proj(decoder_hidden)  # (B, tgt_len, decoder_dim)
        k = self.k_proj(encoder_hidden)  # (B, src_len, decoder_dim)
        v = self.v_proj(encoder_hidden)

        # 分头
        q = q.view(B, tgt_len, self.num_heads, self.head_dim).transpose(1, 2)
        k = k.view(B, src_len, self.num_heads, self.head_dim).transpose(1, 2)
        v = v.view(B, src_len, self.num_heads, self.head_dim).transpose(1, 2)

        # 注意力计算
        attn_scores = torch.matmul(q, k.transpose(-2, -1)) / self.scale

        if key_padding_mask is not None:
            attn_scores = attn_scores.masked_fill(
                key_padding_mask.unsqueeze(1).unsqueeze(2), float('-inf')
            )

        attn_weights = F.softmax(attn_scores, dim=-1)
        attn_weights = self.dropout(attn_weights)
        attn_output = torch.matmul(attn_weights, v)

        # 合并多头
        attn_output = attn_output.transpose(1, 2).contiguous().view(B, tgt_len, -1)
        output = self.out_proj(attn_output)
        return output

5.4 掩码差异

  • Self-Attention:通常需要因果掩码(在 Decoder 自注意力中)和 key_padding_mask。

  • Cross-Attention:只需要 key_padding_mask 来屏蔽 Encoder 的 padding;无需因果掩码。


使用 torch.einsum 实现缩放点积注意力计算,比较与 matmul 实现的异同。

6.1 einsum 实现

torch.einsum 使用爱因斯坦求和约定,通过下标字符串明确指定运算的维度,免去显式的转置和形状记忆。

def scaled_dot_product_einsum(query, key, value, scale=None, mask=None):
    """
    query: (B, n_heads, L_q, head_dim)
    key:   (B, n_heads, L_k, head_dim)
    value: (B, n_heads, L_k, head_dim)
    """
    if scale is None:
        scale = math.sqrt(query.size(-1))
    # 计算注意力分数: 'b h q d, b h k d -> b h q k'
    scores = torch.einsum('b h q d, b h k d -> b h q k', query, key) / scale
    if mask is not None:
        scores = scores.masked_fill(mask == 0, float('-inf'))
    weights = F.softmax(scores, dim=-1)
    # 加权求和: 'b h q k, b h k d -> b h q d'
    output = torch.einsum('b h q k, b h k d -> b h q d', weights, value)
    return output

6.2 与 matmul 的比较

  • 可读性:einsum 直观地表达了张量缩并的语义,尤其适合需要非标准缩并的场景(如批量矩阵乘法中的特定维度对齐)。matmul 则较为隐式,但它是标准的批量矩阵乘法,对于 (B, M, K) @ (B, K, N) 非常高效。

  • 性能:matmul 被底层库(cuBLAS)高度优化,通常比 einsum 快,尤其是大张量。einsum 在某些情况下可能导致临时内存分配和未优化的计算路径。在追求极致性能的生产代码中,推荐使用 matmulF.scaled_dot_product_attention

  • 灵活性:einsum 支持更复杂的张量操作,如可以同时进行转置和缩并,而 matmul 要求输入符合特定形状,必要时需先 transpose

6.3 实践建议

  • 在快速原型开发或教学场景,使用 einsum 可提升代码清晰度。

  • 在模型训练和推理部署中,优先使用 matmul 或 PyTorch 2.0+ 的 scaled_dot_product_attention,后者会自动选择最优实现(包括 Flash Attention)。


实现分组查询注意力(GQA):给定 num_query_heads=8,num_kv_heads=2,正确广播 KV 头至所有 query 头。

7.1 GQA 原理

分组查询注意力(Grouped Query Attention)是多查询注意力(MQA)和多头注意力(MHA)的折中。它将 query 头分成若干组,每组内的 query 头共享一对 K、V 头。这大幅减少了 KV Cache 的大小(从 h 倍减少到 g 倍),同时保持了比 MQA 更好的质量。

num_query_heads=8num_kv_heads=2 时,num_groups=8//2=4。每个 KV 头服务于连续的 4 个 query 头。

7.2 广播机制

  • K 和 V 的初始形状为 (B, num_kv_heads, seq_len, head_dim)

  • 通过 k.unsqueeze(2).expand(-1, -1, num_groups, -1, -1) 在每组内复制。

  • 然后 reshape(B, num_query_heads, seq_len, head_dim),直接与 Q 进行计算。

7.3 完整实现

class GroupedQueryAttention(nn.Module):
    def __init__(self, hidden_dim, num_query_heads, num_kv_heads, dropout=0.1):
        super().__init__()
        assert num_query_heads % num_kv_heads == 0
        self.hidden_dim = hidden_dim
        self.num_query_heads = num_query_heads
        self.num_kv_heads = num_kv_heads
        self.head_dim = hidden_dim // num_query_heads
        self.num_groups = num_query_heads // num_kv_heads

        self.q_proj = nn.Linear(hidden_dim, num_query_heads * self.head_dim, bias=False)
        self.k_proj = nn.Linear(hidden_dim, num_kv_heads * self.head_dim, bias=False)
        self.v_proj = nn.Linear(hidden_dim, num_kv_heads * self.head_dim, bias=False)
        self.out_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x, key_padding_mask=None, attn_mask=None):
        B, seq_len, _ = x.shape

        # 投影与分头
        q = self.q_proj(x).view(B, seq_len, self.num_query_heads, self.head_dim).transpose(1, 2)
        k = self.k_proj(x).view(B, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2)
        v = self.v_proj(x).view(B, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2)

        # 广播 KV 头:在组维度上复制
        # k: (B, 2, L, D) -> (B, 2, 1, L, D) -> (B, 2, 4, L, D) -> (B, 8, L, D)
        k = k.unsqueeze(2).expand(-1, -1, self.num_groups, -1, -1).contiguous()
        k = k.view(B, self.num_query_heads, seq_len, self.head_dim)
        v = v.unsqueeze(2).expand(-1, -1, self.num_groups, -1, -1).contiguous()
        v = v.view(B, self.num_query_heads, seq_len, self.head_dim)

        # 注意力计算
        scale = math.sqrt(self.head_dim)
        attn_scores = torch.matmul(q, k.transpose(-2, -1)) / scale

        if key_padding_mask is not None:
            attn_scores = attn_scores.masked_fill(
                key_padding_mask.unsqueeze(1).unsqueeze(2), float('-inf')
            )
        if attn_mask is not None:
            attn_scores = attn_scores.masked_fill(attn_mask == 0, float('-inf'))

        attn_weights = F.softmax(attn_scores, dim=-1)
        attn_weights = self.dropout(attn_weights)
        attn_output = torch.matmul(attn_weights, v)

        # 合并
        attn_output = attn_output.transpose(1, 2).contiguous().view(B, seq_len, -1)
        return self.out_proj(attn_output)

7.4 性能与内存考量

  • 通过共享 KV 头,KV Cache 的大小减小到原来的 num_kv_heads / num_query_heads = 1/4

  • unsqueezeexpand 过程中,并不分配额外内存(expand 返回视图),但需注意 contiguous() 在某些后续操作前可能需要。

  • GQA 被广泛用于 LLaMA 2/3 等大模型,以平衡推理效率和生成质量。


实现多查询注意力(MQA):将 KV 头数设为 1,分析相比 MHA 节省的 KV Cache 内存。

8.1 MQA 原理

多查询注意力将所有查询头共享同一对 Key 和 Value 头,而查询头保持独立。这极大减少了推理时需缓存的 K、V 张量的大小,是 MHA 向 GQA 演化的极致形态。MQA 在牺牲少量模型质量的前提下,显著提升了自回归解码的吞吐量。

8.2 代码实现

class MultiQueryAttention(nn.Module):
    def __init__(self, hidden_dim, num_heads, dropout=0.1):
        super().__init__()
        self.hidden_dim = hidden_dim
        self.num_heads = num_heads
        self.head_dim = hidden_dim // num_heads

        self.q_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
        self.k_proj = nn.Linear(hidden_dim, self.head_dim, bias=False)  # 只有1个头
        self.v_proj = nn.Linear(hidden_dim, self.head_dim, bias=False)
        self.out_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x, key_padding_mask=None, attn_mask=None):
        B, seq_len, _ = x.shape

        q = self.q_proj(x).view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        k = self.k_proj(x).unsqueeze(1)  # (B, 1, seq_len, head_dim)
        v = self.v_proj(x).unsqueeze(1)

        scale = math.sqrt(self.head_dim)
        attn_scores = torch.matmul(q, k.transpose(-2, -1)) / scale

        if key_padding_mask is not None:
            attn_scores = attn_scores.masked_fill(
                key_padding_mask.unsqueeze(1).unsqueeze(2), float('-inf')
            )
        if attn_mask is not None:
            attn_scores = attn_scores.masked_fill(attn_mask == 0, float('-inf'))

        attn_weights = F.softmax(attn_scores, dim=-1)
        attn_weights = self.dropout(attn_weights)
        attn_output = torch.matmul(attn_weights, v)

        attn_output = attn_output.transpose(1, 2).contiguous().view(B, seq_len, -1)
        return self.out_proj(attn_output)

8.3 KV Cache 内存节省分析

  • MHA 的 KV Cache:每层存储 2 * num_heads * head_dim * seq_len 个元素(K 和 V 各一份),FP16 下字节数为 4 * num_heads * head_dim * seq_len

  • MQA 的 KV Cache:2 * 1 * head_dim * seq_len 元素,即 4 * head_dim * seq_len 字节。

  • 节省比例:MQA 的 KV Cache 大小是 MHA 的 1 / num_heads。例如 32 个查询头时,节省 96.875% 的 KV Cache 内存。这允许在相同显存下支持更大的 batch size 或更长的序列。


在注意力分数上添加 ALiBi 线性偏置,编写生成头相关斜率的函数,并实现无位置编码的长序列注意力。

9.1 ALiBi 原理

ALiBi(Attention with Linear Biases)不依赖可学习的位置编码,而是直接在注意力分数上添加一个与距离成线性关系的负偏置。每个注意力头有一个独特的斜率,控制其对远距离依赖的惩罚程度。这使得模型可以外推至比训练时更长的序列。

斜率通常取几何级数,如对于 num_heads 个头,斜率为 2^(-8/num_heads * i) (i 从 1 开始)或类似方案。

9.2 代码实现

def get_alibi_slopes(num_heads):
    """生成每个头的ALiBi斜率"""
    def power_sequence(n, start=1):
        # 产生 2^(-8/n * i) 序列
        return torch.tensor([2**(-8 / n * (i + 1)) for i in range(n)])
    return power_sequence(num_heads)  # shape (num_heads,)

class AlibiSelfAttention(nn.Module):
    def __init__(self, hidden_dim, num_heads, dropout=0.1):
        super().__init__()
        self.hidden_dim = hidden_dim
        self.num_heads = num_heads
        self.head_dim = hidden_dim // num_heads
        self.slopes = get_alibi_slopes(num_heads)  # (num_heads,)

        self.qkv_proj = nn.Linear(hidden_dim, 3 * hidden_dim, bias=False)
        self.out_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x, key_padding_mask=None):
        B, seq_len, _ = x.shape
        qkv = self.qkv_proj(x)
        q, k, v = qkv.chunk(3, dim=-1)
        q = q.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        k = k.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        v = v.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)

        scale = math.sqrt(self.head_dim)
        attn_scores = torch.matmul(q, k.transpose(-2, -1)) / scale

        # 构建ALiBi偏置矩阵
        # 距离矩阵: (seq_len, seq_len)
        device = x.device
        positions = torch.arange(seq_len, device=device)
        dists = positions.unsqueeze(0) - positions.unsqueeze(1)  # (L, L)
        # 只惩罚未来位置(或双向均可,此处实现单向因果)
        causal_mask = torch.tril(torch.ones(seq_len, seq_len, device=device))
        alibi_bias = -self.slopes.view(-1, 1, 1) * dists.abs().unsqueeze(0)  # (H, L, L)
        alibi_bias = alibi_bias.masked_fill(causal_mask == 0, float('-inf'))

        attn_scores = attn_scores + alibi_bias.unsqueeze(0)  # (B, H, L, L)

        if key_padding_mask is not None:
            attn_scores = attn_scores.masked_fill(
                key_padding_mask.unsqueeze(1).unsqueeze(2), float('-inf')
            )

        attn_weights = F.softmax(attn_scores, dim=-1)
        attn_weights = self.dropout(attn_weights)
        attn_output = torch.matmul(attn_weights, v)

        attn_output = attn_output.transpose(1, 2).contiguous().view(B, seq_len, -1)
        return self.out_proj(attn_output)

9.3 长序列外推

由于 ALiBi 的偏置是线性的,模型在训练时习得每个头对距离的敏感度,推理时遇到更长序列,该线性偏置依然适用,无需额外修改。这相比 RoPE 等需要插值的方法更为简洁。


在自注意力中集成 RoPE(旋转位置编码):实现 apply_rotary_pos_emb 函数,并在计算 Q/K 后施加旋转。

10.1 RoPE 原理

RoPE 通过将每对相邻维度视为复数的实部和虚部,根据位置进行旋转,使得 Query 和 Key 的内积仅依赖于相对位置,而非绝对位置。实现时,先计算每个位置的 cossin 缓存,然后对 Q 和 K 的每个分块进行旋转变换。

10.2 代码实现

def precompute_rope_freqs(dim, seq_len, theta=10000.0):
    """计算旋转频率"""
    freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
    t = torch.arange(seq_len)
    freqs = torch.outer(t, freqs)  # (seq_len, dim//2)
    freqs_cis = torch.polar(torch.ones_like(freqs), freqs)  # 复数表示
    return freqs_cis

def apply_rotary_pos_emb(x, freqs_cis):
    """x: (B, H, L, D) 或 (*, L, D)"""
    # 将 x 转换为复数: (..., L, D) -> (..., L, D/2) 复数
    x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
    # 调整 freqs_cis 形状以广播
    freqs_cis = freqs_cis.unsqueeze(0).unsqueeze(0)  # (1, 1, L, D/2)
    x_rotated = x_complex * freqs_cis  # 复数乘法实现旋转
    # 转回实数
    x_out = torch.view_as_real(x_rotated).flatten(-2)
    return x_out.type_as(x)

# 在注意力层中使用
class RoPESelfAttention(MultiHeadSelfAttention):
    def forward(self, x, freqs_cis, key_padding_mask=None):
        # ... 投影、分头得到 q, k
        q = apply_rotary_pos_emb(q, freqs_cis)
        k = apply_rotary_pos_emb(k, freqs_cis)
        # 继续注意力计算...

10.3 关键点

  • RoPE 只应用于 Q 和 K,不应用于 V。

  • freqs_cis 可预先计算并缓存,推理时根据序列长度动态截取。

  • 对于长序列,RoPE 通常需要配合位置插值(PI)或 NTK 缩放来外推。


实现相对位置编码(如 T5 的 bucket 相对偏置),将可学习的相对偏置加入注意力分数。

11.1 T5 相对位置偏置原理

T5 将相对位置离散化为有限个 bucket,每个 bucket 对应一个可学习的标量偏置。这些偏置被加到注意力分数上,从而模型可以学习不同距离 token 之间的交互模式。

11.2 实现

class T5RelativeBias(nn.Module):
    def __init__(self, num_heads, num_buckets=32, max_distance=128):
        super().__init__()
        self.num_buckets = num_buckets
        self.max_distance = max_distance
        self.relative_bias = nn.Embedding(num_buckets, num_heads)

    @staticmethod
    def _relative_position_bucket(relative_position, num_buckets, max_distance):
        """T5 bucket 函数"""
        ret = 0
        n = -relative_position
        num_buckets //= 2
        ret += (n > max_distance).long() * num_buckets
        n = torch.where(n > max_distance, max_distance, n)
        # 对前后位置采用不同偏移
        return ret + (n * num_buckets // max_distance)

    def forward(self, seq_len, device):
        positions = torch.arange(seq_len, device=device)
        rel_pos = positions.unsqueeze(0) - positions.unsqueeze(1)  # (L, L)
        # 计算每个相对位置的 bucket 索引
        rp_bucket = self._relative_position_bucket(
            rel_pos, self.num_buckets, self.max_distance
        )
        # 取 embedding: (L, L, H) -> (H, L, L)
        values = self.relative_bias(rp_bucket).permute(2, 0, 1).unsqueeze(0)
        return values  # (1, H, L, L)

# 在注意力中使用
class T5Attention(MultiHeadSelfAttention):
    def __init__(self, hidden_dim, num_heads, num_buckets=32, max_distance=128, dropout=0.1):
        super().__init__(hidden_dim, num_heads, dropout)
        self.rel_bias = T5RelativeBias(num_heads, num_buckets, max_distance)

    def forward(self, x, key_padding_mask=None):
        B, seq_len, _ = x.shape
        # ... 得到 attn_scores: (B, H, L, L)
        rel_bias = self.rel_bias(seq_len, x.device)  # (1, H, L, L)
        attn_scores = attn_scores + rel_bias
        # 后续 softmax 等...

11.3 特点

  • T5 的相对偏置是跨层共享的,也有变体每层独立。

  • 离散化 bucket 减少了参数量,并支持超过训练长度的序列。

  • 相比绝对位置编码,相对偏置更好地捕捉 token 间的相对距离。


在注意力计算中引入 dropout,分别对注意力权重和输出进行 dropout,并注意训练/评估模式切换。

12.1 Dropout 作用

  • 注意力权重 dropout:在 softmax 之后、与 V 相乘之前,随机将部分注意力权重置零,防止模型过度依赖特定 token。

  • 输出 dropout:在输出投影之前(或之后)对多头合并后的特征进行 dropout,提升泛化性。

12.2 实现

class AttentionWithDropout(MultiHeadSelfAttention):
    def __init__(self, hidden_dim, num_heads, attn_dropout=0.1, output_dropout=0.1):
        super().__init__(hidden_dim, num_heads, dropout=0.0)  # 父类的 dropout 不用
        self.attn_dropout = nn.Dropout(attn_dropout)
        self.output_dropout = nn.Dropout(output_dropout)

    def forward(self, x, key_padding_mask=None, attn_mask=None):
        B, seq_len, _ = x.shape
        # ... 投影、分头、计算 attn_scores

        # 掩码
        if key_padding_mask is not None:
            attn_scores = attn_scores.masked_fill(
                key_padding_mask.unsqueeze(1).unsqueeze(2), float('-inf'))
        if attn_mask is not None:
            attn_scores = attn_scores.masked_fill(attn_mask == 0, float('-inf'))

        attn_weights = F.softmax(attn_scores, dim=-1)
        # 注意力权重 dropout
        attn_weights = self.attn_dropout(attn_weights)

        attn_output = torch.matmul(attn_weights, v)

        # 合并多头
        attn_output = attn_output.transpose(1, 2).contiguous().view(B, seq_len, -1)
        # 输出 dropout
        attn_output = self.output_dropout(attn_output)

        output = self.out_proj(attn_output)
        return output

12.3 注意事项

  • PyTorch 的 nn.Dropout 会根据 model.train() / model.eval() 自动切换行为。

  • 注意力权重 dropout 是 Transformer 原始实现的重要正则化手段。


手写缩放点积注意力的反向传播,手动推导并实现 backward 中关于 Q、K、V 的梯度计算。

13.1 反向传播推导

设前向过程为:

  1. S = Q @ K^T / sqrt(d)

  2. P = softmax(S) (行方向)

  3. O = P @ V

损失函数关于 O 的梯度为 dO,我们需要求 dQ, dK, dV

推导如下:

  • dV = P^T @ dO

  • dP = dO @ V^T

  • dS = P * (dP - rowsum(dP * P)) (softmax 反向传播的标准结果)

  • dQ = (dS / sqrt(d)) @ K

  • dK = (dS / sqrt(d))^T @ Q

13.2 手动实现

def scaled_dot_product_attention_backward(dO, Q, K, V, P, scale):
    """
    dO: (B, H, Lq, D) 损失对输出的梯度
    Q, K, V: 前向时的张量
    P: softmax 后的注意力权重
    """
    # dV
    dV = torch.matmul(P.transpose(-2, -1), dO)  # (B, H, Lk, D)

    # dP
    dP = torch.matmul(dO, V.transpose(-2, -1))  # (B, H, Lq, Lk)

    # softmax 反向
    # dS = P * (dP - sum(dP * P, dim=-1, keepdim=True))
    sum_term = torch.sum(dP * P, dim=-1, keepdim=True)
    dS = P * (dP - sum_term)

    # dQ, dK
    dQ = torch.matmul(dS / scale, K)  # (B, H, Lq, D)
    dK = torch.matmul(dS.transpose(-2, -1) / scale, Q)  # (B, H, Lk, D)

    return dQ, dK, dV

13.3 验证

可以使用 torch.autograd.gradcheck 对自定义的 Function 进行梯度检验,确保与自动微分的误差在容差范围内。


分析多头自注意力的时间与空间复杂度,写出计算 FLOPs 和显存占用的公式,并编写代码估算。

14.1 复杂度公式

  • 设 B: batch, H: heads, L: seq_len, D: head_dim。

  • 计算 QK^T: B * H * L * L * D FLOPs。

  • softmax: B * H * L * L FLOPs。

  • 加权求和: B * H * L * L * D FLOPs。

  • 总 FLOPs ≈ 2 * B * H * L^2 * D

  • 显存占用(中间激活):主要是 SP 矩阵,每个形状 (B, H, L, L),共 2 * B * H * L^2 个元素。

14.2 代码估算

python

def estimate_attention_cost(B, H, L, D, bytes_per_elem=2): flops = 2 * B * H * L * L * D memory_bytes = 2 * B * H * L * L * bytes_per_elem return flops, memory_bytes

对于 GPT-3 规模(B=8, H=96, L=2048, D=128, FP16),FLOPs 约 2 * 8 * 96 * 2048^2 * 128 ≈ 8.2e11 FLOPs,中间激活显存约 2 * 8 * 96 * 2048^2 * 2 ≈ 1.2 GB。这显示了长序列下注意力对显存的巨大压力,也是 FlashAttention 等优化的驱动力。


实现分块注意力:将序列切分为固定大小的 chunk,逐块计算注意力以减少峰值显存。

15.1 原理

将 Q 和 K、V 按序列长度切分成块,外层循环遍历 Q 块,内层循环遍历 K、V 块。每对块计算局部注意力,并用 online softmax 增量更新输出,避免一次性分配完整的 L x L 矩阵。

15.2 代码实现

def chunked_attention(Q, K, V, chunk_size=256, scale=None):
    """
    Q, K, V: (B, H, L, D)
    """
    B, H, L, D = Q.shape
    if scale is None:
        scale = math.sqrt(D)
    O = torch.zeros_like(Q)
    L = torch.zeros(B, H, L, 1, device=Q.device)  # 归一化累加器
    M = torch.full((B, H, L, 1), -float('inf'), device=Q.device)  # 最大值

    for i in range(0, L, chunk_size):
        Q_chunk = Q[:, :, i:i+chunk_size]  # (B, H, Qc, D)
        Qc = Q_chunk.shape[2]
        for j in range(0, L, chunk_size):
            K_chunk = K[:, :, j:j+chunk_size]  # (B, H, Kc, D)
            V_chunk = V[:, :, j:j+chunk_size]
            Kc = K_chunk.shape[2]

            scores = torch.matmul(Q_chunk, K_chunk.transpose(-2, -1)) / scale  # (B, H, Qc, Kc)

            # Online softmax
            M_prev = M[:, :, i:i+Qc]
            L_prev = L[:, :, i:i+Qc]
            M_cur = torch.max(scores, dim=-1, keepdim=True)[0]
            M_new = torch.maximum(M_prev, M_cur)

            # 修正旧累加值
            L_prev_corrected = L_prev * torch.exp(M_prev - M_new)
            L_cur = torch.exp(scores - M_new).sum(dim=-1, keepdim=True)
            L_new = L_prev_corrected + L_cur

            # 更新输出
            O[:, :, i:i+Qc] = O[:, :, i:i+Qc] * torch.exp(M_prev - M_new) + \
                               torch.matmul(torch.exp(scores - M_new), V_chunk)

            M[:, :, i:i+Qc] = M_new
            L[:, :, i:i+Qc] = L_new

    O = O / L  # 最终归一化
    return O

15.3 效果

峰值显存从 O(L^2) 降至 O(chunk_size^2),使长序列成为可能。


实现简化版 FlashAttention:使用分块 + online softmax 实现数值稳定的注意力,并支持反向计算。

FlashAttention 的关键在于分块计算和重计算。这里展示支持反向传播的简化版(不考虑异步内存拷贝等底层优化)。

class FlashAttention(torch.autograd.Function):
    @staticmethod
    def forward(ctx, Q, K, V, scale):
        B, H, L, D = Q.shape
        if scale is None:
            scale = math.sqrt(D)
        ctx.scale = scale
        chunk_size = min(256, L)
        O = torch.zeros_like(Q)
        L_acc = torch.zeros(B, H, L, 1, device=Q.device)
        M = torch.full((B, H, L, 1), -float('inf'), device=Q.device)

        for i in range(0, L, chunk_size):
            Qi = Q[:, :, i:i+chunk_size]
            for j in range(0, L, chunk_size):
                Kj = K[:, :, j:j+chunk_size]
                Vj = V[:, :, j:j+chunk_size]
                scores = torch.matmul(Qi, Kj.transpose(-2, -1)) / scale
                M_prev = M[:, :, i:i+chunk_size]
                L_prev = L_acc[:, :, i:i+chunk_size]
                M_cur = scores.max(dim=-1, keepdim=True)[0]
                M_new = torch.maximum(M_prev, M_cur)
                L_prev_corrected = L_prev * torch.exp(M_prev - M_new)
                L_cur = torch.exp(scores - M_new).sum(dim=-1, keepdim=True)
                L_new = L_prev_corrected + L_cur
                O[:, :, i:i+chunk_size] = O[:, :, i:i+chunk_size] * torch.exp(M_prev - M_new) + \
                    torch.matmul(torch.exp(scores - M_new), Vj)
                M[:, :, i:i+chunk_size] = M_new
                L_acc[:, :, i:i+chunk_size] = L_new
        O = O / L_acc
        ctx.save_for_backward(Q, K, V, O, L_acc, M)
        return O

    @staticmethod
    def backward(ctx, dO):
        Q, K, V, O, L, M = ctx.saved_tensors
        scale = ctx.scale
        B, H, Lq, D = Q.shape
        Lk = K.shape[2]
        dQ = torch.zeros_like(Q)
        dK = torch.zeros_like(K)
        dV = torch.zeros_like(V)
        chunk_size = min(256, Lq)

        for i in range(0, Lq, chunk_size):
            Qi = Q[:, :, i:i+chunk_size]
            dOi = dO[:, :, i:i+chunk_size]
            Oi = O[:, :, i:i+chunk_size]
            for j in range(0, Lk, chunk_size):
                Kj = K[:, :, j:j+chunk_size]
                Vj = V[:, :, j:j+chunk_size]
                scores = torch.matmul(Qi, Kj.transpose(-2, -1)) / scale
                Pij = torch.exp(scores - M[:, :, i:i+chunk_size])
                # 反向传播局部计算(简化的实现省略细节,实际应与前向对应)
                # ...
                # 这里仅示意结构,完整的反向公式见13题。
        # 由于实现完整手工反向较复杂,此处省略具体累加细节,实际项目推荐使用PyTorch自动微分。
        return dQ, dK, dV, None

在生产环境中,直接使用 PyTorch 2.0+ 的 scaled_dot_product_attention 即可获得 FlashAttention 加速,不需手动实现反向。


实现推理时的 KV Cache:在自回归生成中缓存历史 K、V,并编写单步解码的更新逻辑。

17.1 原理

推理时,每个生成步只需要计算最新 token 的 Q、K、V,并将新的 K、V 追加到缓存中。注意力计算时,Q 只查询所有缓存的 K、V。这避免了重复计算历史 token 的 K、V。

17.2 实现

class SelfAttentionWithCache(MultiHeadSelfAttention):
    def __init__(self, hidden_dim, num_heads, dropout=0.1):
        super().__init__(hidden_dim, num_heads, dropout)
        self.k_cache = None
        self.v_cache = None

    def init_cache(self, batch_size, max_len, device):
        self.k_cache = torch.zeros(batch_size, self.num_heads, max_len, self.head_dim, device=device)
        self.v_cache = torch.zeros_like(self.k_cache)

    def forward(self, x, step, key_padding_mask=None):
        B, seq_len, _ = x.shape  # seq_len 通常为 1
        # 计算 Q,K,V
        qkv = self.qkv_proj(x)
        q, k, v = qkv.chunk(3, dim=-1)
        q = q.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        k = k.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        v = v.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)

        # 更新缓存
        if step == 0:
            self.k_cache[:, :, :seq_len] = k
            self.v_cache[:, :, :seq_len] = v
        else:
            self.k_cache[:, :, step:step+seq_len] = k
            self.v_cache[:, :, step:step+seq_len] = v

        # 注意力计算:用当前的 q 与所有缓存的 k, v
        k_all = self.k_cache[:, :, :step+seq_len]
        v_all = self.v_cache[:, :, :step+seq_len]

        scale = math.sqrt(self.head_dim)
        attn_scores = torch.matmul(q, k_all.transpose(-2, -1)) / scale
        attn_weights = F.softmax(attn_scores, dim=-1)
        attn_output = torch.matmul(attn_weights, v_all)

        attn_output = attn_output.transpose(1, 2).contiguous().view(B, seq_len, -1)
        return self.out_proj(attn_output)

17.3 扩展

实际部署中,缓存通常按层和头管理,且常结合 PagedAttention 等技术优化显存利用。


实现 GQA 或 MQA 下的 KV Cache,确保缓存的 key/value 形状与 query 头数正确广播。

在 GQA/MQA 下,缓存的 K、V 头数少于 Q 头数,推理时需正确广播。

class GQAWithCache(nn.Module):
    def __init__(self, hidden_dim, num_query_heads, num_kv_heads, dropout=0.1):
        super().__init__()
        # ... 同 GQA 初始化
        self.k_cache = None
        self.v_cache = None

    def init_cache(self, batch_size, max_len, device):
        self.k_cache = torch.zeros(batch_size, self.num_kv_heads, max_len, self.head_dim, device=device)
        self.v_cache = torch.zeros_like(self.k_cache)

    def forward(self, x, step):
        B, seq_len, _ = x.shape
        # Q投影: 正常分头
        q = self.q_proj(x).view(B, seq_len, self.num_query_heads, self.head_dim).transpose(1, 2)
        # K,V投影: 分头到 num_kv_heads
        k = self.k_proj(x).view(B, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2)
        v = self.v_proj(x).view(B, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2)

        # 更新缓存
        if step == 0:
            self.k_cache[:, :, :seq_len] = k
            self.v_cache[:, :, :seq_len] = v
        else:
            self.k_cache[:, :, step:step+seq_len] = k
            self.v_cache[:, :, step:step+seq_len] = v

        # 广播 KV 到 query 头数
        k_all = self.k_cache[:, :, :step+seq_len]  # (B, KV_H, L, D)
        v_all = self.v_cache[:, :, :step+seq_len]
        # 复制
        k_all = k_all.unsqueeze(2).expand(-1, -1, self.num_groups, -1, -1)
        k_all = k_all.reshape(B, self.num_query_heads, -1, self.head_dim)
        v_all = v_all.unsqueeze(2).expand(-1, -1, self.num_groups, -1, -1)
        v_all = v_all.reshape(B, self.num_query_heads, -1, self.head_dim)

        scale = math.sqrt(self.head_dim)
        attn_scores = torch.matmul(q, k_all.transpose(-2, -1)) / scale
        attn_weights = F.softmax(attn_scores, dim=-1)
        attn_output = torch.matmul(attn_weights, v_all)
        attn_output = attn_output.transpose(1, 2).contiguous().view(B, seq_len, -1)
        return self.out_proj(attn_output)

18.1 关键点

  • 缓存保持原始 KV 头数,节省显存。

  • 每步解码时,将缓存的 K、V 动态广播至 query 头数,无需重复存储多份。

  • 实际部署中,广播可通过 torch.repeat_interleave 或在注意力核函数内部实现,以进一步优化。


编写可视化代码,绘制给定样本的多头注意力权重热力图,并标注头号和层号。

19.1 原理

可视化注意力权重可以直观地反映模型在处理序列时,各个头关注的不同模式,有助于分析模型是否学到了语法结构、长程依赖或局部注意力。通常取某个头的注意力权重矩阵(形状 (seq_len, seq_len)),绘制热力图,并标注层号和头号。

19.2 实现

import matplotlib.pyplot as plt
import seaborn as sns
import torch

def plot_attention_weights(attn_weights, layer_idx, head_idx, tokens=None, save_path=None):
    """
    attn_weights: (seq_len, seq_len) 或 (batch, seq_len, seq_len),这里取单个样本
    layer_idx: int,层索引
    head_idx: int,头索引
    tokens: list of str,可选的 token 标签
    """
    if attn_weights.dim() == 3:
        attn_weights = attn_weights[0]  # 取第一个 batch
    attn_weights = attn_weights.detach().cpu().numpy()

    plt.figure(figsize=(10, 8))
    ax = sns.heatmap(attn_weights, cmap='viridis', cbar=True,
                     xticklabels=tokens, yticklabels=tokens)
    ax.set_title(f'Attention Weights - Layer {layer_idx}, Head {head_idx}')
    ax.set_xlabel('Key Position')
    ax.set_ylabel('Query Position')
    if tokens:
        plt.xticks(rotation=90)
        plt.yticks(rotation=0)
    if save_path:
        plt.savefig(save_path, dpi=150, bbox_inches='tight')
    plt.show()

对于 Transformer 模型,可以在 forward hook 中捕获每一层的注意力权重,然后调用此函数进行可视化。通常选择特定输入样本,观察不同层的头在不同位置上的关注焦点,例如浅层可能关注局部词窗,深层关注语义相关词。

19.3 扩展

  • 可绘制多个头的子图矩阵,一次性展示一层内所有头的注意力模式。

  • 结合 token 标注(如输入文本分词后的列表),使热力图可解释性更强。


实现滑动窗口注意力:构造带状 mask,使每个位置只能关注前后窗口大小内的 token。

20.1 原理

滑动窗口注意力(Sliding Window Attention)限制每个 Query 位置只能关注 Key 位置中距离不超过 window_size 的部分(通常只关注左侧的 window_size 个 token,或同时关注左右)。这降低了长序列下注意力的计算量,并将复杂度从 O(L^2) 降低到 O(L * window_size)

20.2 实现

def create_sliding_window_mask(seq_len, window_size, device, bidirectional=False):
    """
    生成滑动窗口掩码,1 表示可关注,0 表示屏蔽。
    bidirectional: 是否双向关注(前后窗口)
    """
    mask = torch.zeros(seq_len, seq_len, device=device)
    if bidirectional:
        for i in range(seq_len):
            left = max(0, i - window_size)
            right = min(seq_len, i + window_size + 1)
            mask[i, left:right] = 1.0
    else:  # 单向(因果):只能看左侧 window_size 个 token
        mask = torch.tril(torch.ones(seq_len, seq_len, device=device))
        for i in range(seq_len):
            mask[i, :max(0, i - window_size)] = 0.0
    return mask  # 0/1

# 在注意力中使用
class SlidingWindowAttention(MultiHeadSelfAttention):
    def __init__(self, hidden_dim, num_heads, window_size, dropout=0.1):
        super().__init__(hidden_dim, num_heads, dropout)
        self.window_size = window_size

    def forward(self, x, key_padding_mask=None):
        B, seq_len, _ = x.shape
        qkv = self.qkv_proj(x)
        q, k, v = qkv.chunk(3, dim=-1)
        q = q.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        k = k.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        v = v.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)

        scale = math.sqrt(self.head_dim)
        attn_scores = torch.matmul(q, k.transpose(-2, -1)) / scale

        # 生成滑动窗口掩码
        window_mask = create_sliding_window_mask(seq_len, self.window_size, x.device)
        attn_scores = attn_scores.masked_fill(window_mask == 0, float('-inf'))

        if key_padding_mask is not None:
            attn_scores = attn_scores.masked_fill(
                key_padding_mask.unsqueeze(1).unsqueeze(2), float('-inf'))

        attn_weights = F.softmax(attn_scores, dim=-1)
        attn_weights = self.dropout(attn_weights)
        attn_output = torch.matmul(attn_weights, v)
        attn_output = attn_output.transpose(1, 2).contiguous().view(B, seq_len, -1)
        return self.out_proj(attn_output)

20.3 应用场景

  • 如 Mistral、Longformer 等模型使用滑动窗口注意力处理长序列。

  • 通常与全局注意力(如特定 token 可关注所有位置)结合,以保留长程交互能力。


在 FP16/BF16 混合精度下实现多头自注意力,正确使用 autocast 并处理 scale 问题。

21.1 原理

FP16 计算速度更快但精度有限,易导致数值溢出(尤其是 softmax 的指数操作)。混合精度训练通过 autocast 自动将部分操作(如 matmul)转为 FP16,同时保留关键操作(如 softmax、归一化)在 FP32。缩放因子 scale 通常在 FP16 下需要调整为 FP32 计算,以避免下溢。

21.2 实现

import torch.cuda.amp as amp

class MixedPrecisionAttention(MultiHeadSelfAttention):
    def forward(self, x, key_padding_mask=None, attn_mask=None):
        B, seq_len, _ = x.shape
        with amp.autocast(enabled=True):
            qkv = self.qkv_proj(x)
            q, k, v = qkv.chunk(3, dim=-1)

            q = q.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
            k = k.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
            v = v.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)

            # 缩放因子在 FP32 下计算,避免精度问题
            scale = self.head_dim ** 0.5
            # 将 Q、K 转为 FP32 进行 softmax 前的计算?
            # 实践中,autocast 会自动处理大多数情况,但为确保 softmax 稳定,
            # 注意力分数可以保持在 FP32。
            attn_scores = torch.matmul(q.float(), k.float().transpose(-2, -1)) / scale

            if key_padding_mask is not None:
                attn_scores = attn_scores.masked_fill(
                    key_padding_mask.unsqueeze(1).unsqueeze(2), float('-inf'))
            if attn_mask is not None:
                attn_scores = attn_scores.masked_fill(attn_mask == 0, float('-inf'))

            attn_weights = F.softmax(attn_scores, dim=-1)  # softmax 通常在 FP32
            attn_weights = self.dropout(attn_weights)

            attn_output = torch.matmul(attn_weights, v.float())
            attn_output = attn_output.transpose(1, 2).contiguous().view(B, seq_len, -1)
            output = self.out_proj(attn_output.to(x.dtype))
        return output

21.3 注意事项

  • autocast 上下文会自动将符合条件的操作转为 FP16(如 Linear、matmul),但 softmaxmasked_fill 等通常保留 FP32。

  • 若不显式转换为 FP32,BF16 的动态范围较大,通常更安全;但 FP16 下必须小心 softmax 溢出,因此示例中将 Q、K 显式转为 FP32 计算注意力分数。

  • 混合精度训练时,损失缩放(GradScaler)用于处理 FP16 的梯度下溢。


将 Q、K、V 的线性投影合并为一个大矩阵乘法,再拆分,编写对应代码并分析计算效率。

22.1 原理

将三个独立的线性层 W_q, W_k, W_v 合并为一个大的权重矩阵 W_qkv,通过一次矩阵乘法计算所有投影,再沿特征维度分割为 Q、K、V。这减少了 GPU kernel launch 次数,提升了并行度和内存带宽利用率。

22.2 实现

class MergedQKVAttention(MultiHeadSelfAttention):
    def __init__(self, hidden_dim, num_heads, dropout=0.1):
        super().__init__(hidden_dim, num_heads, dropout)
        # 合并投影矩阵
        self.qkv_proj = nn.Linear(hidden_dim, 3 * hidden_dim, bias=False)

    def forward(self, x, key_padding_mask=None, attn_mask=None):
        B, seq_len, _ = x.shape
        qkv = self.qkv_proj(x)  # (B, seq_len, 3*hidden_dim)
        q, k, v = qkv.chunk(3, dim=-1)  # 每个 (B, seq_len, hidden_dim)

        q = q.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        k = k.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        v = v.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)

        scale = math.sqrt(self.head_dim)
        attn_scores = torch.matmul(q, k.transpose(-2, -1)) / scale
        # ... 掩码、softmax、加权求和、合并、输出投影
        # 省略,同前

22.3 计算效率分析

  • FLOPs:合并后总计算量与三个独立线性层完全相同,因为都是 B * seq_len * hidden_dim * (3 * hidden_dim) 的矩阵乘法。

  • 并行效率:合并后只需一次 GEMM 调用,减少了 GPU 调度开销,提高了计算强度。尤其在小 batch 或短序列时,kernel launch 开销占比高,合并优势明显。

  • 内存访问:权重矩阵 W_qkv 只读取一次,而不是分别读取三次,减少了显存带宽消耗。

  • 现代框架:HuggingFace Transformers 和许多自研模型都采用此优化。


编写单元测试,验证手写多头注意力与 torch.nn.MultiheadAttention 在相同输入下的输出误差。

23.1 测试思路

构造相同输入,分别通过手写 Attention 和 PyTorch 官方 nn.MultiheadAttention,检查输出是否在容差内一致。需注意官方接口的 Q、K、V 输入形状和掩码格式略有不同(官方 expects (L, B, D) 形状,mask 格式也有差异)。

23.2 测试代码

import unittest

class TestMultiHeadAttention(unittest.TestCase):
    def setUp(self):
        self.B, self.L, self.D = 2, 16, 64
        self.H = 4
        self.x = torch.randn(self.B, self.L, self.D)
        self.attn_custom = MultiHeadSelfAttention(self.D, self.H, dropout=0.0)
        self.attn_official = nn.MultiheadAttention(self.D, self.H, dropout=0.0, batch_first=True)

        # 复制相同权重以公平对比
        # 自定义的 qkv_proj 权重需要按顺序复制给官方的 in_proj_weight
        with torch.no_grad():
            # 自定义的 qkv_proj.weight 形状: (3*D, D)
            # 官方的 in_proj_weight 形状: (3*D, D),可直接复制
            self.attn_official.in_proj_weight.copy_(self.attn_custom.qkv_proj.weight)
            self.attn_official.out_proj.weight.copy_(self.attn_custom.out_proj.weight)

    def test_output_equivalence(self):
        # 无掩码
        out_custom = self.attn_custom(self.x)
        out_official, _ = self.attn_official(self.x, self.x, self.x)
        self.assertTrue(torch.allclose(out_custom, out_official, atol=1e-5))

23.3 注意事项

  • 官方 nn.MultiheadAttentionbatch_first=True 时输入形状为 (B, L, D)

  • 如果有偏置,也需要复制。还需注意官方对 key_padding_mask 的解释(True 表示忽略),与自己实现一致。

  • dropout 设为 0 确保确定性。


实现注意力池化:利用一个可学习的 query token(如 CLS)对序列进行注意力加权池化。

24.1 原理

注意力池化(Attention Pooling)不同于平均池化或最大池化,它通过学习一个全局的查询向量 query,与序列中每个位置的 key 进行注意力计算,得到一组权重,再对 value 加权求和。这种方式可以关注序列中最重要的部分,类似于 BERT 的 CLS token 但可独立于下游任务微调。

24.2 实现

class AttentionPooling(nn.Module):
    def __init__(self, hidden_dim, num_heads=1):
        super().__init__()
        self.num_heads = num_heads
        self.head_dim = hidden_dim // num_heads
        # 可学习的 query 向量,形状 (1, 1, hidden_dim) 或 (1, num_heads, head_dim)
        self.query = nn.Parameter(torch.randn(1, 1, hidden_dim))
        self.key_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
        self.value_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
        self.out_proj = nn.Linear(hidden_dim, hidden_dim)

    def forward(self, x, mask=None):
        """
        x: (B, L, D)
        mask: (B, L) 可选
        """
        B, L, D = x.shape
        # 投影 K, V
        K = self.key_proj(x)  # (B, L, D)
        V = self.value_proj(x)
        # 扩展 query
        Q = self.query.expand(B, -1, -1)  # (B, 1, D)

        # 若多头,可 reshape,这里简化为单头或默认 num_heads=1
        # 计算注意力分数
        scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(D)  # (B, 1, L)
        if mask is not None:
            scores = scores.masked_fill(mask.unsqueeze(1), float('-inf'))
        attn_weights = F.softmax(scores, dim=-1)

        # 加权求和
        pooled = torch.matmul(attn_weights, V).squeeze(1)  # (B, D)
        return self.out_proj(pooled)

24.3 应用

  • 可用于替代平均池化,对句子或段落进行聚合。

  • 可配合多个 query token 实现多视角池化(类似于多个 CLS)。