以下是使用 Python 和 PyTorch 手动实现多头注意力机制(Multi-Head Attention)的详细代码及解释:

导入必要的库

python

运行

import torch
import torch.nn as nn

定义缩放点积注意力(Scaled Dot-Product Attention)函数

缩放点积注意力是多头注意力机制的核心部分,它计算 query、key 和 value 之间的注意力权重,并返回加权后的 value。

python

运行

def scaled_dot_product_attention(query, key, value, mask=None):
    d_k = query.size(-1)
    scores = torch.matmul(query, key.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtype=torch.float32))
    if mask is not None:
        scores = scores.masked_fill(mask == 0, float('-inf'))
    attention_weights = torch.softmax(scores, dim=-1)
    return torch.matmul(attention_weights, value)

代码解释

  • d_k 表示 query 和 key 的维度。
  • scores 通过矩阵乘法计算 query 和 key 的转置的乘积,然后除以 \(\sqrt{d_k}\) 进行缩放。
  • 如果提供了掩码(mask),将掩码中值为 0 的位置对应的分数设为负无穷,这样在 softmax 后这些位置的注意力权重接近 0。
  • attention_weights 使用 softmax 函数计算注意力权重。
  • 最后返回注意力权重与 value 的矩阵乘积。

定义多头注意力类(Multi-Head Attention)

python

运行

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, num_heads):
        super(MultiHeadAttention, self).__init__()
        self.num_heads = num_heads
        self.d_model = d_model
        self.d_k = d_model // num_heads
        assert d_model % num_heads == 0

        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)
        self.W_o = nn.Linear(d_model, d_model)

    def split_heads(self, x, batch_size):
        x = x.view(batch_size, -1, self.num_heads, self.d_k)
        return x.transpose(1, 2)

    def forward(self, query, key, value, mask=None):
        batch_size = query.size(0)
        Q = self.split_heads(self.W_q(query), batch_size)
        K = self.split_heads(self.W_k(key), batch_size)
        V = self.split_heads(self.W_v(value), batch_size)

        attention_output = scaled_dot_product_attention(Q, K, V, mask)
        attention_output = attention_output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
        return self.W_o(attention_output)

代码解释

  • 初始化函数 __init__
    • d_model 是输入和输出的维度。
    • num_heads 是头的数量,d_k 是每个头的维度,确保 d_model 能被 num_heads 整除。
    • 定义了四个线性变换层 W_qW_kW_v 和 W_o,分别用于将输入的 query、key、value 进行线性变换,以及将多头注意力的输出进行线性变换得到最终结果。
  • split_heads 方法:将线性变换后的张量按照头的数量进行拆分,调整维度顺序以便后续计算。
  • forward 前向传播函数
    • 对输入的 query、key、value 进行线性变换并拆分成多个头。
    • 调用 scaled_dot_product_attention 函数计算每个头的注意力输出。
    • 将各个头的输出进行维度调整和拼接,再通过 W_o 进行线性变换得到最终的多头注意力输出。

使用示例

python

运行

# 示例输入
batch_size = 2
seq_length = 3
d_model = 6
num_heads = 2
query = torch.rand(batch_size, seq_length, d_model)
key = torch.rand(batch_size, seq_length, d_model)
value = torch.rand(batch_size, seq_length, d_model)

multihead_attn = MultiHeadAttention(d_model, num_heads)
output = multihead_attn(query, key, value)
print(output.shape)

上述示例代码中,先定义了输入张量的形状,然后创建了 MultiHeadAttention 实例,并传入 query、key 和 value 进行计算,最后打印输出张量的形状。通过这样的手动实现,可以更清晰地理解多头注意力机制的工作原理和计算流程。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐