
本文独家改进:全新的聚焦线性注意力模块(Focused Linear Attention),既高效又具有很强的模型表达能力,解决视觉Transformer计算量过大的问题,最终引入到YOLOv8,做到二次创新;

摘要:在Transformer模型应用于视觉领域的过程中,降低自注意力的计算复杂度是一个重要的研究方向。线性注意力通过两个独立的映射函数来近似Softmax操作,具有线性复杂度,能够很好地解决视觉Transformer计算量过大的问题。然而,目前的线性注意力方法整体性能不佳,难以实际应用。本文深入分析了现有线性注意力方法的缺陷,并提出了一个全新的聚焦的线性注意力模块(Focused Linear Attention),同时具有高效性和很强的模型表达能力。

提出了一种新型线性注意力模块,即聚焦的线性注意力。该模块借助聚焦函数获得更加聚焦的注意力分布,借助DWC模块保持特征多样性,既高效又具有很强的模型表达能力。总体来说,该模块具有以下几个优势:
(1) 计算复杂度低。通过改变自注意力机制的矩阵乘法顺序,本文提出的模块能够将计算复杂度降低为线性。此外,不同于以前一些线性注意力模块设计的复杂核函数,本文使用的聚焦函数和DWC的计算开销很小。
(2) 模型表达能力强。以前的线性注意力模块的性能通常不如Softmax注意力机制。但是,在使用聚焦函数和DWC解决两个性能瓶颈后,本文提出的聚焦的线性注意力可以获得比Softmax注意力机制更好的性能。
(3) 能够采用更大的感受野。得益于线性计算复杂度,本文的模块可以自然地采用更大的感受野,而不会增加模型计算量。例如,可以将Swin Transformer的window size由7扩大为56,即直接采用全局自注意力,而完全不引入额外计算量。
(4) 应用的灵活性强。本文提出的模块是常用的Softmax注意力的一个更优的替代, 可以作为一个插件模块应用于各种各样的ViT模型。
该方法在ImageNet上使DeiT、PVT、PVT-v2、Swin Transformer、CSwin Transformer等模型架构取得了显著的性能提升,能够将模型在CPU端加速约2.0倍,在GPU端加速约1.5倍。

核心代码:
class FocusedLinearAttention(nn.Module):
r""" Window based multi-head self attention (W-MSA) module with relative position bias.
It supports both of shifted and non-shifted window.
Args:
dim (int): Number of input channels.
window_size (tuple[int]): The height and width of the window.
num_heads (int): Number of attention heads.
qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set
attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0
proj_drop (float, optional): Dropout ratio of output. Default: 0.0
"""
def __init__(self, dim, window_size=[20, 20], num_heads=8, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0.,
focusing_factor=3, kernel_size=5):
super().__init__()
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
self.dim = dim
self.num_heads = num_heads
head_dim = dim // num_heads
self.focusing_factor = focusing_factor
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
self.attn_drop = nn.Dropout(attn_drop)
self.proj = nn.Linear(dim, dim)
self.proj_drop = nn.Dropout(proj_drop)
self.window_size = window_size
self.positional_encoding = nn.Parameter(torch.zeros(size=(1, window_size[0] * window_size[1], dim)))
self.softmax = nn.Softmax(dim=-1)
self.dwc = nn.Conv2d(in_channels=head_dim, out_channels=head_dim, kernel_size=kernel_size,
groups=head_dim, padding=kernel_size // 2)
self.scale = nn.Parameter(torch.zeros(size=(1, 1, dim)))
def forward(self, x, mask=None):
"""
Args:
x: input features with shape of (num_windows*B, N, C)
mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None
"""
# flatten: [B, C, H, W] -> [B, C, HW]
# transpose: [B, C, HW] -> [B, HW, C]
x = x.flatten(2).transpose(1, 2)
B, N, C = x.shape
qkv = self.qkv(x).reshape(B, N, 3, C).permute(2, 0, 1, 3)
q, k, v = qkv.unbind(0)
k = k + self.positional_encoding[:, :k.shape[1], :]
focusing_factor = self.focusing_factor
kernel_function = nn.ReLU()
q = kernel_function(q) + 1e-6
k = kernel_function(k) + 1e-6
scale = nn.Softplus()(self.scale)
q = q / scale
k = k / scale
q_norm = q.norm(dim=-1, keepdim=True)
k_norm = k.norm(dim=-1, keepdim=True)
if float(focusing_factor) <= 6:
q = q ** focusing_factor
k = k ** focusing_factor
else:
q = (q / q.max(dim=-1, keepdim=True)[0]) ** focusing_factor
k = (k / k.max(dim=-1, keepdim=True)[0]) ** focusing_factor
q = (q / q.norm(dim=-1, keepdim=True)) * q_norm
k = (k / k.norm(dim=-1, keepdim=True)) * k_norm
q, k, v = (rearrange(x, "b n (h c) -> (b h) n c", h=self.num_heads) for x in [q, k, v])
i, j, c, d = q.shape[-2], k.shape[-2], k.shape[-1], v.shape[-1]
z = 1 / (torch.einsum("b i c, b c -> b i", q, k.sum(dim=1)) + 1e-6)
if i * j * (c + d) > c * d * (i + j):
kv = torch.einsum("b j c, b j d -> b c d", k, v)
x = torch.einsum("b i c, b c d, b i -> b i d", q, kv, z)
else:
qk = torch.einsum("b i c, b j c -> b i j", q, k)
x = torch.einsum("b i j, b j d, b i -> b i d", qk, v, z)
num = int(v.shape[1] ** 0.5)
feature_map = rearrange(v, "b (w h) c -> b c w h", w=num, h=num)
feature_map = rearrange(self.dwc(feature_map), "b c w h -> b (w h) c")
x = x + feature_map
x = rearrange(x, "(b h) n c -> b n (h c)", h=self.num_heads)
x = self.proj(x)
x = self.proj_drop(x)
x = rearrange(x, "b (w h) c -> b c w h", b=B, c=self.dim, w=num, h=num)
return x关注私信获取源码
原创声明:本文系作者授权腾讯云开发者社区发表,未经许可,不得转载。
如有侵权,请联系 cloudcommunity@tencent.com 删除。