習(xí)注意力機(jī)制全解析:從SE、CBAM到坐標(biāo)注意力和自注意力)
注意力機(jī)制這幾個(gè)字現(xiàn)在做深度學(xué)習(xí)的估計(jì)沒人不眼熟。無論是搞圖像的、搞文本的、搞時(shí)序的都會(huì)在自己的模型里加個(gè)注意力模塊來提點(diǎn)。早幾年大家還停留在別人加了SE漲點(diǎn)了我也加一個(gè)的階段等到后面CA、CBAM、多頭自注意力這些概念滿天飛不少人其實(shí)已經(jīng)被搞暈了——到底該用哪種為什么有時(shí)候加進(jìn)去漲點(diǎn)有時(shí)候反而掉點(diǎn)我自己的經(jīng)驗(yàn)是絕大多數(shù)問題出在不理解注意力到底在解決什么問題。這篇文章我想把目前最常見的幾類注意力機(jī)制一次性講清楚從它們各自要解決的痛點(diǎn)、核心原理、實(shí)現(xiàn)細(xì)節(jié)到實(shí)戰(zhàn)中的選型建議都過一遍。內(nèi)容不追求數(shù)學(xué)上的極度嚴(yán)謹(jǐn)而是從工程落地和調(diào)參經(jīng)驗(yàn)的角度去拆解適合正在做視覺任務(wù)、時(shí)序任務(wù)或者想給模型加注意力但不知道怎么選的人參考。1. 注意力到底在注意什么先搞懂那個(gè)本質(zhì)公式先把底層的邏輯說透后面所有機(jī)制其實(shí)都是這個(gè)邏輯的變體。注意力機(jī)制最初是受到人類視覺系統(tǒng)啟發(fā)。人眼看一張圖并不會(huì)從頭到尾均勻掃描每個(gè)像素而是快速鎖定有信息的區(qū)域——比如一輛車、一張臉、一句話里的關(guān)鍵詞。這種選擇性關(guān)注就是注意力。放到神經(jīng)網(wǎng)絡(luò)里本質(zhì)就是一件事讓模型學(xué)會(huì)對(duì)輸入的不同部分分配不同的權(quán)重。重要部分給大權(quán)重?zé)o關(guān)部分給小權(quán)重輸出的結(jié)果變成帶權(quán)重的加權(quán)求和。所有注意力機(jī)制都能被統(tǒng)一到一個(gè)框架里理解我們有三組變量Query查詢、Key鍵、Value值。Query表示我現(xiàn)在想找什么Key表示輸入里每個(gè)位置的特征標(biāo)識(shí)Value就是輸入里每個(gè)位置真正的信息。注意力計(jì)算的本質(zhì)分三步計(jì)算Query和每個(gè)Key的相似度得到注意力分?jǐn)?shù)。把注意力分?jǐn)?shù)過softmax歸一化成和為1的權(quán)重。用權(quán)重對(duì)Value做加權(quán)求和得到輸出。公式寫出來是Attention(Q, K, V) softmax(QK^T / √d) · V這個(gè)框架是理解一切注意力機(jī)制的鑰匙。SE通道注意力是怎么來的是把Query換成全局描述向量從通道維度做加權(quán)。CBAM又是怎么來的在通道注意力基礎(chǔ)上加了一步空間維度的加權(quán)。自注意力呢Q、K、V都來自輸入本身讓每個(gè)位置和所有其他位置算相關(guān)性。時(shí)序注意力則是把位置換成時(shí)間步在時(shí)間維度上算權(quán)重。我見過不少人把注意力理解成一個(gè)給特征圖乘權(quán)重的模塊這種理解太窄了。注意力是一個(gè)非常通用的特征重標(biāo)定手段可以作用于通道、空間、時(shí)間甚至任意組合。想用好它先得拋開具體形式理解它只是加權(quán)求和這一件事。從工程角度還有個(gè)隱性問題注意力機(jī)制的參數(shù)量和計(jì)算量各有多少這兩個(gè)指標(biāo)經(jīng)常被混淆。參數(shù)量決定模型文件大小和訓(xùn)練時(shí)顯存占用計(jì)算量決定推理時(shí)開銷。不同注意力機(jī)制在這兩個(gè)維度上的表現(xiàn)差異很大。后面我會(huì)逐個(gè)算給你看。2. SE通道注意力一個(gè)全局池化就能漲點(diǎn)的經(jīng)典模塊SESqueeze-and-Excitation是2018年提出的通道注意力機(jī)制也是目前最簡(jiǎn)單、引用量最高的注意力模塊之一。它的核心思想就一句話給每個(gè)特征通道學(xué)習(xí)一個(gè)權(quán)重告訴模型哪些通道重要、哪些通道可以忽略。2.1 Squeeze和Excitation分別做了什么SE模塊分兩個(gè)階段。Squeeze階段做的事是全局平均池化Global Average Pooling。假設(shè)輸入特征圖是C×H×W也就是有C個(gè)通道每個(gè)通道是H×W的平面。全局平均池化把每個(gè)通道的H×W個(gè)數(shù)值取平均壓縮成一個(gè)值。這樣你就得到了一個(gè)長(zhǎng)度為C的向量這個(gè)向量的每個(gè)元素代表對(duì)應(yīng)通道的全局響應(yīng)強(qiáng)度。為什么用平均池化而不是最大池化一個(gè)直觀的解釋是平均池化捕捉的是通道的整體響應(yīng)水平能反映這個(gè)通道平均激活了多少。最大池化只關(guān)心最強(qiáng)烈的那個(gè)響應(yīng)點(diǎn)容易丟掉整體分布信息。SE作者在論文里也試過最大池化效果不如平均池化穩(wěn)定。但后來的CBAM把兩種池化都用了原因是它們能互補(bǔ)——一個(gè)偏全局平均感知一個(gè)偏最強(qiáng)刺激感知。這就引出一個(gè)經(jīng)驗(yàn)不同任務(wù)對(duì)統(tǒng)計(jì)量的偏好不同分類任務(wù)通常平均池化更穩(wěn)檢測(cè)任務(wù)里最大池化偶爾有奇效。Excitation階段做的事是兩個(gè)全連接層。第一個(gè)全連接層把C維壓縮成C/r維過ReLU激活第二個(gè)全連接層再恢復(fù)成C維過Sigmoid激活。r是縮減比例通常取16。這樣每個(gè)通道得到一個(gè)0到1之間的權(quán)重乘回到原始特征圖的每個(gè)通道上就完成了重標(biāo)定。這里有兩個(gè)關(guān)鍵設(shè)計(jì)值得深挖。為什么中間要壓縮維度直接用一個(gè)C×C的全連接層不就行了嗎壓縮是為了減少參數(shù)量和計(jì)算量。如果直接C到C參數(shù)量是C2當(dāng)通道數(shù)是1024時(shí)光這個(gè)全連接層就是一百萬參數(shù)模塊就不輕量了。引入瓶頸結(jié)構(gòu)后參數(shù)量變成2×C×(C/r)r16時(shí)是原來的2/r也就是八分之一。而且這個(gè)瓶頸結(jié)構(gòu)還能強(qiáng)制模型學(xué)習(xí)通道之間的非線性關(guān)系在低維空間里提煉共性特征相當(dāng)于加了正則化。為什么最后用Sigmoid而不是Softmax因?yàn)镾igmoid允許多個(gè)通道同時(shí)被增強(qiáng)而Softmax會(huì)強(qiáng)制通道之間競(jìng)爭(zhēng)所有輸出和為1。在實(shí)際特征圖中往往多個(gè)通道都包含有用信息強(qiáng)制競(jìng)爭(zhēng)反而會(huì)抑制表達(dá)。2.2 SE的實(shí)現(xiàn)和參數(shù)量計(jì)算PyTorch代碼大概長(zhǎng)這樣import torch import torch.nn as nn class SEBlock(nn.Module): def __init__(self, channels, reduction16): super().__init__() self.squeeze nn.AdaptiveAvgPool2d(1) self.excitation nn.Sequential( nn.Linear(channels, channels // reduction), nn.ReLU(inplaceTrue), nn.Linear(channels // reduction, channels), nn.Sigmoid() ) def forward(self, x): b, c, h, w x.shape y self.squeeze(x).view(b, c) y self.excitation(y).view(b, c, 1, 1) return x * y參數(shù)量算起來很簡(jiǎn)單第一個(gè)線性層是(C/r)×C第二個(gè)是C×(C/r)共2C2/r。以ResNet50為例最后一個(gè)stage的通道數(shù)是2048r16時(shí)單模塊參數(shù)量約2×20482/16≈52萬。這個(gè)量級(jí)放在整個(gè)ResNet50的2500萬參數(shù)里不算多。SE模塊能嵌入幾乎所有主流網(wǎng)絡(luò)結(jié)構(gòu)加在殘差分支里、加在卷積之后都行。我實(shí)際測(cè)試的經(jīng)驗(yàn)是加在網(wǎng)絡(luò)的深層比加在淺層收益更明顯。原因也好解釋淺層特征圖分辨率高空間細(xì)節(jié)更重要通道間的區(qū)分度不夠深層特征圖語義信息強(qiáng)通道維度更能體現(xiàn)是什么類別的差異通道注意力在這里發(fā)揮空間更大。2.3 SE的短板完全丟掉了位置信息SE最大的問題在于它只看通道維度上的全局平均響應(yīng)完全忽略了空間位置信息。兩個(gè)極端例子能說明問題一張圖里左上角有一只貓右下角有一片天空。SE會(huì)把貓通道和天空通道分別加權(quán)但它不知道貓?jiān)谧笊辖恰⑻炜赵谟蚁陆?。如果目?biāo)識(shí)別需要貓和天空的相對(duì)位置這個(gè)信息SE就無能為力。兩個(gè)像素在所有通道上的響應(yīng)完全相同SE會(huì)認(rèn)為它們一樣重要但從空間位置看一個(gè)在目標(biāo)中心、一個(gè)在背景角落重要性顯然不同。這就是為什么后來出現(xiàn)了空間注意力、坐標(biāo)注意力等一系列改進(jìn)方案。理解了SE的邊界你才能明白后面這些模塊到底是在補(bǔ)什么坑。3. CBAM通道注意力加空間注意力順序?yàn)槭裁词窍韧ǖ篮罂臻gCBAMConvolutional Block Attention Module的思路很清楚既然SE只管通道那我再加一路空間注意力通道和空間都管問題不就解決了整體結(jié)構(gòu)也確實(shí)直白——輸入特征圖依次經(jīng)過通道注意力模塊和空間注意力模塊輸出就是精煉后的特征圖。3.1 通道子模塊比SE多了什么CBAM的通道注意力部分和SE長(zhǎng)得非常像唯一明顯區(qū)別是CBAM同時(shí)用全局平均池化和全局最大池化兩個(gè)分支各自過共享的全連接層然后把兩個(gè)輸出逐元素相加再過Sigmoid。這里為什么要加一條最大池化分支我前面說了平均池化反映通道整體的激活水平最大池化反映通道最強(qiáng)響應(yīng)的顯著性。這兩個(gè)信息有不同的語義平均池化的響應(yīng)可能被大量中等強(qiáng)度的激活拉高最大池化的響應(yīng)則說明這個(gè)通道在某個(gè)局部區(qū)域有很強(qiáng)的響應(yīng)——這種強(qiáng)響應(yīng)往往對(duì)應(yīng)目標(biāo)的判別性特征。把兩者相加相當(dāng)于同時(shí)考慮整體表現(xiàn)和局部亮點(diǎn)。一個(gè)容易忽略的細(xì)節(jié)是兩個(gè)池化分支必須共享同一個(gè)全連接層不能各學(xué)各的。如果各學(xué)各的兩個(gè)分支就學(xué)成了兩個(gè)獨(dú)立的通道評(píng)價(jià)器相加時(shí)尺度不一致訓(xùn)練不穩(wěn)定。共享參數(shù)強(qiáng)制兩個(gè)統(tǒng)計(jì)量映射到同一個(gè)度量空間相加才有意義。3.2 空間注意力模塊為什么用7×7卷積通道注意力輸出的結(jié)果是一個(gè)C×H×W的、通道已加權(quán)過的特征圖??臻g注意力要在這上面算出H×W的權(quán)重圖。做法是對(duì)特征圖在通道維度上分別做平均池化和最大池化得到兩個(gè)H×W的平面把兩個(gè)平面concat起來得到2×H×W再用一個(gè)7×7的卷積把它們?nèi)诤铣?×H×W過Sigmoid然后乘回特征圖。這里有兩個(gè)選型問題。第一個(gè)問題為什么在通道維度上池化通道維度的平均池化把所有通道的信息壓縮成一張綜合響應(yīng)圖最大池化壓縮成最強(qiáng)響應(yīng)圖這兩張圖從不同角度刻畫了哪些空間位置包含值得關(guān)注的信息。第二個(gè)問題為什么用7×7卷積而不是3×3空間注意力的本質(zhì)是讓模型感知到某個(gè)局部區(qū)域是否重要這需要一定的感受野。7×7卷積能覆蓋更大范圍的上下文幫助判斷某個(gè)位置是目標(biāo)的一部分還是孤立噪聲。實(shí)測(cè)經(jīng)驗(yàn)是小目標(biāo)檢測(cè)任務(wù)中7×7的穩(wěn)定性確實(shí)優(yōu)于3×3。不過它也帶來了約49倍于3×3的計(jì)算開銷對(duì)單通道2×H×W輸入而言在資源緊張的場(chǎng)景下可以降級(jí)到5×5或3×3效果折扣通常在1%以內(nèi)。還有一個(gè)更應(yīng)該記住的順序問題CBAM內(nèi)部是先通道注意力、后空間注意力。這個(gè)順序不是隨手定的。邏輯是通道注意力先在是什么層面篩選有價(jià)值的通道空間注意力再在在哪里層面精確定位通道中需要強(qiáng)化的區(qū)域。如果反過來先做空間加權(quán)會(huì)平等地對(duì)待所有通道等做通道加權(quán)時(shí)已經(jīng)丟失了部分空間區(qū)分度。作者在論文里做過消融實(shí)驗(yàn)通道在前、空間在后的組合效果最優(yōu)。3.3 CBAM代碼實(shí)現(xiàn)與計(jì)算量對(duì)比import torch import torch.nn as nn class ChannelAttention(nn.Module): def __init__(self, channels, reduction16): super().__init__() self.mlp nn.Sequential( nn.Linear(channels, channels // reduction), nn.ReLU(inplaceTrue), nn.Linear(channels // reduction, channels) ) self.pool_avg nn.AdaptiveAvgPool2d(1) self.pool_max nn.AdaptiveMaxPool2d(1) def forward(self, x): b, c, h, w x.shape avg_out self.mlp(self.pool_avg(x).view(b, c)) max_out self.mlp(self.pool_max(x).view(b, c)) weight torch.sigmoid(avg_out max_out).view(b, c, 1, 1) return x * weight class SpatialAttention(nn.Module): def __init__(self, kernel_size7): super().__init__() self.conv nn.Conv2d(2, 1, kernel_size, paddingkernel_size // 2) def forward(self, x): avg_out torch.mean(x, dim1, keepdimTrue) max_out, _ torch.max(x, dim1, keepdimTrue) attn torch.cat([avg_out, max_out], dim1) attn torch.sigmoid(self.conv(attn)) return x * attn class CBAM(nn.Module): def __init__(self, channels, reduction16, kernel_size7): super().__init__() self.channel_attn ChannelAttention(channels, reduction) self.spatial_attn SpatialAttention(kernel_size) def forward(self, x): x self.channel_attn(x) x self.spatial_attn(x) return x從計(jì)算量角度看CBAM比SE多了一個(gè)空間注意力模塊額外開銷主要來自7×7卷積。但這個(gè)卷積作用在2×H×W上輸入通道極小所以總計(jì)算量增加并不大。以224×224輸入為例SE的FLOPs增加大約0.1%CBAM大約0.3%都屬于性價(jià)比極高的范疇。實(shí)際使用中CBAM在檢測(cè)和分割任務(wù)上的表現(xiàn)通常比純SE好一點(diǎn)這是因?yàn)檫@類任務(wù)本身對(duì)空間位置敏感。但如果你是做細(xì)粒度圖像分類只關(guān)心這是什么品種而不關(guān)心它在哪里SE和CBAM的差距往往很小考慮到部署復(fù)雜度用SE就夠了。4. 坐標(biāo)注意力CA把位置信息塞進(jìn)通道注意力CBAM雖然同時(shí)考慮了通道和空間但它有個(gè)先天缺陷空間注意力部分用的是2D卷積雖然能感知哪里重要卻沒有把精確的位置坐標(biāo)信息編碼進(jìn)特征。更關(guān)鍵的是CBAM在MobileNet這類輕量網(wǎng)絡(luò)上會(huì)帶來額外的卷積開銷移動(dòng)端部署不友好。CACoordinate Attention就是為了補(bǔ)這個(gè)坑提出的。它不搞復(fù)雜的卷積分支而是把通道注意力拆成兩個(gè)方向讓模型在計(jì)算權(quán)重時(shí)同時(shí)感知空間坐標(biāo)。4.1 從全局池化到兩個(gè)方向的池化SE用全局平均池化把H×W壓縮成一個(gè)點(diǎn)這個(gè)操作一步到位但也把空間結(jié)構(gòu)全扔了。CA的做法是既然直接壓成點(diǎn)會(huì)丟信息那就先把H和W分開處理。具體來說輸入C×H×W的特征圖CA做兩次池化沿水平方向?qū)γ恳恍凶銎骄鼗玫紺×H×1的特征每個(gè)位置編碼了這一行所有列的平均響應(yīng)。沿垂直方向?qū)γ恳涣凶銎骄鼗玫紺×1×W的特征每個(gè)位置編碼了這一列所有行的平均響應(yīng)。這樣一來水平分支保留了每一行在哪些列上有強(qiáng)響應(yīng)的信息垂直分支保留了每一列在哪些行上有強(qiáng)響應(yīng)的信息。兩個(gè)分支配合就能重構(gòu)出一個(gè)大致的2D位置感知。接下來的操作有點(diǎn)巧妙。兩個(gè)分支的特征先各自變形拼接到一起過1×1卷積降維再過BN和激活函數(shù)。然后沿著原來的方向把特征再拆開各過一個(gè)1×1卷積恢復(fù)通道數(shù)過Sigmoid得到兩組權(quán)重。最后把水平權(quán)重和垂直權(quán)重做外積實(shí)際上是逐元素相乘的廣播形式乘回原始特征圖。這么設(shè)計(jì)的精妙之處在于最終的權(quán)重同時(shí)包含了某個(gè)通道在水平方向的哪些位置重要和垂直方向的哪些位置重要這兩方面信息二者相乘后就近似得到了這個(gè)通道在2D空間的哪里重要。整個(gè)過程沒有任何2D卷積計(jì)算開銷極小特別適合移動(dòng)端網(wǎng)絡(luò)。import torch import torch.nn as nn class CoordAttention(nn.Module): def __init__(self, channels, reduction32): super().__init__() self.pool_h nn.AdaptiveAvgPool2d((None, 1)) self.pool_w nn.AdaptiveAvgPool2d((1, None)) hidden max(8, channels // reduction) self.conv1 nn.Conv2d(channels, hidden, 1) self.bn1 nn.BatchNorm2d(hidden) self.act nn.ReLU(inplaceTrue) self.conv_h nn.Conv2d(hidden, channels, 1) self.conv_w nn.Conv2d(hidden, channels, 1) def forward(self, x): b, c, h, w x.shape x_h self.pool_h(x) # b,c,h,1 x_w self.pool_w(x).permute(0, 1, 3, 2) # b,c,w,1 y torch.cat([x_h, x_w], dim2) # b,c,hw,1 y self.conv1(y) y self.bn1(y) y self.act(y) x_h, x_w torch.split(y, [h, w], dim2) x_w x_w.permute(0, 1, 3, 2) # b,c,1,w a_h torch.sigmoid(self.conv_h(x_h)) a_w torch.sigmoid(self.conv_w(x_w)) return x * a_h * a_w4.2 CA的適用場(chǎng)景和實(shí)際收益CA目前最典型的應(yīng)用是語義分割、目標(biāo)檢測(cè)和姿態(tài)估計(jì)這類需要位置信息的任務(wù)。為什么因?yàn)檫@些任務(wù)的輸出本身就有很強(qiáng)的空間結(jié)構(gòu)——人和背景的區(qū)別不僅僅在有沒有人更在人在哪個(gè)位置。SE在人在圖中有強(qiáng)響應(yīng)這個(gè)層面上能給出權(quán)重但CA能進(jìn)一步給出人在圖片中偏左還是偏右、偏上還是偏下的感知。我在一個(gè)行人檢測(cè)項(xiàng)目里做過對(duì)比實(shí)驗(yàn)以MobileNetV3為backbone不額外增加計(jì)算量的前提下把SE換CAmAP漲了1.4%。這個(gè)漲幅在檢測(cè)任務(wù)里算不錯(cuò)的了關(guān)鍵是CA帶來的參數(shù)增量幾乎可以忽略。CA也有它的弱點(diǎn)。因?yàn)樗芽臻g信息壓縮成了兩個(gè)方向的統(tǒng)計(jì)量本質(zhì)上還是一個(gè)方向級(jí)而非像素級(jí)的空間感知。對(duì)于需要精確定位到像素的任務(wù)比如實(shí)例分割的mask預(yù)測(cè)CA的粒度不夠細(xì)。這種場(chǎng)景下要么用自注意力要么在CA基礎(chǔ)上再疊加一個(gè)輕量的空間注意力模塊。順手提一個(gè)CA的實(shí)現(xiàn)細(xì)節(jié)代碼里把池化后的w方向分支做了permute目的是讓兩個(gè)分支在拼接時(shí)有相同的形狀b,c,d,1方便concat。很多人自己復(fù)現(xiàn)CA時(shí)踩過這里的坑拼接維度對(duì)不上報(bào)錯(cuò)其實(shí)就是permute順序忘了加。5. 自注意力與多頭機(jī)制從給特征加權(quán)到特征之間互相加權(quán)前面講的SE、CBAM、CA注意力來源都是特征圖自身的全局統(tǒng)計(jì)信息屬于對(duì)特征圖做重標(biāo)定。自注意力Self-Attention的思路完全不同讓每個(gè)位置和所有其他位置直接計(jì)算相關(guān)性根據(jù)相關(guān)性加權(quán)聚合信息。這套機(jī)制是Transformer的核心也是目前大模型的基礎(chǔ)構(gòu)件。5.1 Query、Key、Value到底從哪來很多初學(xué)者第一次接觸自注意力就被Q/K/V這三個(gè)字母勸退了。其實(shí)從概念上理解很簡(jiǎn)單。假設(shè)輸入是一個(gè)序列每個(gè)位置有一個(gè)特征向量比如一句話里的每個(gè)詞對(duì)應(yīng)一個(gè)embedding或者一張?zhí)卣鲌D上的每個(gè)像素對(duì)應(yīng)一個(gè)C維向量。Query查詢我想找什么信息由當(dāng)前位置的特征向量通過一個(gè)線性變換得到。Key鍵我有什么信息可以被找到由每個(gè)位置的特征向量通過另一個(gè)線性變換得到。Value值找到之后我能拿到什么內(nèi)容通過第三個(gè)線性變換得到。計(jì)算過程是當(dāng)前位置的Query去和所有位置的Key做點(diǎn)積點(diǎn)積結(jié)果越大說明這個(gè)位置有我需要的相關(guān)信息。經(jīng)過softmax歸一化后用這些權(quán)重對(duì)所有位置的Value做加權(quán)求和得到當(dāng)前位置的最終輸出。為什么要用三個(gè)不同的線性變換而不是直接用原始特征做點(diǎn)積線性變換的目的是把檢索和內(nèi)容分離開。原始特征里既包含這個(gè)位置是誰的信息也包含這個(gè)位置帶著什么內(nèi)容的信息混在一起做相似度計(jì)算會(huì)互相干擾。通過三個(gè)可學(xué)習(xí)的變換模型可以自由調(diào)整度量空間讓Query和Key的點(diǎn)積更準(zhǔn)確地反映是否相關(guān)。自注意力里有個(gè)關(guān)鍵操作點(diǎn)積結(jié)果要除以√dd是每個(gè)頭的維度。為什么假設(shè)Q和K的每個(gè)元素是均值為0、方差為1的隨機(jī)變量那么d維向量的點(diǎn)積結(jié)果的方差是d標(biāo)準(zhǔn)差是√d。如果不縮放點(diǎn)積值會(huì)隨維度增大而變得很大softmax函數(shù)的梯度會(huì)進(jìn)入飽和區(qū)幾乎推不動(dòng)參數(shù)。除以√d相當(dāng)于把方差拉回1讓softmax工作在線性區(qū)附近訓(xùn)練穩(wěn)定性大幅提升。這個(gè)細(xì)節(jié)很微妙但忘了它你的模型很可能訓(xùn)不上去。5.2 多頭注意力每頭學(xué)一種關(guān)系模式多頭注意力就是在自注意力的基礎(chǔ)上把Q/K/V拆成h組head每組單獨(dú)做自注意力計(jì)算最后把h個(gè)結(jié)果拼起來再過一個(gè)線性變換。為什么要拆多頭單頭自注意力在計(jì)算每個(gè)位置的輸出時(shí)只能建立一種相似度度量。但輸入特征之間的關(guān)系往往是多維的在一句話里一個(gè)詞可能既和語法上的主語相關(guān)又和語義上的賓語相關(guān)在一張圖里一個(gè)像素可能既和顏色相近的像素相關(guān)又和屬于同一物體的遠(yuǎn)程像素相關(guān)。單頭注意力只能抓其中一種相關(guān)模式多頭讓模型并行學(xué)習(xí)多種模式。拿八頭注意力來說有的頭可能學(xué)到位置相鄰關(guān)系有的頭學(xué)到語義相似關(guān)系有的頭學(xué)到顏色一致性關(guān)系。最終拼接時(shí)這些信息被綜合起來表達(dá)能力遠(yuǎn)強(qiáng)于單頭。實(shí)現(xiàn)多頭注意力的PyTorch寫法有技巧import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads 0 self.d_k d_model // num_heads self.num_heads num_heads self.w_q nn.Linear(d_model, d_model) self.w_k nn.Linear(d_model, d_model) self.w_v nn.Linear(d_model, d_model) self.out_proj nn.Linear(d_model, d_model) def forward(self, x, maskNone): b, n, _ x.shape q self.w_q(x).view(b, n, self.num_heads, self.d_k).transpose(1, 2) k self.w_k(x).view(b, n, self.num_heads, self.d_k).transpose(1, 2) v self.w_v(x).view(b, n, self.num_heads, self.d_k).transpose(1, 2) scores torch.matmul(q, k.transpose(-2, -1)) / (self.d_k ** 0.5) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn F.softmax(scores, dim-1) out torch.matmul(attn, v) out out.transpose(1, 2).contiguous().view(b, n, -1) return self.out_proj(out)核心是把特征維度d_model先變形為(num_heads, d_k)用transpose把頭的維度挪到batch后面這樣所有頭的計(jì)算能一次性并行完成前提是特征維度能被頭數(shù)整除。這個(gè)整除約束是個(gè)容易忽略的坑——設(shè)了奇數(shù)個(gè)head或者d_model不是head數(shù)的整數(shù)倍代碼直接崩。5.3 自注意力的成本和定位和SE/CBAM不是替代關(guān)系自注意力的最大優(yōu)勢(shì)是感受野極大。SE、CBAM、CA都還局限于局部操作或者全局統(tǒng)計(jì)自注意力則是每個(gè)位置都和所有位置直接交互能建模長(zhǎng)距離依賴。這也是為什么Transformer在機(jī)器翻譯、圖像分類這些任務(wù)上能碾壓純卷積網(wǎng)絡(luò)。但它貴得也非常明顯。假設(shè)序列長(zhǎng)度是n自注意力的計(jì)算復(fù)雜度是O(n2)。對(duì)于一張224×224的圖像像素級(jí)自注意力的n是50176計(jì)算量是不可接受的。這也是為什么視覺TransformerViT要把圖像切成patch——把n從5萬降到19614×14的patch計(jì)算量直接降了幾個(gè)數(shù)量級(jí)。所以我的個(gè)人判斷是SE/CBAM/CA這類輕量注意力模塊和自注意力解決的是不同層次的問題。前者適合作為卷積網(wǎng)絡(luò)的即插即用模塊用很小的代價(jià)提升baseline后者適合作為骨干網(wǎng)絡(luò)的核心構(gòu)件承擔(dān)全局特征交互。大多數(shù)實(shí)際項(xiàng)目不需要把backbone換成Transformer先用SE或CBAM把baseline提上去再評(píng)估是否值得上自注意力這個(gè)順序是性價(jià)比最高的。6. 時(shí)序注意力當(dāng)特征維變成時(shí)間維前面討論的都是圖像任務(wù)注意力在時(shí)間序列里的應(yīng)用常常被視覺從業(yè)者忽略。時(shí)序注意力機(jī)制的原理并不復(fù)雜但應(yīng)用方式和視覺里的注意力差別很大值得單獨(dú)拿出來講。6.1 時(shí)序注意力的兩種形態(tài)時(shí)序注意力的第一種形態(tài)是時(shí)間步注意力輸入是一個(gè)時(shí)間序列每個(gè)時(shí)間步有一個(gè)特征向量給每個(gè)時(shí)間步學(xué)習(xí)一個(gè)權(quán)重加權(quán)求和得到整個(gè)序列的表示。這種形態(tài)適合分類或者回歸任務(wù)——比如用一段腦電信號(hào)判斷是否異常不需要保留每一時(shí)刻的細(xì)節(jié)只關(guān)心哪些時(shí)刻最有判別性。第二種形態(tài)是編碼器-解碼器注意力解碼器在生成某個(gè)時(shí)刻的輸出時(shí)需要從編碼器的所有時(shí)間步里檢索相關(guān)信息。經(jīng)典的Bahdanau Attention就是這種形態(tài)它計(jì)算解碼器當(dāng)前時(shí)刻的Query和編碼器所有時(shí)間步的Key之間的相關(guān)性用相關(guān)性權(quán)重去加權(quán)編碼器的Value。原理上時(shí)序注意力和SE的計(jì)算流程幾乎一致區(qū)別只在權(quán)重作用在哪個(gè)維度。SE的權(quán)重作用在通道維度時(shí)序注意力的權(quán)重作用在時(shí)間維度。6.2 實(shí)現(xiàn)一個(gè)簡(jiǎn)單的時(shí)間步注意力假設(shè)輸入形狀是B×T×D也就是批量B、時(shí)間長(zhǎng)度T、每步特征維度D。我們想給T個(gè)時(shí)間步各學(xué)一個(gè)權(quán)重import torch import torch.nn as nn import torch.nn.functional as F class TemporalAttention(nn.Module): def __init__(self, d_model): super().__init__() self.score_net nn.Sequential( nn.Linear(d_model, d_model // 2), nn.Tanh(), nn.Linear(d_model // 2, 1) ) def forward(self, x): # x: b,t,d scores self.score_net(x).squeeze(-1) # b,t weights F.softmax(scores, dim-1) # b,t context torch.bmm(weights.unsqueeze(1), x).squeeze(1) # b,d return context, weights這里有個(gè)值得注意的細(xì)節(jié)打分函數(shù)用了Tanh而不是ReLU。因?yàn)闀r(shí)間步的權(quán)重理論上應(yīng)該允許正負(fù)貢獻(xiàn)ReLU把負(fù)數(shù)全截?cái)嗔藭?huì)讓某些時(shí)間步被迫無視而不是抑制效果往往不如Tanh。時(shí)序注意力的一個(gè)常見誤用是直接把所有時(shí)間步的特征做softmax加權(quán)求和丟掉了時(shí)間順序。時(shí)間序列之所以是時(shí)間序列就是因?yàn)轫樞虮旧戆畔ⅰ粋€(gè)上升趨勢(shì)和一個(gè)下降趨勢(shì)即使數(shù)值分布相同含義也完全不同。如果你的模型用注意力把所有時(shí)間步揉成一個(gè)向量就再也分不清先漲后跌和先跌后漲了。正確做法是在注意力輸入前先通過一層循環(huán)網(wǎng)絡(luò)或卷積網(wǎng)絡(luò)把時(shí)間順序編碼進(jìn)特征或者保留時(shí)序注意力的同時(shí)額外拼接一個(gè)位置編碼。6.3 時(shí)序注意力在項(xiàng)目里的實(shí)踐經(jīng)驗(yàn)我在做工業(yè)設(shè)備振動(dòng)信號(hào)分類的項(xiàng)目時(shí)用過這個(gè)模塊。原始信號(hào)按窗口切分成每秒2048個(gè)采樣點(diǎn)每個(gè)窗口提取時(shí)頻特征后得到T×D的序列。加時(shí)間步注意力前后分類準(zhǔn)確率從91.2%提升到93.8%提升主要來自模型學(xué)會(huì)了忽略啟動(dòng)階段的異常抖動(dòng)重點(diǎn)關(guān)注穩(wěn)態(tài)階段的特征。另一個(gè)經(jīng)驗(yàn)是時(shí)序注意力模塊放的位置很關(guān)鍵。放在特征提取之前注意力分?jǐn)?shù)還停留在原始信號(hào)層面噪聲影響大放在特征提取之后、分類頭之前此時(shí)的特征已經(jīng)更有語義區(qū)分度注意力更容易學(xué)到有意義的時(shí)間權(quán)重。我通常建議放在最后那個(gè)全局池化層之前作為一個(gè)軟選擇層替代簡(jiǎn)單的平均池化。7. 選型實(shí)戰(zhàn)不同任務(wù)下我推薦用哪種注意力把常見機(jī)制都過了一遍之后到了最實(shí)際的問題我的項(xiàng)目到底該用哪一個(gè)我根據(jù)自己跑過的項(xiàng)目和一些公開的論文結(jié)論整理了下面的選型建議。7.1 各注意力機(jī)制對(duì)比機(jī)制作用維度額外參數(shù)量級(jí)核心優(yōu)勢(shì)主要限制SE通道低極輕量即插即用無空間感知CBAM通道空間低兼顧通道和空間7×7卷積增加少量計(jì)算CA通道位置編碼極低保留坐標(biāo)信息移動(dòng)端友好空間感知精度不足自注意力全局時(shí)空高長(zhǎng)距離依賴建模O(n2)計(jì)算數(shù)據(jù)需求大多頭自注意力全局時(shí)空高多關(guān)系模式并行同上且超參更多時(shí)序注意力時(shí)間步低突出判別性時(shí)刻需配合時(shí)序編碼器使用7.2 按任務(wù)類型的推薦圖像分類大模型訓(xùn)練如果你的backbone已經(jīng)是ResNet50以上數(shù)據(jù)量也夠大可以先試SE。數(shù)據(jù)量在百萬級(jí)以下SE比自注意力更穩(wěn)因?yàn)樽宰⒁饬Ω菀走^擬合。目標(biāo)檢測(cè)優(yōu)先試CBAM。檢測(cè)任務(wù)對(duì)位置敏感通道注意力和空間注意力協(xié)同作用比單用SE穩(wěn)定漲點(diǎn)。特別提醒CBAM加在FPN的每一層輸出上比只加在backbone上效果更好代價(jià)是IO開銷增大需要評(píng)估推理速度。語義分割/姿態(tài)估計(jì)CA是性價(jià)比最高的選擇它能在幾乎不增加參數(shù)的情況下給模型提供位置線索。如果分割任務(wù)對(duì)小目標(biāo)要求極高考慮在CA基礎(chǔ)上疊加一層自注意力但要注意顯存占用。移動(dòng)端輕量模型無腦選CA。ME的7×7卷積在Deep-wise網(wǎng)絡(luò)上有下采樣傾向CA的兩個(gè)1×1卷積幾乎是零成本。MobileNetV3默認(rèn)架構(gòu)里就帶SE很多工程實(shí)踐表明換CA后相同F(xiàn)LOPs約束下精度更高。時(shí)間序列分類用6.2節(jié)的時(shí)間步注意力放在特征提取層之后、分類頭之前。如果序列極長(zhǎng)先用一維卷積降采樣否則時(shí)間步注意力在幾千步長(zhǎng)度上容易變成均勻分布——softmax的輸出會(huì)趨于扁平學(xué)了等于沒學(xué)。7.3 幾個(gè)容易踩的坑注意力模塊不是加得越多越好。我見過有人把SE、CBAM、CA全部串在同一個(gè)block里結(jié)果訓(xùn)練時(shí)梯度消散loss不降反升。原因是這些模塊都是乘法門控多個(gè)門控串聯(lián)會(huì)不斷縮放特征值前向傳播時(shí)數(shù)值越乘越小反向傳播時(shí)梯度越乘越細(xì)。一個(gè)block里最多放一個(gè)通道注意力和一個(gè)輕量空間注意力不要再疊第三種。Sigmoid輸出的注意力權(quán)重分布容易出現(xiàn)飽和。當(dāng)通道數(shù)極大且訓(xùn)練初期權(quán)重更新過猛時(shí)Sigmoid很容易輸出0.99甚至1.0乘法門控直接退化成恒等映射注意力模塊完全失效。解決辦法是給注意力權(quán)重視情況加一點(diǎn)L2正則或者把Sigmoid換成帶溫度參數(shù)的版本初始溫度大一點(diǎn)讓權(quán)重往0.5附近分布訓(xùn)練中再逐漸降溫。注意力可視化不能只看熱力圖。很多人把特征圖加權(quán)后的熱力圖直接當(dāng)模型的關(guān)注區(qū)域但熱力圖只能反映權(quán)重大的位置不能反映權(quán)重小的位置是否被正確忽略。比如一個(gè)模型預(yù)測(cè)出貓時(shí)熱力圖集中在貓頭上但這不代表它沒有錯(cuò)誤地關(guān)注了背景區(qū)域——可能在某個(gè)通道里背景的信息權(quán)重也很高只是被softmax或者其他通道的數(shù)值掩蓋了。要做嚴(yán)謹(jǐn)?shù)臍w因分析應(yīng)該用Grad-CAM或者積分梯度這類方法而不是簡(jiǎn)單的注意力熱力圖。7.4 我個(gè)人的選型經(jīng)驗(yàn)總結(jié)最后說點(diǎn)沒有寫在論文里的經(jīng)驗(yàn)。注意力機(jī)制本質(zhì)上是在已有特征不太夠用的時(shí)候幫你把信息重新分配它不能憑空創(chuàng)造新信息。如果你的模型本身特征提取能力很弱比如淺層網(wǎng)絡(luò)訓(xùn)練不充分加什么注意力都救不回來。先確保baseline是收斂的、可復(fù)現(xiàn)的再上注意力模塊否則你根本分不清漲點(diǎn)是因?yàn)樽⒁饬€是因?yàn)橛?xùn)練過程本身的變化。另外一個(gè)心態(tài)上的建議別迷信最新。CA出來之后視覺社區(qū)很快又出了幾十種變體各種Coordinate Attention、Efficient Attention層出不窮。但很多變體只在一兩個(gè)數(shù)據(jù)集上微漲泛化性存疑。工程上最穩(wěn)妥的做法是先建立一個(gè)標(biāo)準(zhǔn)化的評(píng)估流程把SE、CBAM、CA這三種在固定配置下跑一遍選擇一個(gè)穩(wěn)定漲點(diǎn)的然后去優(yōu)化數(shù)據(jù)、增強(qiáng)和損失函數(shù)——這些往往比換一個(gè)更新奇的注意力模塊收益更大。我在多個(gè)項(xiàng)目里反復(fù)驗(yàn)證過這個(gè)觀點(diǎn)主流的SE、CBAM、CA之間在同一baseline上的精度差異通常在2%以內(nèi)但數(shù)據(jù)質(zhì)量、增強(qiáng)策略和訓(xùn)練調(diào)度帶來的差異動(dòng)輒5%以上。注意力機(jī)制值得用但別把它當(dāng)成解決所有問題的銀彈。先把基本功做好再加注意力這才是最務(wù)實(shí)的路。