技巧:MixMatch-pytorch 中正確計(jì)算 BatchNorm 的進(jìn)階指南)
interleave 交錯(cuò)技巧MixMatch-pytorch 中正確計(jì)算 BatchNorm 的進(jìn)階指南【免費(fèi)下載鏈接】MixMatch-pytorchCode for MixMatch - A Holistic Approach to Semi-Supervised Learning項(xiàng)目地址: https://gitcode.com/gh_mirrors/mi/MixMatch-pytorchMixMatch-pytorch 是對經(jīng)典半監(jiān)督學(xué)習(xí)論文《MixMatch: A Holistic Approach to Semi-Supervised Learning》的非官方 PyTorch 復(fù)現(xiàn)。在它的訓(xùn)練代碼里藏著一個(gè)讓無數(shù)新手困惑的細(xì)節(jié)——interleave 交錯(cuò)技巧它直接決定了混合 batch 下的BatchNorm統(tǒng)計(jì)量是否算得正確進(jìn)而影響最終精度。本文用最通俗的方式帶你徹底看懂這個(gè)進(jìn)階技巧的原理與源碼實(shí)現(xiàn)。什么是 MixMatch半監(jiān)督學(xué)習(xí)為什么需要混合 batch半監(jiān)督學(xué)習(xí)的場景很現(xiàn)實(shí)標(biāo)注數(shù)據(jù)貴無標(biāo)注數(shù)據(jù)多。MixMatch 的思路是把兩者揉在一起訓(xùn)練流程分三步兩次增強(qiáng)每個(gè)無標(biāo)注樣本做兩次隨機(jī)增強(qiáng)得到 U1、U2標(biāo)簽猜測 銳化用模型預(yù)測兩個(gè)增強(qiáng)版本的類別概率并取平均再用溫度 T 銳化出偽標(biāo)簽mixup 混合把有標(biāo)注樣本 X 與兩份無標(biāo)注樣本 U1、U2 線性混合得到混合后的訓(xùn)練樣本。于是每個(gè)訓(xùn)練 batch 里實(shí)際是X U1 U2 三個(gè) batch_size 的數(shù)據(jù)本項(xiàng)目默認(rèn) batch-size64合計(jì) 192 張圖。問題就從這里開始出現(xiàn)了。BatchNorm 的坑為什么直接一起算會出錯(cuò)BatchNorm 在訓(xùn)練模式下會統(tǒng)計(jì)當(dāng)前 mini-batch 內(nèi)樣本的均值和方差來歸一化特征。同樣的輸入放進(jìn)不同的 batch輸出就會不同。而 MixMatch 的混合 batch 恰好踩中了兩個(gè)坑整批一起前向如果把 3B 個(gè)樣本一次性塞進(jìn)網(wǎng)絡(luò)BN 統(tǒng)計(jì)的是 3B 個(gè)樣本的統(tǒng)計(jì)量和正常訓(xùn)練時(shí)每批 B 個(gè)樣本的統(tǒng)計(jì)口徑完全不一致拆開分別前向如果按原順序拆成 3 份三份的構(gòu)成各不相同第一份偏有標(biāo)注分布后兩份偏無標(biāo)注分布三次前向的 BN 統(tǒng)計(jì)量彼此漂移梯度會變得很不穩(wěn)定。結(jié)論很明確必須讓每一次前向傳播都看到構(gòu)成一致、來源均衡的 batch才能得到正確的 BatchNorm 計(jì)算。這正是 interleave 交錯(cuò)技巧的用武之地。interleave 交錯(cuò)技巧的核心原理interleave 的思路可以概括為一句話把三份數(shù)據(jù)各自切碎再交叉重組成三份拼盤。假設(shè)三份原始 chunk 各有 3 個(gè)子塊原始 Chunk子塊 1子塊 2子塊 3第 1 份偏有標(biāo)注A0A1A2第 2 份U1B0B1B2第 3 份U2C0C1C2interleave 會做兩次換位A1?B1、A2?C2重組后變成重組 Chunk子塊 1子塊 2子塊 3第 1 份A0B1C2第 2 份B0A1C1第 3 份C0B2A2現(xiàn)在每一份 batch 都同時(shí)包含來自 A、B、C 三個(gè)來源的子塊組成完全均衡。三次前向傳播各自獨(dú)立計(jì)算 BatchNorm 統(tǒng)計(jì)量口徑一致這就是代碼注釋里correct batchnorm calculation的含義。更巧妙的是interleave 是一種對合操作——前向傳播之后再調(diào)用一次 interleave就能把順序完全還原然后按位置切出有標(biāo)注的 logits 和無標(biāo)注的 logits 分別計(jì)算損失完美閉環(huán)。源碼解析train.py 中的 interleave 實(shí)現(xiàn)整個(gè)技巧在train.py里只有兩個(gè)函數(shù)卻非常精煉。先看調(diào)用處的關(guān)鍵邏輯# mixup 之后把 3B 的混合 batch 拆成 3 份再交錯(cuò)重組 mixed_input list(torch.split(mixed_input, batch_size)) mixed_input interleave(mixed_input, batch_size) # 3 次前向傳播每次都是一個(gè)構(gòu)成均衡的 batch logits [model(mixed_input[0])] for input in mixed_input[1:]: logits.append(model(input)) # 前向完成后再次 interleave把順序還原回去 logits interleave(logits, batch_size) logits_x logits[0] logits_u torch.cat(logits[1:], dim0)再看實(shí)現(xiàn)本體。interleave_offsets負(fù)責(zé)把每個(gè) batch 盡量均勻地切成 nu1 個(gè)子塊def interleave_offsets(batch, nu): groups [batch // (nu 1)] * (nu 1) for x in range(batch - sum(groups)): groups[-x - 1] 1 offsets [0] for g in groups: offsets.append(offsets[-1] g) return offsetsinterleave則完成切碎 → 換位 → 重組三步def interleave(xy, batch): nu len(xy) - 1 offsets interleave_offsets(batch, nu) xy [[v[offsets[p]:offsets[p 1]] for p in range(nu 1)] for v in xy] for i in range(1, nu 1): xy[0][i], xy[i][i] xy[i][i], xy[0][i] return [torch.cat(v, dim0) for v in xy]如果 batch_size 不能被 nu1 整除比如 64 不能被 3 整除interleave_offsets會把余數(shù)均勻分配到靠后的子塊保證任何 batch_size 都能用。復(fù)現(xiàn)指南一行命令跑通 CIFAR-10 半監(jiān)督訓(xùn)練想親手體驗(yàn) interleave 交錯(cuò)技巧的效果先克隆倉庫git clone https://gitcode.com/gh_mirrors/mi/MixMatch-pytorch安裝依賴PyTorch、torchvision、tensorboardX、progress、matplotlib、numpy后只需一行命令即可用 250 張有標(biāo)注圖片開始訓(xùn)練python train.py --gpu 0 --n-labeled 250 --out cifar10250模型是 WideResNet-28-2定義在models/wideresnet.py數(shù)據(jù)管線在dataset/cifar10.py無標(biāo)注樣本通過TransformTwice做兩次增強(qiáng)訓(xùn)練主循環(huán)在train.py。README 中給出的復(fù)現(xiàn)精度如下有標(biāo)注樣本數(shù)250500100020004000論文精度88.9290.3592.2592.9793.76本項(xiàng)目精度88.7188.9690.5292.2393.52在只有 250 張標(biāo)注圖片的情況下就能達(dá)到約 88.7% 的準(zhǔn)確率足見 MixMatch 與 interleave 配合的價(jià)值。常見問題與避坑小貼士Q去掉 interleave 直接前向會怎樣ABN 統(tǒng)計(jì)量口徑混亂訓(xùn)練不穩(wěn)定精度會明顯下降——這正是很多復(fù)現(xiàn)跑不出來的原因之一。Q為什么第二次調(diào)用 interleave 能還原順序Ainterleave 是對合操作兩次應(yīng)用等價(jià)于恒等映射前向后再交錯(cuò)一次即可恢復(fù)原始的樣本分組。Q除 interleave 外還有哪些關(guān)鍵細(xì)節(jié)A評估時(shí)使用指數(shù)移動平均模型WeightEMAdecay 0.999銳化溫度 T0.5無標(biāo)注損失權(quán)重 λ_u 隨訓(xùn)練線性 ramp-up這些在train.py中都有體現(xiàn)。 小貼士如果你想驗(yàn)證 interleave 的作用可以臨時(shí)把訓(xùn)練循環(huán)里兩處interleave調(diào)用注釋掉對比精度你會直觀感受到這個(gè)不起眼技巧的分量。小結(jié)interleave 交錯(cuò)技巧是 MixMatch-pytorch 中最值得學(xué)習(xí)的工程細(xì)節(jié)之一它用極簡的代碼解決了混合 batch 下 BatchNorm 統(tǒng)計(jì)量不一致這個(gè)隱蔽問題保證了每一次前向傳播都在構(gòu)成均衡的 batch 上計(jì)算統(tǒng)計(jì)量。理解了它你就掌握了半監(jiān)督學(xué)習(xí)復(fù)現(xiàn)中最關(guān)鍵的一環(huán)也為后續(xù)閱讀其他一致性正則方法如 FixMatch、FlexMatch打下了堅(jiān)實(shí)基礎(chǔ)。【免費(fèi)下載鏈接】MixMatch-pytorchCode for MixMatch - A Holistic Approach to Semi-Supervised Learning項(xiàng)目地址: https://gitcode.com/gh_mirrors/mi/MixMatch-pytorch創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考