1 Sparse Attention

主要思想是认为多数情况下,长距离的注意力是少数的,所以削减远距离注意力,成为稀疏注意力机制。

如下为原本的注意力权重分布:

image-esjx.png

从注意力矩阵上看除了相对距离不超过k的、相对距离为k,2k,3k,…的注意力都设为0 ,这样一来Attention就具有“局部紧密相关和远程稀疏相关”的特性,如下为稀疏注意力机制的权重分布:

image-ytoc.png

# 简单的稀疏注意力实现,只计算Window内的attention​
def apply_sparse_attention(self, attention):   ​
    N, heads, query_len, key_len = attention.shape                  ​
    for i in range(N):             ​
        for j in range(heads):                 ​
            for q in range(query_len):                     ​
                start = max(0, q - self.window_size)                     ​
                end = min(key_len, q + self.window_size + 1)                     ​
                attention[i, j, q, :start] = 0                     ​
                attention[i, j, q, end:] = 0​
    return attention

2 Linear Attention

主要思想就是将softmax拿掉,然后先算K转置V,这样算法复杂度从 O(N^2d) 变为 O (Nd^2)

注意这里先先对QK使用不同层的softmax,从而达到先矩阵乘再softmax的效果

image-botd.png

def linear_attn(q, k, v, kv_mask = None):​
    dim = q.shape[-1]​
​
    if exists(kv_mask):​
        mask_value = max_neg_value(q)​
        mask = kv_mask[:, None, :, None]​
        k = k.masked_fill_(~mask, mask_value)​
        v = v.masked_fill_(~mask, 0.)​
        del mask​
​
    q = q.softmax(dim=-1)​
    k = k.softmax(dim=-2)​
​
    q = q * dim ** -0.5​
​
    context = einsum('bhnd,bhne->bhde', k, v)​
    attn = einsum('bhnd,bhde->bhne', q, context)​
    return attn.reshape(*q.shape)