1. 导入模块

import torch.nn.functional as F
import math
  • import torch.nn.functional as F:导入PyTorch的神经网络函数库,用于调用softmax等函数。

  • import math:导入数学库,用于计算平方根。

注意:代码中使用了torch.matmul等函数,因此实际运行时还需要import torch,但此处未显式写出,我们假设已导入。


2. 缩放点积注意力函数 

def attention(Q, K, V, mask=None):
"""缩放点积注意力"""

定义函数attention,接受查询张量Q、键张量K、值张量V和一个可选的掩码mask

文档字符串,该函数实现的是缩放点积注意力机制。

    scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(Q.size(-1))
  • torch.matmul, K.transpose(-2, -1)):计算查询和键的点积。Q的形状通常为(batch_size, ..., seq_len_q, d_k)K的形状为(batch_size, ..., seq_len_k, d_k)。对K的最后两个维度转置,得到形状(batch_size, ..., d_k, seq_len_k),然后矩阵乘法得到形状(batch_size, ..., seq_len_q, seq_len_k)的注意力分数矩阵。

  • / math.sqrt(Q.size(-1)):除以d_k的平方根进行缩放,防止点积过大导致softmax梯度消失。Q.size(-1)即为d_k

    if mask is not None:
        scores = scores.masked_fill(mask == 0, -1e9)

如果提供了掩码,则将scores中对应mask == 0的位置填充为一个极小的负数(-1e9)。这样在后续softmax中,这些位置的权重趋近于0,起到遮蔽作用。

    attention_weights = F.softmax(scores, dim=-1)

在最后一个维度(seq_len_k)上应用softmax,将分数转换为概率分布,得到注意力权重,形状与scores相同。

output = torch.matmul(attention_weights, V)

将注意力权重与值张量V相乘,得到加权求和后的输出。V的形状为(batch_size, ..., seq_len_k, d_v),乘法结果形状为(batch_size, ..., seq_len_q, d_v)。

    return output, attention_weights

返回注意力输出和注意力权重。


3. 多头注意力类 

class MultiHeadAttention(torch.nn.Module):

定义一个继承自torch.nn.Module的类,实现多头注意力机制。

    def __init__(self, d_model, n_heads):
        super().__init__()

构造函数,接受模型维度d_model和头数n_heads。首先调用父类构造函数初始化。

       self.d_model = d_model
       self.n_heads = n_heads

 保存型维度和头数。

        self.d_k = d_model // n_heads

计算每个头的维度,假设d_model能被n_heads整除。

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

创建三个线性变换层,分别用于将输入的查询、键、值从d_model维映射到d_model维。每个头的线性变换是共享的,但后续会通过重塑来分头处理。

   self.W_o = torch.nn.Linear(d_model, d_model)

创建输出线性变换层,用于将多头合并后的结果再映射回d_model维。

  def forward(self, query, key, value, mask=None):

定义前向传播函数,接受查询、键、值以及可选的掩码。通常输入形状为(batch_size, seq_len, d_model)。

       batch_size = query.size(0)

 获取批次大小,假设输入的第一维是batch维度。

  # 线性变换并重塑
        Q = self.W_q(query).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
  • self.W_q(query):对查询进行线性变换,输出形状(batch_size, seq_len_q, d_model)

  • .view(batch_size, -1, self.n_heads, self.d_k):重塑为(batch_size, seq_len_q, n_heads, d_k)

  • .transpose(1, 2):交换第1维(seq_len_q)和第2维(n_heads),得到形状(batch_size, n_heads, seq_len_q, d_k)。这样每个头独立处理。

  K = self.W_k(key).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)

 类似地对键进行处理,得到形状(batch_size, n_heads, seq_len_k, d_k)。

  V = self.W_v(value).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)

   类似地对值进行处理,得到形状(batch_size, n_heads, seq_len_k, d_k)。 

     # 注意力计算
     output, attention_weights = attention(Q, K, V, mask)

调用前面定义的attention函数,传入多头形式的Q、K、V和掩码。返回的输出形状为(batch_size, n_heads, seq_len_q, d_k),注意力权重形状为(batch_size, n_heads, seq_len_q, seq_len_k)

  # 合并多头
  output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
  • .transpose(1, 2):将维度从(batch_size, n_heads, seq_len_q, d_k)转换为(batch_size, seq_len_q, n_heads, d_k)

  • .contiguous():确保张量在内存中是连续的,因为transpose可能产生非连续张量,后续view操作需要连续性。

  • .view(batch_size, -1, self.d_model):将后两维合并,n_heads * d_k = d_model,得到形状(batch_size, seq_len_q, d_model)

return self.W_o(output), attention_weights

  • self.W_o(output):对合并后的输出应用输出线性变换,得到最终的多头注意力输出,形状为(batch_size, seq_len_q, d_model)

  • 同时返回注意力权重,用于可视化或分析。

4.注意力机制完整实现
 

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

def attention(Q, K, V, mask=None):
    """缩放点积注意力"""
    scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(Q.size(-1))
    
    if mask is not None:
        scores = scores.masked_fill(mask == 0, -1e9)
    
    attention_weights = F.softmax(scores, dim=-1)
    output = torch.matmul(attention_weights, V)
    
    return output, attention_weights

class MultiHeadAttention(torch.nn.Module):
    def __init__(self, d_model, n_heads):
        super().__init__()
        self.d_model = d_model
        self.n_heads = n_heads
        self.d_k = d_model // n_heads
        
        self.W_q = torch.nn.Linear(d_model, d_model)
        self.W_k = torch.nn.Linear(d_model, d_model)
        self.W_v = torch.nn.Linear(d_model, d_model)
        self.W_o = torch.nn.Linear(d_model, d_model)
    
    def forward(self, query, key, value, mask=None):
        batch_size = query.size(0)
        
        # 线性变换并重塑
        Q = self.W_q(query).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
        K = self.W_k(key).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
        V = self.W_v(value).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
        
        # 注意力计算
        output, attention_weights = attention(Q, K, V, mask)
        
        # 合并多头
        output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
        
        return self.W_o(output), attention_weights

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐