:PyTorch預(yù)訓(xùn)練模型微調(diào)與避坑指南)
簡介這是一份基于 PyTorch 的 Vision TransformerViT實現(xiàn)面向深度學(xué)習(xí)研究者與工程師提供從原始 JAX/Flax 權(quán)重轉(zhuǎn)換而來的預(yù)訓(xùn)練模型可直接用于圖像分類、特征提取以及下游任務(wù)微調(diào)。壓縮包約 173KB共 35 個文件其中以 22 個 Python 源碼為主涵蓋模型定義、訓(xùn)練、評估、數(shù)據(jù)加載與配置管理另有 README、Markdown 說明、requirements 環(huán)境文件、YAML 配置及 Notebook 示例等目錄結(jié)構(gòu)清晰便于二次開發(fā)。資源描述了與原始模型相當?shù)慕Y(jié)果支持 ImageNet2012 等數(shù)據(jù)集并附帶微調(diào)與評估腳本適合有一定 PyTorch 基礎(chǔ)、希望快速復(fù)現(xiàn) ViT 論文或?qū)⑵鋺?yīng)用于自身視覺任務(wù)的開發(fā)者。已有 7800 余人瀏覽學(xué)習(xí)是入門 Vision Transformer 并獲取可用預(yù)訓(xùn)練權(quán)重的實用參考。 我做視覺這兩年最常用的框架就是PyTorch而ViT相關(guān)的項目里vision-transformer-pytorch這個庫我?guī)缀跏欠磸?fù)在用。很多朋友第一次接觸Vision TransformerViT時被論文里一堆概念唬住總覺得這是個很復(fù)雜的模型。但如果你上手跑一遍這個庫會發(fā)現(xiàn)ViT的思路其實非常直接把圖像切成小塊當成一串“視覺單詞”送給Transformer去處理。而且這個項目坐標很明確——Pytorch 預(yù)訓(xùn)練模型正好是現(xiàn)在視覺任務(wù)落地最常用的一套組合。這個項目解決的核心問題就是讓大家不用重復(fù)造輪子。它提供了完整的ViT模型實現(xiàn)結(jié)構(gòu)清晰、參數(shù)可調(diào)并且可以配合預(yù)訓(xùn)練權(quán)重直接使用。對我個人來說它最大的價值在于既能幫新手理解ViT內(nèi)部到底發(fā)生了什么也能讓有經(jīng)驗的工程師在幾行代碼內(nèi)完成模型搭建和遷移學(xué)習(xí)。無論你是打算做圖像分類還是把ViT當作骨干網(wǎng)絡(luò)接進檢測分割框架這篇文章都適用。1. ViT模型為什么值得關(guān)注從CNN到Transformer的范式轉(zhuǎn)移1.1 圖像能否直接當序列處理在ViT出現(xiàn)之前視覺模型幾乎被卷積神經(jīng)網(wǎng)絡(luò)CNN統(tǒng)治。CNN的核心假設(shè)是局部性和平移不變性相鄰像素關(guān)系更密切同一個卷積核在整張圖上滑動。這個假設(shè)在ImageNet這類中大規(guī)模數(shù)據(jù)集上非常好用因為卷積天然帶有先驗不需要太多數(shù)據(jù)就能學(xué)會。但Transformer的想法完全不同。它最初在NLP里證明了一件事只要數(shù)據(jù)足夠多你不給模型任何結(jié)構(gòu)先驗讓注意力機制自己去找全局關(guān)系效果反而可能更好。于是就有了一個自然的問題圖像能不能也當成一個token序列來處理ViT回答這個問題的方式很直接——把圖像切成固定大小的patch每個patch線性映射成一個向量再按順序拼起來當作一個“句子”來處理。我最早看到這個思路的時候也挺驚訝整個ViT居然沒有一個卷積層全靠自注意力在建模。也正是這種極簡讓它在超大數(shù)據(jù)集上表現(xiàn)驚人。論文里用JFT-300M這種量級的數(shù)據(jù)集做完預(yù)訓(xùn)練之后ViT在ImageNet上的精度能超過同等規(guī)模的ResNet和EfficientNet。所以它不是換了個結(jié)構(gòu)而是換了一種視覺建模的范式。1.2 vision-transformer-pytorch解決的三個實際問題用這個庫一年多我感受到的突出價值有三點。第一實現(xiàn)與論文高度對齊可讀性強。它不是把ViT封裝成黑盒而是把patch embedding、transformer encoder、分類頭拆成模塊改哪里都一目了然。我在對比不同層數(shù)、不同head數(shù)量對精度影響的時候基本就是改幾個參數(shù)的事。第二原生PyTorch生態(tài)無縫銜接。庫本身不依賴timm或者更上層的框架直接就能和你自己的訓(xùn)練管線集成。我經(jīng)常需要把ViT輸出的特征接給檢測頭或者分割頭這個庫的數(shù)據(jù)流非常透明改起來很順手。第三社區(qū)認可度高坑少。這個項目在GitHub上star量很大用的人多意味著踩坑經(jīng)驗多。比如位置編碼的維度問題、patch size選擇問題網(wǎng)上一搜就有很多討論遇到bug不至于孤立無援。對做工程和做研究的人來說這種成熟度很重要。2. 核心架構(gòu)拆解ViT到底在做什么2.1 Patch Embedding圖像是如何變成Token序列的ViT最核心的一個操作就是Patch Embedding。以最常見的ViT-B/16為例輸入是224x224的RGB圖像patch_size設(shè)為16那么圖像會被劃分成(224/16)(224/16)1414196個patch。每個patch大小為16x16x3把它展平成長度為768的向量再經(jīng)過一個線性投影層映射到768維的embedding空間。很多人第一次看會疑惑為什么不直接展平再用全連接其實這里的線性投影本質(zhì)就是一個1x1卷積或者一個reshape加Linear它的作用是讓每個patch的原始像素映射到更適合Transformer處理的語義空間。實際操作中很多實現(xiàn)直接用nn.Conv2d(in_channels3, out_channelsdim, kernel_sizepatch_size, stridepatch_size)來一步完成切patch和投影效率更高。這也是為什么你會看到有些代碼里Patch Embedding層長得像卷積層但它做的事情其實就是“切塊線性變換”。這里我想強調(diào)一個點patch size的選擇直接影響序列長度。patch越小序列越長計算量越大但細節(jié)保留越多。ViT-B/32用32的patch序列長度只有491個token速度快很多但精度略降。實際工程里如果顯存有限又不想掉太多精度可以考慮用大patch或者保持patch不變減少層數(shù)。2.2 位置編碼、CLS Token與Transformer Encoderpatch被映射成token之后接下來的問題很關(guān)鍵Transformer本身是順序無關(guān)的它不知道哪個token在圖像的哪個位置。ViT的做法是加一個1D可學(xué)習(xí)的位置編碼向量直接加到所有token的embedding上。這里沒有用NLP里常見的2D位置編碼因為論文實驗發(fā)現(xiàn)1D可學(xué)習(xí)編碼對效果影響不大但實現(xiàn)更簡單。ViT還在序列最前面插入了一個特殊的CLS token它的作用和BERT里的CLS一樣用于匯聚全局信息。在Transformer編碼若干層之后模型拿出CLS token對應(yīng)的輸出向量接一個分類頭完成最終分類。我實際中還見過一些改造方案比如直接對所有token做全局平均池化再分類效果有時也不差但標準ViT用的是CLS token方案。接下來是Transformer Encoder。以ViT-Base為例包含12層Encoder每層由多頭自注意力12個head、MLPhidden size從768擴展到3072再降回來、LayerNorm和殘差連接組成。值得注意的是ViT用的是Pre-LayerNorm結(jié)構(gòu)也就是每個子層注意力或MLP之前先做歸一化。這個細節(jié)影響穩(wěn)定性訓(xùn)練大模型時尤其明顯。我自己的經(jīng)驗是這種設(shè)計配合較大的學(xué)習(xí)率也能保持穩(wěn)定微調(diào)時不容易崩。2.3 模型規(guī)格怎么選Base、Large與HugeViT官方發(fā)布了幾個規(guī)格Base86M參數(shù)、Large307M參數(shù)、Huge632M參數(shù)。視覺任務(wù)里用最多的就是Base它和ResNet50規(guī)模差不多但效果更好。如果資源充足、任務(wù)復(fù)雜Large往往能帶來明顯提升Huge則適合在超大數(shù)據(jù)集上從頭訓(xùn)練一般做遷移學(xué)習(xí)的用不起。我選型時通常會先想清楚數(shù)據(jù)集規(guī)模。數(shù)據(jù)集只有幾千張圖直接用Base甚至Small版本配合在ImageNet上預(yù)訓(xùn)練的權(quán)重效果往往比從零訓(xùn)練要穩(wěn)得多。所以這里就引出了下一部分的重點預(yù)訓(xùn)練模型到底怎么選、怎么用。3. 預(yù)訓(xùn)練模型的選擇與微調(diào)實戰(zhàn)思路3.1 三種獲取預(yù)訓(xùn)練權(quán)重的方式對比標題里特別提到了“帶有預(yù)訓(xùn)練模型”這一點其實是很多人最關(guān)心的。vision-transformer-pytorch庫本身側(cè)重于提供模型結(jié)構(gòu)而預(yù)訓(xùn)練權(quán)重通常可以從下面三個渠道獲取獲取渠道是否攜帶官方預(yù)訓(xùn)練權(quán)重適用場景vit-pytorch庫否只提供模型結(jié)構(gòu)學(xué)習(xí)結(jié)構(gòu)、自定義改造timm是ImageNet預(yù)訓(xùn)練權(quán)重日常分類、工程落地HuggingFace transformers是Google官方權(quán)重研究復(fù)現(xiàn)、需要官方預(yù)處理我個人的建議是追求簡單就直接用timm一行代碼搞定下載和加載追求跟原論文對齊就去HuggingFace。用timm加載預(yù)訓(xùn)練權(quán)重的方法非常直接import timm model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes1000) model.eval()如果你要做的分類任務(wù)類別數(shù)不是1000可以直接在create_model時指定num_classestimm會自動把最后的分類頭替換成對應(yīng)數(shù)量微調(diào)時非常省事。另外timm還提供了豐富的數(shù)據(jù)增強策略、EMA等訓(xùn)練工具對訓(xùn)練精度的提升很友好。如果你更希望復(fù)現(xiàn)原論文的預(yù)處理流程可以用HuggingFace的transformersfrom transformers import ViTForImageClassification, ViTImageProcessor model ViTForImageClassification.from_pretrained(google/vit-base-patch16-224) processor ViTImageProcessor.from_pretrained(google/vit-base-patch16-224)這里的processor封裝了標準化和尺寸調(diào)整邏輯拿過來就能用不用自己糾結(jié)normalize參數(shù)。不過要注意transformers加載這種大權(quán)重時會從HuggingFace服務(wù)器下載雖然大部分時候沒問題但偶爾會卡住這時候可以設(shè)置環(huán)境變量HF_ENDPOINThttps://hf-mirror.com用國內(nèi)鏡像加速下載。注意預(yù)訓(xùn)練權(quán)重的下載只是第一步真正決定模型效果的是后續(xù)微調(diào)策略。3.2 微調(diào)策略凍結(jié)層、學(xué)習(xí)率與數(shù)據(jù)規(guī)模拿到預(yù)訓(xùn)練模型后第一個選擇就是凍結(jié)還是不凍結(jié)。如果數(shù)據(jù)集比較小、只有幾千張我建議凍結(jié)前11層Encoder只微調(diào)最后一層和分類頭。實操中Encoder底層學(xué)習(xí)到的都是一些基礎(chǔ)紋理、邊緣特征這些對于任何視覺任務(wù)都是通用的不需要重新學(xué)。凍結(jié)之后可以大幅減少顯存占用和訓(xùn)練時間加快收斂。如果數(shù)據(jù)集中等幾萬張或者目標域和ImageNet差異很大比如醫(yī)學(xué)影像、衛(wèi)星圖我建議全部微調(diào)但把學(xué)習(xí)率調(diào)低一些。常規(guī)的做法是整體學(xué)習(xí)率設(shè)0.0001左右分類頭學(xué)習(xí)率可以稍微放大到0.001。ViT對學(xué)習(xí)率比較敏感一開始就用太大學(xué)習(xí)率很容易出現(xiàn)loss震蕩甚至不收斂。另外一個小技巧是如果顯存允許可以先把模型在較大分辨率如384x384上微調(diào)幾輪效果通常會比224更好。因為ViT沒有卷積的局部先驗更大分辨率意味著更多patch更多細節(jié)。代價就是序列長度變長顯存和速度都翻幾倍。我自己做細粒度分類時這個技巧帶來的精度提升很明顯。4. 實操記錄從安裝到自定義數(shù)據(jù)集微調(diào)4.1 環(huán)境準備和安裝先說我實測過的環(huán)境組合Python 3.10PyTorch 2.1以上CUDA 11.8顯卡是RTX 3090。PyTorch的安裝直接用官方命令就行如果下載慢可以把pip源換成清華源或者阿里源。裝完之后裝依賴pip install torch torchvision timm pip install vit-pytorch后面這個vit_pytorch就是lucidrains的版本也是很多博客里提到的vision-transformer-pytorch的PyTorch實現(xiàn)。不過我想特別提醒一句vit_pytorch這個包默認不攜帶預(yù)訓(xùn)練權(quán)重它的作用是快速構(gòu)建模型結(jié)構(gòu)你想直接跑預(yù)訓(xùn)練推理還是要配合timm或者transformers。很多人一開始沒搞清楚這點裝了個vit-pytorch然后發(fā)現(xiàn)模型是隨機初始化的以為庫有問題。4.2 快速推理用預(yù)訓(xùn)練權(quán)重分類一張圖下面這段代碼是我在項目里驗證一張新圖時常用的模板基于timm實現(xiàn)簡單可靠import torch import timm from PIL import Image from torchvision import transforms device torch.device(cuda if torch.cuda.is_available() else cpu) model timm.create_model(vit_base_patch16_224, pretrainedTrue) model model.to(device) model.eval() transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) img Image.open(test.jpg).convert(RGB) x transform(img).unsqueeze(0).to(device) with torch.no_grad(): out model(x) prob torch.softmax(out, dim1) top5 torch.topk(prob, 5)輸出top5之后去查一下ImageNet的類別索引表就能知道模型預(yù)測的是什么類。這里最常踩的坑有兩個一是忘了convert(RGB)導(dǎo)致灰度圖或RGBA圖報錯二是忘了加batch維度unsqueeze(0)少寫就報維度錯誤。我剛開始跑的時候就在這兩處浪費過時間。4.3 自定義數(shù)據(jù)集微調(diào)假設(shè)現(xiàn)在你有一個10類的自定義數(shù)據(jù)集目錄結(jié)構(gòu)大概是train/class1、train/class2這樣。用torchvision的ImageFolder讀進來然后替換最后的分類頭就可以開始微調(diào)import torch import torch.nn as nn import timm from torchvision import datasets, transforms from torch.utils.data import DataLoader model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes10) model.to(device) train_transform transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) train_ds datasets.ImageFolder(train, transformtrain_transform) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4) optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay0.05) criterion nn.CrossEntropyLoss() for epoch in range(10): model.train() for x, y in train_loader: x, y x.to(device), y.to(device) loss criterion(model(x), y) optimizer.zero_grad() loss.backward() optimizer.step() print(fepoch {epoch}, loss {loss.item():.4f})這段代碼只是最基礎(chǔ)的訓(xùn)練循環(huán)。實際項目里我還會加warmup、余弦退火、Mixup和數(shù)據(jù)增強。ViT在中小數(shù)據(jù)集上容易過擬合所以增強策略比CNN時代更講究。還有一個細節(jié)AdamW的weight_decay我習(xí)慣設(shè)0.05這是從DeiT論文里來的經(jīng)驗值實測比0.01穩(wěn)定。5. 避坑指南使用ViT時最容易踩的坑5.1 輸入尺寸和歸一化必須匹配ViT模型對輸入尺寸非常敏感這一點比CNN嚴格得多。timm里的vit_base_patch16_224要求輸入224x224如果你給它喂512x512的圖patch數(shù)量就變了位置編碼的維度對不上通常運行到模型內(nèi)部就會直接報維度不匹配的錯誤。我的建議是把預(yù)處理統(tǒng)一封裝成一個函數(shù)Resize到224歸一化參數(shù)就用ImageNet默認的mean和std不要自己隨便改。另外一旦換了數(shù)據(jù)集別忘了重新檢查normalize參數(shù)是否匹配——尤其是醫(yī)學(xué)圖像或者紅外圖像它們的像素分布和自然圖像差別很大。5.2 預(yù)訓(xùn)練權(quán)重下載失敗怎么辦這個問題在HuggingFace的權(quán)重上尤其常見。下載到一半斷掉、網(wǎng)絡(luò)超時、緩存損壞都會導(dǎo)致無法加載。我遇到這種情況一般分兩步排查先看報錯是不是SSL或者超時如果是就說明是網(wǎng)絡(luò)問題可以設(shè)置HF_ENDPOINThttps://hf-mirror.com再重新拉取如果報錯是鍵名不匹配那大概率是模型結(jié)構(gòu)定義和權(quán)重來源版本不一致比如用了patch16的定義去加載patch32的權(quán)重。權(quán)重和模型結(jié)構(gòu)的匹配相當重要。提示HuggingFace的緩存目錄通常在~/.cache/huggingface刪掉對應(yīng)模型的緩存再重新下載可以解決很多奇怪的加載問題。下載慢但沒報錯時多試幾次或者手動下載后放到緩存目錄也行。5.3 顯存不夠用先別急著換顯卡ViT雖然參數(shù)不算特別多但自注意力的計算復(fù)雜度是序列長度的平方。224分辨率下197個token還好一旦輸入到448x448token數(shù)變成(448/16)^21785計算量增長非常明顯。如果顯存爆了最直接的思路是減小batch size或者把patch_size從16改成32。還有一個實用技巧是開啟梯度累積用多個小batch累加梯度模擬大batch效果能接近但省顯存。訓(xùn)練速度慢的話優(yōu)先檢查是不是數(shù)據(jù)加載瓶頸。num_workers調(diào)大或者用pin_memoryTrue經(jīng)常能把GPU利用率拉滿。我見過很多新人把num_workers默認0跑GPU利用率低得可憐改到4或8之后速度立竿見影。相比一上來就換卡這招劃算得多。5.4 位置編碼與輸入分辨率不匹配如果你想在384x384或更大的分辨率上微調(diào)直接用vit_base_patch16_224的權(quán)重會報位置編碼維度不匹配。因為預(yù)訓(xùn)練權(quán)重的position embedding是1x197x768而384x384對應(yīng)的是1x577x768。解決辦法是插值調(diào)整位置編碼的尺寸。timm里有些版本支持img_size參數(shù)或者在create_model時指定img_size384但不是所有實現(xiàn)都自動處理。如果要手動插值可以這樣import torch from vit_pytorch import ViT model ViT(...) pos_embed model.pos_embedding # 形狀 [1, 197, 768] new_pos_embed torch.nn.functional.interpolate( pos_embed.permute(0, 2, 1).unsqueeze(0), size(577,), modelinear, align_cornersFalse, ).squeeze(0).permute(0, 2, 1) model.pos_embedding torch.nn.Parameter(new_pos_embed)不過說實話如果只是做普通分類任務(wù)我不建議手動插值直接用timm里帶384后綴的模型比如vit_base_patch16_384會省心很多位置編碼部分timm已經(jīng)處理好了。把上面這些坑都踩過一遍之后我對ViT的理解反而更深了。說實話ViT的門檻并不在模型本身而在于各種細節(jié)patch怎么切、位置編碼怎么處理、預(yù)訓(xùn)練權(quán)重怎么融合、微調(diào)參數(shù)怎么調(diào)。把這些細節(jié)弄明白它就是非常趁手的視覺骨干網(wǎng)絡(luò)。如果你剛開始接觸ViT我的建議是先按著代碼把模型結(jié)構(gòu)打印出來一個個模塊核對再跑通預(yù)訓(xùn)練推理最后再上自己的數(shù)據(jù)。這個過程走一遍比干啃論文有用得多。最后分享一個我自己的習(xí)慣任何新項目要采用ViT我都會先用Base模型加ImageNet預(yù)訓(xùn)練跑一個baseline確認任務(wù)可行之后再根據(jù)顯存和精度需求決定要不要換Large或者調(diào)分辨率。不要一開始就上大模型否則調(diào)參和排錯的成本會高到讓你懷疑人生。希望這篇內(nèi)容對你有幫助有問題歡迎留言交流。本文還有配套的精品資源點擊獲取