1 Sparse Attention
主要思想是认为多数情况下,长距离的注意力是少数的,所以削减远距离注意力,成为稀疏注意力机制。
如下为原本的注意力权重分布:

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

# 简单的稀疏注意力实现,只计算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的效果

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)