普通conv和fc層的凍結方式:
# 凍結引數
for i, p in enumerate(self.model.parameters()):
if i <= 66:
p.requires_grad = False
# 驗證一下是否成功凍結引數
for k, v in self.model.named_parameters():
print("k:{} v:{} ".format(k, v.requires_grad))
注意:model.parameters()都在梯度回傳的更新程序中,所以可以用param.requires_grad = False的方式凍結,但是對于一些BN層的引數,比如BN層的runing_mean和runing_var,這兩個值是前向計算統計得來的,并沒有在梯度回傳的更新程序中,所以,param.requires_grad=False對它們不起任何作用!
踩坑:
我的目的:在共用一個主干網路的多任務學習中,完全凍結其中一個表現較好的任務1分支,只訓練其他兩個任務:任務2分支和任務3分支,
結果:我以為用 “param.requires_grad=False” 的方式可以凍結任務1分支的所有引數,然后我發現我錯了,凍結完,在驗證程序中,我發現任務1的表現居然變差了,
驗證:列印引數值,發現任務1的卷積層和全連接層引數不變(被成功凍結),只有BN層的runing_mean和runing_var發生了改變(未被凍結),應該就是他們的問題,
凍結BN層的runing_mean和runing_var的方法可以參照:
def fix_bn(m):
classname = m.__class__.__name__
if classname.find('BatchNorm') != -1:
m.eval()
model = models.resnet50(pretrained=True)
model.cuda()
model.train()
model.apply(fix_bn) # fix batchnorm
總結:當需要保持某一分支的分類性能不變時(我的意思是后續都不對這個分支進行訓練了,只拿來驗證和測驗),除了要凍結可回傳梯度的權重值,還要凍結上述BN層的值,當然如果網路還要繼續訓練的話,也可以不凍結BN層的runing_mean和runing_var,如果凍結網路后某一分支的性能突然變差,可以考慮一下試試凍結BN層的runing_mean和runing_var~
轉載請註明出處,本文鏈接:https://www.uj5u.com/qita/437013.html
標籤:AI
上一篇:Windows10環境下自己配置Pytracking詳細流程(有參考博客)
下一篇:ORB-SLAM2的安裝與運行
