36、CrissCrossAttention模块
论文《CROSSFORMER: A VERSATILE VISION TRANSFORMER HINGING ON CROSS-SCALE ATTENTION》
1、作用
CrossFormer通过跨尺度的特征提取和注意力机制,有效处理计算机视觉任务。它克服了现有视觉Transformer无法在不同尺度间建立有效交互的限制,提升了模型对图像的理解能力。
2、机制
1、跨尺度嵌入层(CEL):
通过不同尺度的内核采样图像补丁并将它们合并,为自注意力模块提供了跨尺度特征。
2、长短距离注意力(LSDA):
将自注意力模块分为短距离注意力(SDA)和长距离注意力(LDA)两部分,既降低了计算成本,又保留了不同尺度的特征。
3、动态位置偏差(DPB):
提出了一种动态位置偏差模块,使得相对位置偏差可以应用于不同大小的图像,提高了模型的灵活性和适用性。
3、独特优势
1、跨尺度交互:CrossFormer通过CEL和LSDA实现了特征在不同尺度间的有效交互,这对于处理具有不同尺寸对象的图像至关重要。
2、灵活性和适用性:通过动态位置偏差模块,CrossFormer能够适应不同尺寸的输入图像,提高了模型在各种视觉任务上的适用性。
3、优异的性能:广泛的实验表明,CrossFormer在图像分类、对象检测、实例分割和语义分割等任务上超越了其他最先进的视觉Transformer模型。
4、代码
import torch
import torch.nn as nn
from torch.nn import Softmax
# 定义一个无限小的矩阵,用于在注意力矩阵中屏蔽特定位置
def INF(B, H, W):
return -torch.diag(torch.tensor(float("inf")).repeat(H), 0).unsqueeze(0).repeat(B * W, 1, 1)
class CrissCrossAttention(nn.Module):
""" Criss-Cross Attention Module"""
def __init__(self, in_dim):
super(CrissCrossAttention, self).__init__()
# Q, K, V转换层
self.query_conv = nn.Conv2d(in_channels=in_dim, out_channels=in_dim // 8, kernel_size=1)
self.key_conv = nn.Conv2d(in_channels=in_dim, out_channels=in_dim // 8, kernel_size=1)
self.value_conv = nn.Conv2d(in_channels=in_dim, out_channels=in_dim, kernel_size=1)
# 使用softmax对注意力分数进行归一化
self.softmax = Softmax(dim=3)
self.INF = INF
# 学习一个缩放参数,用于调节注意力的影响
self.gamma = nn.Parameter(torch.zeros(1))
def forward(self, x):
m_batchsize, _, height, width = x.size()
# 计算查询(Q)、键(K)、值(V)矩阵
proj_query = self.query_conv(x)
proj_query_H = proj_query.permute(0, 3, 1, 2).contiguous().view(m_batchsize * width, -1, height).permute(0, 2, 1)
proj_query_W = proj_query.permute(0, 2, 1, 3).contiguous().view(m_batchsize * height, -1, width).permute(0, 2, 1)
proj_key = self.key_conv(x)
proj_key_H = proj_key.permute(0, 3, 1, 2).contiguous().view(m_batchsize * width, -1, height)
proj_key_W = proj_key.permute(0, 2, 1, 3).contiguous().view(m_batchsize * height, -1, width)
proj_value = self.value_conv(x)
proj_value_H = proj_value.permute(0, 3, 1, 2).contiguous().view(m_batchsize * width, -1, height)
proj_value_W = proj_value.permute(0, 2, 1, 3).contiguous().view(m_batchsize * height, -1, width)
# 计算垂直和水平方向上的注意力分数,并应用无穷小掩码屏蔽自注意
energy_H = (torch.bmm(proj_query_H, proj_key_H) + self.INF(m_batchsize, height, width)).view(m_batchsize, width, height, height).permute(0, 2, 1, 3)
energy_W = torch.bmm(proj_query_W, proj_key_W).view(m_batchsize, height, width, width)
# 在垂直和水平方向上应用softmax归一化
concate = self.softmax(torch.cat([energy_H, energy_W], 3))
# 分离垂直和水平方向上的注意力,应用到值(V)矩阵上
att_H = concate[:, :, :, 0:height].permute(0, 2, 1, 3).contiguous().view(m_batchsize * width, height, height)
att_W = concate[:, :, :, height:height + width].contiguous().view(m_batchsize * height, width, width)
# 计算最终的输出,加上输入x以应用残差连接
out_H = torch.bmm(proj_value_H, att_H.permute(0, 2, 1)).view(m_batchsize, width, -1, height).permute(0, 2, 3, 1)
out_W = torch.bmm(proj_value_W, att_W.permute(0, 2, 1)).view(m_batchsize, height, -1, width).permute(0, 2, 1, 3)
return self.gamma * (out_H + out_W) + x
if __name__ == '__main__':
block = CrissCrossAttention(64)
input = torch.rand(1, 64, 64, 64)
output = block(input)
print( output.shape) # 打印输出形状