BERT多头自注意力机制代码实现详解:从理论到实践

1次阅读
没有评论

共计 2813 个字符,预计需要花费 8 分钟才能阅读完成。

image.webp

背景介绍

自注意力机制(Self-Attention)是 Transformer 架构的核心组件,也是 BERT 等预训练模型成功的关键。它通过计算输入序列中每个元素与其他元素的相关性,动态地学习上下文表示。相比传统的 RNN 和 CNN,自注意力机制能够更好地捕捉长距离依赖关系,并且具有天然的并行计算优势。

BERT 多头自注意力机制代码实现详解:从理论到实践

核心概念

Query、Key 和 Value 矩阵

  1. Query(Q):表示当前需要关注的元素,用于与其他元素的 Key 进行匹配。
  2. Key(K):表示其他元素的标识,用于与 Query 计算相似度。
  3. Value(V):包含实际的信息内容,根据相似度权重进行加权求和。

多头注意力

多头注意力(Multi-Head Attention)是将输入线性投影到多个子空间,每个子空间独立计算注意力,最后将结果拼接。这种方式可以让模型同时关注不同位置的不同特征。

代码实现

以下是一个完整的 BERT 多头自注意力机制的 PyTorch 实现,包含详细注释:

import torch
import torch.nn as nn
import math

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

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads

        # 线性变换层,用于生成 Q、K、V
        self.q_linear = nn.Linear(embed_dim, embed_dim)
        self.k_linear = nn.Linear(embed_dim, embed_dim)
        self.v_linear = nn.Linear(embed_dim, embed_dim)

        # 输出层和 dropout
        self.out_linear = nn.Linear(embed_dim, embed_dim)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x, mask=None):
        batch_size = x.size(0)

        # 生成 Q、K、V
        q = self.q_linear(x)
        k = self.k_linear(x)
        v = self.v_linear(x)

        # 分割多头
        q = q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
        k = k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
        v = v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)

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

        # 应用 mask(如有)if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)

        # 计算注意力权重
        attention = torch.softmax(scores, dim=-1)
        attention = self.dropout(attention)

        # 加权求和
        context = torch.matmul(attention, v)

        # 合并多头
        context = context.transpose(1, 2).contiguous().view(batch_size, -1, self.embed_dim)

        # 输出
        output = self.out_linear(context)

        return output, attention

# 示例用法
if __name__ == "__main__":
    embed_dim = 512
    num_heads = 8
    batch_size = 2
    seq_len = 10

    # 随机生成输入
    x = torch.randn(batch_size, seq_len, embed_dim)

    # 创建多头注意力层
    mha = MultiHeadAttention(embed_dim, num_heads)

    # 前向传播
    output, attention_weights = mha(x)

    print(f"输入形状: {x.shape}")
    print(f"输出形状: {output.shape}")
    print(f"注意力权重形状: {attention_weights.shape}")

常见问题

  1. 维度不匹配
  2. 确保 embed_dim 能被 num_heads 整除,否则会报错。
  3. 多头分割后,注意调整张量的形状和维度顺序。

  4. 权重初始化不当

  5. 线性变换层的权重需要合理初始化(如 Xavier 初始化)。
  6. 偏差项(bias)通常初始化为零。

  7. 注意力分数溢出

  8. 计算分数时,记得除以sqrt(head_dim),避免 softmax 后梯度消失。

  9. mask 应用错误

  10. 确保 mask 的形状与注意力分数匹配。
  11. 将 mask 中需要忽略的位置设为 0,并替换为极小的负值(如-1e9)。

优化建议

  1. 性能调优
  2. 使用 torch.einsum 替代 matmultranspose,减少显存占用。
  3. 开启 PyTorch 的自动混合精度(AMP),加速计算。

  4. 内存优化

  5. 对于长序列,考虑使用稀疏注意力或分块计算。
  6. 梯度检查点(Gradient Checkpointing)可以减少训练时的显存占用。

  7. 数值稳定性

  8. 在 softmax 前对分数做masked_fill,避免无效位置影响计算结果。
  9. 使用 log_softmax 替代softmax,提高数值稳定性。

实践建议

  1. 应用场景
  2. 文本分类、命名实体识别等任务中,可以直接使用 BERT 的多头注意力。
  3. 对于生成任务(如机器翻译),需要结合编码器 - 解码器注意力。

  4. 注意事项

  5. 训练时注意学习率设置,多头注意力对学习率比较敏感。
  6. 推理时可以通过缓存 Key 和 Value,减少重复计算。

延伸阅读

  1. Attention Is All You Need – Transformer 原论文
  2. BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding – BERT 原论文
  3. The Illustrated Transformer – 图解 Transformer

思考题

  1. 如何修改代码实现相对位置编码(Relative Positional Encoding)?
  2. 多头注意力中,为什么需要将 embed_dim 分成 num_heads 份?
  3. 如何实现跨语言的注意力机制(如翻译任务中的源语言和目标语言)?
正文完
 0
评论(没有评论)