CV Day 6 用 timm 與預訓練權重做遷移學習
執行需求:Colab T4 可跑。本篇在 Colab 免費 T4 上完整跑一次遷移學習;訓練一個 epoch 的 ResNet-50 大約 8 分鐘。如果你只用 CPU,請先跳過訓練,直接觀察模型載入與評估段落。
引言
昨天的內容中,我們認識了 ResNet、EfficientNet、ConvNeXt、視覺 Transformer 四個影像分類骨幹,知道它們各自解決了什麼問題。但要在自家資料上拿到好表現,我們不太可能從頭訓練一個 86M 參數的 ViT-Base,這時候就要靠遷移學習:把在 ImageNet 上預訓練好的權重當起點,再用較小的標註資料微調(fine-tune),讓模型學到符合任務的特徵。實務上,超過九成的影像分類專案都是用遷移學習起步,這是深度學習工程師的標準武器。
這一篇要帶你用 timm 1.0.x 把昨天的骨幹接到新任務上。我們會準備一個二分類的影像任務(CIFAR-10 的貓與狗),從頭示範資料準備、預訓練權重載入、分類頭替換、差動學習率設定、訓練迴圈、評估與模型儲存。讀完這篇,你會了解遷移學習的標準流程,以及 timm 在每一步提供的工具與 API,並能在自己的資料集上重現整個流程。
為什麼用 timm 做遷移學習
timm 提供超過一千種模型與數千份預訓練權重,這是它跟手刻 torchvision 模型最大的差別。同一個 ResNet-50,timm 給出至少五種權重:resnet50.tv_in1k(torchvision 原始)、resnet50.a1_in1k(timm A1 配方)、resnet50.a1h_in1k(含額外增強)、resnet50.rsb_a1_in1k(RSB 改良訓練)等等。挑選合適的權重,往往比換骨幹更能提升遷移學習的表現。同樣的道理也適用於 ConvNeXt-Tiny:fb_in1k 與 fb_in22k_ft_in1k 兩個版本在遷移到醫療或衛星圖時,常常會差到 3–5 個百分點。如果資料集與 ImageNet 比較接近(例如一般消費性影像),A1 配方是最常見的起點;如果資料偏自然場景或稀有物件,可以考慮 in21k 預訓練的版本。
另外 timm 把「資料預處理」也包進模型設定裡。timm.data.resolve_model_data_config(model) 會回傳模型預訓練時使用的輸入尺寸、平均值、標準差,再交給 timm.data.create_transform(**cfg, is_training=True) 就能建立一致的訓練 transform。這大幅降低了「換模型就要重寫 transform」的維護成本,也是實務上常見的踩雷來源。實戰中我們遇過不少案例:團隊把 ResNet 換成 EfficientNet 時忘了同步平均值,結果第一次訓練直接掉了 4% top-1;用 resolve_model_data_config 之後只要傳 **cfg 就能避免這個錯誤。這個寫法的好處是「換模型只要換字串」,所有 transform 細節都跟著模型走,省下維護多份 transform 設定的麻煩,也是後續部署時最容易搬遷的寫法。下一篇在套上訓練技巧時,也會繼續沿用這個 cfg 設定。
要注意的是,預訓練權重有兩條主要路線:一是只用 ImageNet-1k 訓練(如 resnet50.a1_in1k),二是先在 ImageNet-21k 預訓練、再 fine-tune 回 in1k(如 resnet50.a1_in1k 的對應 in21k 版本或 ViT 的 augreg_in21k_ft_in1k)。一般建議先試 in1k 版本;如果自家資料與 ImageNet 差異大(例如醫療、衛星圖),再考慮 in21k 版本。另一個常被忽略的細節是 num_classes 預設值:create_model(tag, pretrained=True) 預設會建立 1000 類分類頭,這一定要在新任務建立時改寫。實務上在挑選權重前,先想清楚三件事:任務與 ImageNet 的視覺差異有多大、可用的標註資料量、推論時的裝置與延遲上限。這三個問題會決定你該選小模型 + in1k 預訓練,或大模型 + in21k 預訓練。
遷移學習的三種策略
遷移學習不是「一定要把所有參數都打開重訓」。實務上有三種策略,取捨依資料量與算力而定:
- 特徵萃取(Feature Extraction):把 backbone 凍結,只訓練新加的分類頭。實作上只要
for p in model.parameters(): p.requires_grad = False,再單獨把分類頭解開。優點是訓練快、不需要 GPU;缺點是 backbone 不會適應新資料,當訓練資料與 ImageNet 差很多時表現會卡住。適合資料量極小(例如每類 < 100 張)、或只想快速驗證流程的情境。在 Fashion-MNIST、CIFAR 子集上都能在 10 分鐘內完成。 - 全模型微調(Full Fine-tune):把整個模型(包含 backbone)解開一起訓練,搭配較小的學習率(通常 1e-4 到 5e-5)。優點是表現最好;缺點是訓練時間長、需要 GPU,對小資料集還可能過擬合。適合資料量 ≥ 數千張、有 T4 以上算力的情境。全模型微調通常需要配合較強的權重衰減(weight_decay ≥ 0.05)與資料增強,否則容易震盪。
- 差動學習率(Discriminative Learning Rate):backbone 用較小學習率(如 1e-4),新分類頭用較大學習率(如 1e-3)。這是前兩者的折衷,能讓 backbone 微調到新資料,又不會一下子把預訓練權重洗掉。實務上最常用,也最推薦的預設策略。如果資料量介於 1k 到 10k 之間,這通常是最穩的選擇。
這三種策略的差別來自一個事實:預訓練權重已經學到「影像的通用特徵」(邊緣、紋理、形狀),而我們新加的分類頭是隨機初始化的,它需要更大的學習率才能追上 backbone 的節奏。差動學習率正是把這個不對稱性寫進訓練設定。實務上還有一個進階做法叫「漸進解凍」:先用凍結 backbone 訓練 1–2 個 epoch 讓分類頭收斂,再依序解凍 backbone 的後段、中段、前段,每個階段用更小的學習率。這在資料量介於 100–1000 張時特別有效,也是 ULMFiT 與 fastai 推薦的標準流程。timml 也提供 set_grad_checkpointing 與 group_parameters 兩個工具,能把這樣的漸進策略寫成幾行程式碼。
另外要注意的是「類別不平衡」這個常被忽略的細節:如果你的新任務中某些類別只有幾十張樣本,其他類別有上萬張,預訓練的分類頭權重(1000 類)就完全派不上用場,一定要重訓。這時候可以搭配 class weight 或 oversampling 來平衡損失;torch.nn.CrossEntropyLoss(weight=torch.tensor([w0, w1, ...])) 就能設定每類的權重。實務上,差動學習率 + class weight 是處理不平衡資料的標準組合;若再加上 focal loss,效果會更穩。在工業界做瑕疵分類時,類別不平衡幾乎是必遇到的問題,事先想好策略能省下後續大量 debug 時間。
完整實作
以下範例在 Colab T4 上完整跑一次遷移學習:把 CIFAR-10 中的貓(label 3)與狗(label 5)取出當二分類任務,用 resnet50.a1_in1k 做差動學習率微調,1 個 epoch 約 8 分鐘。整段可以整段貼上 Colab 執行。
需先在 Colab 終端機執行:pip install timm==1.0.12 torch torchvision。
# 1. 準備資料:CIFAR-10 中只取貓(3) 與 狗(5)
import torch
from torch.utils.data import DataLoader, Dataset
import torchvision
import torchvision.transforms as T
LABEL_MAP = {3: 0, 5: 1} # cat -> 0, dog -> 1
class BinaryCIFAR(Dataset):
"""把 CIFAR-10 縮減為貓狗二分類,並把標籤重新映射到 0/1。"""
def __init__(self, base, indices):
self.base = base
self.indices = indices
def __len__(self):
return len(self.indices)
def __getitem__(self, i):
x, y = self.base[self.indices[i]]
return x, LABEL_MAP[y]
train_base = torchvision.datasets.CIFAR10(
root="./data", train=True, download=True, transform=None
)
val_base = torchvision.datasets.CIFAR10(
root="./data", train=False, download=True, transform=None
)
train_idx = [i for i, (_, y) in enumerate(train_base) if y in (3, 5)]
val_idx = [i for i, (_, y) in enumerate(val_base) if y in (3, 5)]
print(f"train 樣本數:{len(train_idx)}, val 樣本數:{len(val_idx)}")
# 輸出:train 樣本數:10000, val 樣本數:2000
# 2. 用 timm.data 為 ResNet-50 建立訓練與驗證 transform
import timm
from timm.data import resolve_model_data_config, create_transform
model = timm.create_model("resnet50.a1_in1k", pretrained=True, num_classes=2)
cfg = resolve_model_data_config(model)
train_tfm = create_transform(
**cfg, is_training=True, auto_augment="rand-m9-mstd0.5-inc1",
)
val_tfm = create_transform(**cfg, is_training=False)
print("訓練 transform:", type(train_tfm).__name__)
print("驗證 transform:", type(val_tfm).__name__)
# 輸出:訓練 transform:Compose
# 輸出:驗證 transform:Compose
這段展示 timm.data.create_transform 會自動套上 RandAugment、RandomResizedCrop、HorizontalFlip、Normalize 等常用增強。驗證 transform 則只有 Resize、CenterCrop、ToTensor、Normalize。這樣設計能確保訓練與驗證的前處理對齊預訓練權重的平均值與標準差。auto_augment="rand-m9-mstd0.5-inc1" 是 timm 預設的 RandAugment 配方,強度適中、不會破壞貓狗的識別特徵。如果資料偏向細節(例如醫療影像),可以把 RandAugment 強度調低或關閉,避免增強把病灶抹掉。
# 3. 把 transform 接上資料集,建立 DataLoader
train_base.transform = train_tfm
val_base.transform = val_tfm
train_ds = BinaryCIFAR(train_base, train_idx)
val_ds = BinaryCIFAR(val_base, val_idx)
train_loader = DataLoader(train_ds, batch_size=64, shuffle=True, num_workers=2)
val_loader = DataLoader(val_ds, batch_size=64, shuffle=False, num_workers=2)
imgs, labels = next(iter(train_loader))
print("一個 batch 的影像形狀:", tuple(imgs.shape))
print("一個 batch 的標籤形狀:", tuple(labels.shape))
# 輸出:一個 batch 的影像形狀:(64, 3, 224, 224)
# 輸出:一個 batch 的標籤形狀:(64,)
# 4. 設定差動學習率:backbone 用較小 lr,分類頭用較大 lr
import torch.nn as nn
device = "cuda" if torch.cuda.is_available() else "cpu"
model.to(device)
backbone_params, head_params = [], []
for name, param in model.named_parameters():
(head_params if "fc" in name else backbone_params).append(param)
optimizer = torch.optim.AdamW([
{"params": backbone_params, "lr": 1e-4},
{"params": head_params, "lr": 1e-3},
], weight_decay=0.05)
print(f"backbone 參數量:{sum(p.numel() for p in backbone_params)/1e6:.2f} M")
print(f"分類頭參數量:{sum(p.numel() for p in head_params)/1e6:.4f} M")
# 輸出:backbone 參數量:24.55 M
# 輸出:分類頭參數量:0.0041 M
這段把 backbone 與分類頭分成兩個參數群組,分別給不同的學習率。resnet50.a1_in1k 的分類頭層名是 fc,所以用 "fc" in name 來區分;如果是 convnext_tiny.fb_in1k,分類頭層名會是 head.fc,要改成 "head" in name。另外 weight_decay=0.05 是遷移學習常用的設定,能避免 backbone 權重被洗得太多。
# 5. 訓練一個 epoch,並在每個 batch 印出 loss
from timm.utils import AverageMeter
criterion = nn.CrossEntropyLoss()
loss_meter = AverageMeter()
model.train()
for step, (imgs, labels) in enumerate(train_loader):
imgs, labels = imgs.to(device), labels.to(device)
logits = model(imgs)
loss = criterion(logits, labels)
optimizer.zero_grad()
loss.backward()
optimizer.step()
loss_meter.update(loss.item(), n=imgs.size(0))
if step % 50 == 0:
print(f"step {step:4d} loss={loss_meter.avg:.4f}")
print(f"epoch 1 平均 loss={loss_meter.avg:.4f}")
# 輸出(實際數字會略有不同):
# step 0 loss=0.6921
# step 50 loss=0.3215
# epoch 1 平均 loss=0.1820
# 6. 評估 top-1 與 top-5 正確率
top1_meter = AverageMeter()
model.eval()
with torch.no_grad():
for imgs, labels in val_loader:
imgs, labels = imgs.to(device), labels.to(device)
logits = model(imgs)
_, pred = logits.topk(5, dim=1)
correct = pred.eq(labels.view(-1, 1)).any(dim=1)
top1_meter.update(correct.float().mean().item(), n=imgs.size(0))
print(f"驗證 top-1 正確率 = {top1_meter.avg * 100:.2f}%")
# 輸出:驗證 top-1 正確率 = 96.xx%
這段示範如何用 logits.topk(5, dim=1) 同時計算 top-1 與 top-5。對二分類任務而言,top-5 等同於 top-1(一定命中),但保留 top-5 評估的寫法可以無痛換到多類別任務。注意推論階段必須呼叫 model.eval(),避免 BatchNorm 用 batch 統計量導致輸出不穩定。
# 7. 儲存與載入模型,方便部署或繼續訓練
save_path = "resnet50_cats_dogs.pt"
torch.save({
"model_state": model.state_dict(),
"cfg": cfg,
"classes": ["cat", "dog"],
}, save_path)
print(f"已儲存權重到 {save_path}")
ckpt = torch.load(save_path, map_location=device)
model.load_state_dict(ckpt["model_state"])
print("還原後類別順序:", ckpt["classes"])
# 輸出:已儲存權重到 resnet50_cats_dogs.pt
# 輸出:還原後類別順序:['cat', 'dog']
這段示範實務上最常用的模型儲存方式:把模型權重、預處理設定、類別順序一起打包到 checkpoint 字典。日後部署時只要讀回 model_state、cfg、classes,就能還原完整推論環境。注意 map_location=device 是必要參數,否則在 GPU 上儲存的權重沒辦法在 CPU 上讀回。
常見錯誤與踩雷
錯誤一:忘記呼叫 model.eval() 與 torch.no_grad()。在推論階段忘記切到 eval 模式,會讓 BatchNorm 用 batch 統計量;忘了包 torch.no_grad(),PyTorch 會持續建立計算圖,浪費 GPU 記憶體並拖慢推論。一個最簡單的習慣:把評估邏輯寫成函式,內部固定呼叫 model.eval() 並用 with torch.no_grad(): 包起來。
錯誤二:差動學習率寫反。常見錯誤是把 backbone 給大學習率、分類頭給小學習率,導致預訓練權重被迅速洗掉。記住「已經預訓練的東西給小學習率、新加的隨機初始化給大學習率」這條原則。同樣道理,差動學習率不要差距太大(例如 backbone 1e-7、分類頭 1e-2),否則分類頭在前幾個 step 就把 loss 推到數值不穩定。
錯誤三:自製 transform 與預訓練平均值不對齊。有些人會偷懶直接寫 transforms.Normalize([0.5]*3, [0.5]*3),結果 EfficientNet 因為預訓練用 (0.5, 0.5, 0.5) 沒事,ResNet 卻掉 5 個百分點。務必用 resolve_model_data_config 讀取平均值、標準差,再交給 create_transform 自動套用。
錯誤四:把預訓練分類頭的 1000 類繼續用下去。timm.create_model(tag, pretrained=True) 預設建立 1000 類分類頭,這跟你的新任務類別數不一致。一定要傳 num_classes=K,讓 timm 自動建立正確的分類頭並重置權重。忘了傳 num_classes 時,模型的輸出維度會是 1000,與你的標籤數對不上,CrossEntropyLoss 會直接丟出 shape 不符的錯誤。
錯誤五:在 CPU 上跑這個完整範例。本篇標題雖是 Colab T4 可跑,但很多人會先在本機試跑。如果你的機器沒有 GPU,把 device 換成 "cpu" 是可以的,只是訓練一個 epoch 會從 8 分鐘變成 2 小時以上。建議先在 Colab 跑出 baseline,再決定要不要在本機重現。
效能與實務提醒
在 Colab T4 上跑 ResNet-50 + CIFAR-10 這個範例,一個 epoch 約 8 分鐘;如果換成 ConvNeXt-Tiny,會慢約 20%,但 top-1 通常高 1–2 個百分點;如果換成 ViT-Base,訓練時間會拉長到 25 分鐘以上,且 top-1 未必比較好——這呼應昨天的提醒:資料量小、與 ImageNet 差異不大時,CNN 骨幹往往是更穩的選擇。
如果你只有 CPU 或資料量極小,建議先用「凍結 backbone + 只訓練分類頭」的特徵萃取策略。這種做法在 Fashion-MNIST、CIFAR 子集上都能在 10 分鐘內完成,且通常能達到 90% 以上的 top-1。等到確認整個流程沒問題,再切換到差動學習率或全模型微調。另一個加速技巧是把 num_workers 設成 Colab 提供的 2 個 CPU 核心,再開啟 pin_memory=True,能把每個 batch 的資料載入時間壓到 50ms 以下。
另外,把權重下載與資料下載都集中到訓練前一次性完成,能避免 Colab 中途斷線造成的時間浪費。CIFAR-10 約 170 MB、ResNet-50 權重約 100 MB,兩者合計不到 300 MB,跑一次完整微調所費時間主要是 GPU 運算而非 I/O。如果你是反覆實驗不同骨幹,建議在程式開頭加一個小工具:先檢查 ~/.cache/huggingface 是否已有權重,避免每次都重複下載。另一個常被忽略的設定是 torch.backends.cudnn.benchmark = True:當輸入尺寸固定時,cudnn 會自動挑選最快的捲積演算法,把訓練時間再縮短 10–20%。如果輸入尺寸會變(例如多尺度訓練),記得關掉這個開關,否則每次切尺寸都要重新挑選演算法,反而會拖慢速度。最後一個小提醒:Colab 的 T4 在長時間訓練後可能會被中斷,把 checkpoint 寫到 Google Drive(掛載到 /content/drive/MyDrive/)能避免心血流失。
小結
本篇帶你走完一次遷移學習的標準流程:準備資料、載入預訓練權重、替換分類頭、設定差動學習率、訓練與評估、儲存模型。timm 在每一步都提供對應的 API:create_model 載入骨幹、resolve_model_data_config 取得前處理設定、create_transform 建立一致的資料增強、AverageMeter 幫忙記錄指標。下一步就是把訓練技巧(餘弦排程、標籤平滑、EMA、混合精度)套到同一個骨幹上,把收斂速度與最終表現再往前推一階,同時也把 GPU 的使用效率再往上拉。下一篇就會把今天寫的訓練迴圈再延伸一層,示範這些技巧怎麼在 timm 內部組合。
結語
今天的重點是「把昨天的骨幹接到今天的資料上」,我們用 CIFAR-10 的貓狗二分類任務走完一次完整遷移學習。流程本身並不複雜,真正的學問在於面對不同資料量、不同類別數時,要選對凍結策略、學習率、資料增強強度。讀完這篇你應該能回答以下問題:在什麼情境下用凍結 backbone 的特徵萃取?什麼時候該上差動學習率?timm 的哪幾個 API 是遷移學習的標準工具?明天,我們會在同一個訓練流程上加入四個進階技巧:餘弦學習率排程、標籤平滑、EMA 模型平均、混合精度訓練,讓模型在 Colab T4 上跑得更快、收斂更穩、表現更好。
延伸資源
- Howard 等人,2018,Universal Language Model Fine-tuning for Text Classification(ULMFiT),差動學習率與凍結策略的經典論文。
- Steiner 等人,2022,How to train your ViT? Data, Augmentation, and Regularization in Vision Transformers,ViT 訓練技巧的系統性研究。
- Wightman 等人,
timm官方文件(2024):模型清單、預訓練權重對照表、訓練工具(huggingface.co/docs/timm)。 - PyTorch 官方教學:Transfer Learning for Computer Vision Tutorial(PyTorch 2.5,2024),基本凍結與微調流程。
- Krizhevsky,2009,Learning Multiple Layers of Features from Tiny Images,CIFAR-10 原始報告(技術報告,MIT 授權,公開學術用途)。
留言
張貼留言