神经网络学习笔记5——Swin-Transformer网络 系列文章目录神经网络学习笔记1——ResNet残差网络、Batch Normalization理解与代码神经网络学习笔记2——VGGNet神经网络结构与感受野理解与代码参考博客1参考博客2文章目录系列文章目录前言一、Patch Merging操作二、W-MSA、SW-MSA与cyclic shift窗口设计1、窗口自注意力W-MSA2、移动窗口自注意力SW-MSA3、窗口移动优化cyclic shift三、Swin Transformer Blocks模块四、整体模型理解五、相对位置偏置Relative Position Bias参数六、Swin transformer类别前言swin-transformer是什么Swim Transformer是特为视觉领域设计的一种分层Transformer结构。Swin 的两大特性是滑动窗口和分层表示。滑动窗口在局部不重叠的窗口中计算自注意力并允许跨窗口连接。分层结构允许模型适配不同尺度的图片并且计算复杂度与图像大小呈线性关系也因此被人成为披着transformer皮的CNN。Swin Transformer借鉴了CNN的分层结构不仅能够做分类还能够和CNN一样扩展到下游任务可以用于计算机视觉任务的通用主干网络可以用于图像分类、图像分割、目标检测等一系列视觉下游任务。它以VIT作为起点设计思想吸取了resnet的精华从局部到全局将transformer设计成逐步扩大感受野的工具它的成功背后绝不是偶然而是厚厚的积累与沉淀。解决了什么问题与NLP领域不同视觉领域同类的物体在不同图像上/同一图像上的尺度会相差巨大同一张多行人图中行人会有大有小有近有远同一语义下的不同目标可能尺度差距变化会很大相较于文本图像的尺寸过大计算复杂度较高相比之前的ViT做了两个改进引入CNN中常用的层次化构建方式构建层次化Transformer 引入locality思想对无重合的window区域内进行self-attention计算。优点是什么提出了一种层级式网络结构解决视觉图像的多尺度问题提供各个尺度的维度信息提出Shifted Windows移动窗口带来了更大的效率移动操作让相邻窗口得到交互极大降低了transformer的计算复杂度计算复杂度是线性增长而不是平方式增长可以广泛应用到所有计算机视觉领域结论是什么Transformer完全可以在各个领域取代CNN被人称为CV领域的新方向新时代效果怎么样在ImageNet上并非SOTA仅与EfficientNet的性能差不多swin-transformer的优点不是在于分类在分类上的提升不是太多而在检测、分割等下游任务中有巨大的提升.该论文是在2021年3月发表的一经发表就已在多项视觉任务中霸榜。一、Patch Merging操作ViT用的是16×16的patch size也就是16倍的下采样率从低到高这些token每个patch的尺寸并不会发生改变通过全局自注意力操作来实现全局建模可是面对多尺寸的目标的学习会较差一单一尺寸处理为主。且面对大图片时序列长度还是过大计算复杂度平方式递增。在密集预测型任务如检测和分割或者说在落地项目中使用的图片多尺度问题是很重要的问题成熟的模型都会有专门的多尺度特征处理方法。Swin transformer是在小窗口中进行自注意力窗口概念在第二块这些patch组成的小窗口和ViT的patch不同是相较独立的。比如4倍下采样中将特征图划分成了多个不相交的小窗口区域Multi-Head Self-Attention只在每个窗口patch内进行。面对不相交的窗口如何传递信息如何学习多尺度信息它提出了patch merging简单来说就是由小窗口patch合成大窗口patch增大感受野再通过序号选取的方式去提取出深度特征图模拟出一种类似池化的操作。详细来说就是通过一个Patch Merging层进行下采样如下图所示比如想下采样两倍先将四个小patch合成大patch再通过小patch身上的序号1、2、3、4进行提取提取的时候是每隔一个点选一个也就是选择同序号同样序号位置上的 patch 就会被 merge。经过提取之后原来的这个张量就变成了四个张量在深度方向进行concat拼接维度从h × w × c变为h/2 × w/2 × 4c然后在通过一个LayerNorm层。因为要类比CNN模式每次经过pooling后通道数只会翻倍所以这里也只想让他翻2倍而不是变成4倍所以紧接着又再做了一次操作就是在 c 的维度上用一个1x1的卷积或者全连接层把通道数降下来变成2c最后就得到了h/2 × w/2 × 2c的输出。即通过Patch Merging层后feature map的高和宽会减半深度会翻倍。二、W-MSA、SW-MSA与cyclic shift窗口设计1、窗口自注意力W-MSA像ViT的全局自注意力的计算会导致平方倍的复杂度同样当去做视觉里的下游任务尤其是密集预测型的任务或者说遇到非常大尺寸的图片时候这种全局算自注意力的计算复杂度对比卷积就会有很大算力差别。文章提出用窗口的方式去做自注意力也就是Windows Multi-head Self-AttentionW-MSAW-MSA模块是为了减少计算量。如下图所示左侧使用的是普通的Multi-head Self-AttentionMSA模块对于feature map中的每个像素或称作token或patch与Class序列在Self-Attention计算过程中需要和所有的像素去计算全局。但在图右侧将特征图拆分成一个个不重叠的window使用W-MSA模块时首先将feature map按照M×M例子中的M2大小划分成一个个Windows然后单独对每个Windows内部进行Self-Attention。假设在Swin transformer中输入224×224×3的图片那么一个patch的大小划分为4×4那么就有56×56个patch而每7个patch就组成一个窗口也就是一个窗口有7×7个patch一个224×224×3的图片会有8×864个窗口。原论文中有给出下面两个公式这里忽略了Softmax的计算复杂度h代表feature map的高度w代表feature map的宽度C代表feature map的深度M代表每个窗口Windows的大小patch为单位对比公式1和公式2虽然这两个公式前面这两项是一样的只有后面从 (hw) ^ 2变成了 M^2 * h * w看起来好像差别不大但其实如果仔细带入数字进去计算就会发现计算复杂的差距是相当巨大的因为这里的 hw 如果是56*56的话 M^2 其实只有49所以是相差了几十甚至上百倍。2、移动窗口自注意力SW-MSAtransformer初衷是理解上下文是一种信息的传递交互采用W-MSA模块时只会在每个窗口内进行自注意力计算所以窗口与窗口之间是无法进行信息传递的。为了解决这个问题作者引入了Shifted Windows Multi-Head Self-AttentionSW-MSA模块即进行偏移的W-MSA。根据左右两幅图对比能够发现窗口发生了偏移可以理解成窗口从左上角分别向右侧和下方各偏移了 M /2 个patch。在L1层使用的是W-MSAL11层使用的是SW_MSA在L1时每个窗口里的patch只能和同一个窗口里的patch相互学习而到了L11层时由于窗口的移动导致一些patch进入新的窗口这些带有上一层窗口信息的patch可以和别的带有上一层前窗口信息的patch相互学习。这就是跨窗连接cross window connection操作使得窗口与窗口之间有着交互。再结合合并Patch Merging操作在最后几层的时候每个patch已经与特征图绝大部分的patch有过交流也就是感受野已经很大了可以看见图片的绝大部分了。这些局部注意力信息最终会扩散到全局变相达到全局注意力的效果。简单来说就L11层中心的4×4窗口学习融合的信息是L1层四个窗口的信息因为中心的4×4窗口来源组成是L1层四个窗口的patchL1层四个窗口的patch经过W-MSA时就已经学习到所在窗口的信息带有自己窗口的信息。所以L11层中心的4×4窗口的学习就是L1层四个窗口的绝大部分相邻信息融合。3、窗口移动优化cyclic shift其实SW-MSA窗口的移动也存在问题虽然实现让窗口里的patch可以和其他窗口的patch相互通信交流到别的窗口的信息。可是移动的前后却带有一个问题就比如移动前L1是四个窗口每个窗口都是16个patch移动后的L11是九个窗口每个窗口大小不一分别是4\8\16个patch。有一种简单的做法就是补零比如把四个patch的窗口补多12个0补成16个patch的窗口格式这样4补12,8补8补完后就得到9个窗口再将9个窗口打成batch进行学习。虽然做法直白简单但是一个batch里的窗口从4个提升到9个实际上计算量提高了复杂度也提高了。Swin transformer提出利用掩码做一次循环移位cyclic shift具体的做法就是给L11的上方和左边的残缺窗口临时编号为A、B、C。把A、B、C残缺窗口移动到L11的下方和右边A是对角B、C是对面。将新建的特征图在划分为4个窗口其中原中心16patch窗口不变其他的残缺窗口拼接成新的16patch窗口。这种操作即实现了不同窗口的patch交流又不会像补零操作那样窗口增加计算复杂度提高。但是又产生新问题就是原中心16patch窗口是不变的里面的patch是本来就是像素意义上的邻居是有关系的可以两两相互做自注意力。可是对于另外3个拼接16patch窗口来说它们是来自不同区域的特征图如果它们之间做自注意力那么学习出来的特征可能是混乱的也就是说它们之间不能当做一个纯粹的窗口去做自注意力。如何处理拼接窗口Swin transformer提出利用掩码masked操作比如这里有一个已经进过移动拼接的14×14×3的特征图0号窗口占7×7个patch1号与3号是4×7个patch2号与6号是3×7个patch4号4×45号和7号是3×48号是3×3一共就是14×14个patch窗口从左上角分别向右侧和下方各偏移了 M /2 个patch。0号窗口是一个完整的窗口可以直接使用自注意力3号和6号是属于拼接窗口它自身不可以直接做自注意力。所以先执行前面的操作将3号方块和6号方块的patch提取出来拉长为一个向量A这个向量A中3号patch的值有4×728个6号patch的值有3×721个。再通过向量A进行转置操作得到向量B。通过向量A、B的矩阵乘法进行自注意力计算得到自注意力矩阵C矩阵C中可以具体区分成四种类型分别是向量A的所有3号patch值与向量B的所有3号patch值相乘3×3向量A的所有3号patch值与向量B的所有6号patch值相乘3×6向量A的所有6号patch值与向量B的所有3号patch值相乘6×3向量A的所有6号patch值与向量B的所有6号patch值相乘6×6其中3×3和6×6是符合自注意力理念的3×6和6×3是拼接的混乱值所以我们只需要3×3和6×6的数据而3×6和6×3是需要masked掉的。那么如何去处理3号方块6号方块的窗口呢Swin transformer提出一个巧妙的思路就是使用一个掩码模板矩阵D让矩阵C与矩阵D相加本来矩阵C里的那些数值是很小的值大概是0点以下的值3×3和6×6的数据加上0是不会变化的而3×6和6×3的数据加上-100则会变成一个很大的负数这是将这些值都进行softmax操作那么那些负数就会归为0剩下的也就是我们所需要的3×3和6×6数据。讲完了36窗口那么继续看看12窗口结合上面的思路可以发现12窗口和36窗口是很不一样的这个不一样产生于拉直向量A上。可以看见向量A里1号patch的值和2号patch的值是交错排序的这也导致转置向量B以及自注意矩阵C的变化。主要说说矩阵C的变化它依然是分成4种类型1×1,1×2,2×1,2×2但不再是集中化和区域化了而是一个横竖条纹围棋格式的矩阵这种变化也导致了掩码模板矩阵D的设计由于这种格式比较麻烦我就没有专门一个个画出来可以参考一下Swin transformer提供的掩码模板。至于4578窗口其实就是36和12窗口的合体我就画了一个拉直向量A的图具体可以自己去理解需要结合36和12的规律。具体的掩码模板在上图有。做完了多头自注意力后需要把拼接的特征图还原回去以保证它的相对位置不变语义信息不变。如果不还原的话那么循环轮到下一次Blocks模块时学习的W-MSA是混乱的学习SW-MSA时又将移动过的特征图继续拆分拼接向右下角拼接多轮下来学到的特征会越来越混乱特征图也会处于不停打乱的状态。三、Swin Transformer Blocks模块Swin Transformer Blocks有两种结构区别在于窗口多头自注意力的计算一个使用了W-MSA结构一个使用了SW-MSA结构。而且这两个结构是成对使用的先使用一个W-MSA结构再使用一个SW-MSA结构。所以堆叠Swin Transformer Block的次数都是偶数在整体模型里Swin Transformer Blocks下的×2、×6就是因为成对使用的意思。结合图片和公式进行前向模拟传入输入格式为[H/4n, W/4n, nC]的序列Z(l-1)。传入LayerNorm层归一化后执行窗口自注意力W-MSA操作。W-MSA输出与Z(l-1)相加输出为Z’l。再传入LayerNorm层归一化后执行MLP操作注意MLP Block输入的通道深度会×4输出再÷4。MLP输出与Z’l相加,输出l。传入输入格式为[H/4n, W/4n, nC]的序列Zl。传入LayerNorm层归一化后执行移动窗口自注意力SW-MSA操作。SW-MSA输出与Zl相加输出为Z’l。再传入LayerNorm层归一化后执行MLP操作。MLP输出与Z’l相加,输出l。四、整体模型理解前向过程输入一张大小为H×W×3大小的图片Images。执行patch partition也就是将图片划分为H/4×W/4×48个patchH/4×W/4×48H×W×3但不设置ClS token。执行Linear Embedding对每个像素的channel数据做线性变换也就是说要把向量的维度变成一个预先设置好的值CC值的设置与它的类型有关就行resnet有18\34\50\101\152类型这个C是一个超参数。即将图像的shape由 [H/4, W/4, 48]变成了 [H/4, W/4, C]。输入的数据[H/4, W/4]拉直变成H/4× W/4HW/16序列长度C变成每个token的向量维度。因为[HW/16]个patch的序列长度比较长比如说输入224×224×3的图片那么patch就划分为56×56×48那么序列长度就是3136patchC为96对比ViT的196patch来说太长了所以就引入了给予窗口的自注意力计算。每个窗口一般设为7×749个patch的序列长度。[H/4, W/4, C]序列输入Stage1的Block输出也是 [H/4, W/4, C]×2。构建层级式transformer来提取多尺度信息把Stage1的Block输出的[H/4, W/4, C]传入Patch Merging模块进行类似池化下采样操作。它执行Patch Merging操作后输出的值是[H/8, W/8, 4C]模拟卷积模型的通道深度翻倍效果进行1×1卷积将4C降为2C。[H/8, W/8, 2C]序列输入Stage2的Block输出也是[H/8, W/8, 2C]×2。把Stage3的Block输出的[H/8, W/8, 2C]传入Patch Merging模块执行Patch Merging操作后输出的值是[H/16, W/16, 8C]模拟卷积模型的通道深度翻倍效果进行1×1卷积将8C降为4C。[H/16, W/16, 4C]序列输入Stage3的Block输出也是[H/16, W/16, 4C]×6。把Stage4的Block输出的[H/16, W/16, 4C]传入Patch Merging模块执行Patch Merging操作后输出的值是[H/32, W/32, 16C]模拟卷积模型的通道深度翻倍效果进行1×1卷积将16C降为8C。[H/32, W/32, 8C]序列输入Stage4的Block输出也是[H/32, W/32, 8C]×2。以上就是整体骨干网络 如果是用于图片分类为了和卷积神经网络保持一致Swin Transformer这篇论文并没有像 ViT 一样使用 CLS token而是接上一个Layer Norm层、全局池化层以及全连接层得到最终输出。(作者这个图里并没有画因为 Swin Transformer的本意并不是只做分类它还会去做检测和分割所以说它只画了backbone的部分没有去画最后的分类头或者检测头)参考别的类似图五、相对位置偏置Relative Position Bias参数Swin transformer在做实验的时候表示做SW-MSA会比只做W-MSA好做相对位置rel. pos.会比做绝对位置abs. pos.和没有位置no pos.好。回到论文公式会发现它的自注意力公式之前我们讲的多了一个B操作这个B就是相对位置偏置。借用并参考这位博主的理解与图1、假设window的大小M×M是2×2patch计算window内的自注意力时先计算相对位置索引这里的索引并不是偏置B而是一个构成B的要素。2、这里要区分出一个概念就是绝对位置与相对位置相对位置是可以结合参考系的方式理解就是参考主体减去自身与参考客体计算而来。3、不同的参考主体patch都可以计算出一种对应的相对位置索引将每种计算出来的索引展平并拼接到一个矩阵矩阵大小为(M×)2。、我们可以观察矩阵里的相对位置索引分布规律就比如右边位置的位置概念红色patch右边是蓝色以红色为参考主体蓝色的相对位置是[0,-1]。又比如黄色patch的右边是绿色相对位置也会是[0,-1]。仔细观察后会发现上下左右左上左下右上右下等位置概念是相同的索引的。5、通过对行列做哈希变换将2D的相对位置索引变为1D以精简计算作者这里采用的哈希公式是(xM-1)×(2M-1)(yM-1)其中x为行y为列号M为窗口大小。6、将2D降为1D可以用相加或相乘等简单方法实现但是会出现重复值比如说红色patch右方是[0,-1]下方[-1,0]在2D时还是较明显的但是用相加-1或相乘0时就会出现不同输入却输出相同的干扰这时可以使用哈希算法来解决这个问题。7、实现相对位置的特有值相对位置索引总共有(2M-1)×(2M-1)种那么就可以随机生成(2M-1)*(2M-1)个随机相对位置偏置(nn.Parameter可学参数)根据相对位置索引去获取对应的相对位置偏置也就是公式里面的B进行多头自注意力的计算。六、Swin transformer类别win. sz. 7x7表示使用的窗口Windows的大小dim表示feature map的channel深度或者说token的向量长度head表示多头注意力模块中head的个数未完待续。。。