ALiBi 位置编码详细介绍

1. 背景

ALiBi(Attention with Linear Biases)位置编码是一种先进的相对位置编码技术,旨在提高模型在处理超出训练时最大序列长度的数据时的外推性能。传统的基于正弦函数的位置嵌入(sinusoidal position embedding)在模型外推时表现不佳,主要体现在当输入文本长度超过训练时的最大序列长度时,模型的困惑度(perplexity)急剧上升。

2. 核心思想

ALiBi 位置编码的核心思想是不给词向量加入位置嵌入向量,而是用一个与 query 和 key 之间的距离成比例的“惩罚项”来偏置 query-key 的注意力得分。具体来说,ALiBi 通过向自注意力机制的每个输入添加一个线性偏置项来实现这一目标。这个线性偏置项是基于输入的位置计算的,因此可以反映出输入之间的相对位置信息。

3. 实现方式

ALiBi 的实现主要包括以下几个步骤:

  1. 计算线性偏置:

    • 对于每个输入位置 ( i ) 和每个输出位置 ( j ),计算一个线性偏置 ( b_{ij} = i - j )。这个偏置反映了输入和输出之间的相对位置。
    • 例如,距离当前位置前1位为 -1,前2位为 -2,这些常数要乘上权重 ( m )。
  2. 添加偏置:

    • 将这个偏置添加到自注意力机制的输入上。具体来说,如果自注意力机制的输入是一个矩阵 ( X ),那么新的输入就是 ( X + b \cdot m )。
    • 例如,softmax(q_i K^T + m \cdot [- (i - 1), ..., -2, -1, 0])。
  3. 权重 ( m ) 的选择:

    • 权重 ( m ) 是训练前固定的头部特定斜率,其不同头部的值被选择为几何序列。例如,对于头部值 8,( m ) 可能是:
      • 第一个头部有一个相对较大的 ( m ),因此它更多地惩罚相距较远的标记,并专注于最近的标记。
      • 第八个头有最小的 ( m ),使其能够处理更远的标记。
4. 优势
  • 外推能力:ALiBi 位置编码显著提升了模型在处理超出训练时最大序列长度的数据时的外推性能。实验表明,ALiBi 在长文本生成任务中的表现优于传统的正弦位置编码和 RoPE。
  • 训练效率:ALiBi 可以加快 11% 的训练速度,同时减少 11% 的内存使用。
  • 灵活性:ALiBi 通过直接在注意力分数上添加偏置,能够有效地放大远距离的内积值,从而更好地处理长序列。
5. 实验结果
  • 长文本测试基准 LongBench:通过采用 ALiBi 作为位置编码的百川模型进行实验,在长文本测试基准 LongBench 上进行评估,结果显示 ALiBi 在长文本生成任务中的表现及其优势。
  • 外推能力对比:ALiBi 的外推能力优于正弦位置编码、RoPE 和 T5 偏置方法。具体来说,ALiBi 在处理更长序列时的困惑度(perplexity)显著低于其他方法。
6. 代码实现示例
import torch
import torch.nn as nn
import torch.nn.functional as F

class ALiBiPositionEncoding(nn.Module):
    def __init__(self, num_heads, max_seq_len):
        super(ALiBiPositionEncoding, self).__init__()
        self.num_heads = num_heads
        self.max_seq_len = max_seq_len
        self.slopes = torch.tensor([1.0 / (2 ** (i / (num_heads - 1))) for i in range(num_heads)])

    def forward(self, query, key):
        batch_size, seq_len, embed_dim = query.size()
        query = query.view(batch_size, seq_len, self.num_heads, -1).transpose(1, 2)
        key = key.view(batch_size, seq_len, self.num_heads, -1).transpose(1, 2)

        # 计算线性偏置
        relative_positions = torch.arange(seq_len).unsqueeze(0) - torch.arange(seq_len).unsqueeze(1)
        relative_positions = relative_positions.unsqueeze(0).expand(self.num_heads, -1, -1)
        biases = self.slopes.unsqueeze(-1).unsqueeze(-1) * relative_positions

        # 添加偏置
        scores = torch.einsum("bhld,bhmd->bhlm", query, key) + biases
        scores = scores / (embed_dim ** 0.5)

        return scores

class MultiHeadAttentionWithALiBi(nn.Module):
    def __init__(self, embed_dim, num_heads, max_seq_len):
        super(MultiHeadAttentionWithALiBi, self).__init__()
        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads
        self.q_proj = nn.Linear(embed_dim, embed_dim)
        self.k_proj = nn.Linear(embed_dim, embed_dim)
        self.v_proj = nn.Linear(embed_dim, embed_dim)
        self.out_proj = nn.Linear(embed_dim, embed_dim)
        self.alibi = ALiBiPositionEncoding(num_heads, max_seq_len)

    def forward(self, query, key, value, key_padding_mask=None):
        batch_size, seq_len, embed_dim = query.size()
        q = self.q_proj(query)
        k = self.k_proj(key)
        v = self.v_proj(value)
        q = q.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        k = k.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
        v = v.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2)

        scores = self.alibi(q, k)
        if key_padding_mask is not None:
            scores = scores.masked_fill(key_padding_mask.unsqueeze(1).unsqueeze(2), float("-inf"))
        attn_weights = F.softmax(scores, dim=-1)
        output = torch.einsum("bhlm,bhmd->bhld", attn_weights, v)
        output = output.transpose(1, 2).reshape(batch_size, seq_len, embed_dim)
        output = self.out_proj(output)

        return output

希望这些解释和代码示例能帮助你更好地理解 ALiBi 位置编码!如果有任何进一步的问题,请随时提问。

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐