跳到主要內容

Day 34 遷移學習

Day 34 遷移學習

引言

前一篇我們用 CNN 在 CIFAR-10 上從頭訓練一個影像分類模型,從資料增強到卷積層設計都親手做了一遍。不過實務上,自己從零訓練一個大型模型既耗時又吃資源,這時遷移學習(Transfer Learning)就派上用場了。遷移學習的核心精神是:拿別人已經訓練好的模型當起點,再針對自己的任務做微調。這種「站在巨人肩膀上」的做法,是現代深度學習最常見的工作流程之一。

這一篇會介紹遷移學習的兩種常見策略——特徵萃取(Feature Extraction)與微調(Fine-tuning),並用 torchvision 提供的 ResNet18 預訓練權重,示範怎麼在自家的小資料集上動手改造模型。完成今天的範例後,你就能用幾行程式碼,就讓小資料集也能交出不錯的成效,省下從零訓練的時間與運算資源。

什麼是遷移學習?

遷移學習是把一個任務上訓練好的模型,搬到另一個相關任務上使用。概念來自人類學習:學會騎腳踏車之後,再學騎機車會快很多,因為兩者都需要保持平衡。在深度學習裡,大型模型(例如在 ImageNet 上訓練的 ResNet、EfficientNet)已經學會了豐富的視覺特徵——邊緣、紋理、形狀甚至語意,這些知識對新的影像任務常常也很有用。想像一下,一個看過幾百萬張照片的模型,它對「什麼是線條」、「什麼是圓弧」的理解通常已經相當成熟;面對新任務時,只要再加上一點新資料做調整,就能把那些能力接著用下去。

什麼時候適合用遷移學習?當自己的資料集規模不大(幾百到幾萬張)、訓練資源有限、或想快速得到一個還不錯的 baseline,都可以考慮。當資料量真的很大、任務又和預訓練資料差很遠時,從頭訓練才有可能比遷移學習更划算。例如把一個針對自然照片訓練的模型,搬到醫療影像或衛星圖上,效果通常還是很不錯;但若換到風格完全不同的小樣本任務,遷移的效果就可能打折扣。

遷移學習主要有兩種做法:

  • 特徵萃取:把預訓練模型當作固定的特徵抽取器,只訓練自己新增的分類頭(classifier head)。這種方式訓練速度快,適合資料很少的場景。
  • 微調(Fine-tuning):解凍一部分預訓練層,讓它們在新資料上繼續學習。通常會用比較小的學習率,避免破壞原本學到的特徵。

兩種做法各有取捨:特徵萃取安全、穩定、省資源;微調彈性更高,但風險也比較大,需要對訓練過程多一份觀察。

載入預訓練的 ResNet18

torchvision 把許多經典模型打包好,只要一行就能下載預訓練權重。第一次執行時會自動從網路下載權重檔(檔名為 resnet18-f37072fd.pth),下載完成後會快取在 ~/.cache/torch/hub/checkpoints/,之後執行就不需要再連線。

import torch
from torchvision import models

# 載入在 ImageNet 上預訓練的 ResNet18
weights = models.ResNet18_Weights.IMAGENET1K_V1
model = models.resnet18(weights=weights)

print(model.fc)  # 最後的全連接層,原本輸出 1000 個類別
print("參數總數:", sum(p.numel() for p in model.parameters()))

ResNet18 預設會輸出 1000 個類別(對應 ImageNet 的 1000 個分類)。如果我們要做的是貓狗二分類,就要把最後一層 fc 換成輸出 2 個類別。除了最後一層以外,前面所有卷積層其實都可以直接拿來用,這也是遷移學習能省事的關鍵。

替換分類頭:特徵萃取策略

特徵萃取的做法非常直接:把 model.fc 換成新的線性層,並把整個 backbone 的參數凍結(requires_grad=False),這樣訓練時就只會更新新分類頭的權重。在資料量很少的場景,這個策略可以避免過擬合,也省下大量訓練時間。

import torch.nn as nn

# 凍結所有預訓練參數
for param in model.parameters():
    param.requires_grad = False

# 替換最後一層,輸出改成 2 類
model.fc = nn.Linear(in_features=512, out_features=2)

# 確認只有最後一層會更新
trainable = [p for p in model.parameters() if p.requires_grad]
print("可訓練參數數量:", sum(p.numel() for p in trainable))
print("總參數數量  :", sum(p.numel() for p in model.parameters()))

從輸出可以看到,可訓練的參數只有最後那 1024 個(512*2 + 2 個偏差值),其他四千多萬個參數全部凍結。這種「大頭不動、小頭訓練」的設定,是遷移學習最常見的入門組合。記得在這種情況下,optimizer 一定要用 filter(lambda p: p.requires_grad, model.parameters()),把可訓練的參數單獨抓出來,否則會出現「沒凍結也沒更新」的怪現象。

微調策略:解凍部分層一起訓練

如果資料量稍微多一些,可以考慮微調:解凍 backbone 的後半段,讓模型在新任務上做更細緻的調整。一般建議從比較後面的層開始解凍,因為前面的層通常負責比較低階的特徵(邊緣、顏色),跨任務的通用性高;後面的層則負責比較語意化的特徵(物件部件、整體輪廓),換任務時通常需要再學習。

# 重新載入一個乾淨的預訓練模型
model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)

# 先全部凍結
for param in model.parameters():
    param.requires_grad = False

# 解凍 layer4 與新的分類頭
for param in model.layer4.parameters():
    param.requires_grad = True
model.fc = nn.Linear(512, 2)

trainable = [p for p in model.parameters() if p.requires_grad]
print("可訓練參數數量:", sum(p.numel() for p in trainable))

執行後會看到可訓練參數明顯增加,但相對整個模型來說還是少數。微調時有兩個小眉角:學習率要比從頭訓練小很多(通常是原本的 1/10),不然很容易把原本學好的權重「推歪」;另外也常用差異化學習率—— backbone 用小學習率、新分類頭用較大學習率——讓兩邊的更新幅度更平衡。PyTorch 的 optimizer 支援傳入一個由參數群組組成的串列,每個群組各自設定 lr、weight_decay,這就是實作差異化學習率的標準做法。

訓練與評估迴圈

不管是特徵萃取還是微調,後面的訓練流程和 Day 30 介紹的基本一致:定義 optimizer、把資料與模型搬到裝置上、跑每個 epoch 的訓練與驗證。下方示範一個常見的寫法,並示範怎麼在凍結 backbone 的情況下設定 optimizer。

import torch.optim as optim

device = "cuda" if torch.cuda.is_available() else "cpu"
model.to(device)

# 只把可訓練參數傳給 optimizer
optimizer = optim.Adam(
    filter(lambda p: p.requires_grad, model.parameters()),
    lr=1e-3,
)
criterion = nn.CrossEntropyLoss()

model.train()
for images, labels in train_loader:
    images, labels = images.to(device), labels.to(device)
    optimizer.zero_grad()
    outputs = model(images)
    loss = criterion(outputs, labels)
    loss.backward()
    optimizer.step()

model.eval()
correct, total = 0, 0
with torch.no_grad():
    for images, labels in val_loader:
        images, labels = images.to(device), labels.to(device)
        preds = model(images).argmax(dim=1)
        correct += (preds == labels).sum().item()
        total += labels.size(0)
print(f"驗證正確率:{correct / total:.4f}")

實際跑下來,遷移學習在小型資料集上常常只要幾個 epoch 就能突破 90% 的正確率,比起從頭訓練一個 CNN 省下非常可觀的時間。這也是為什麼在業界,特別是在沒有大型 GPU 資源或標註資料有限的情況下,遷移學習幾乎是影像任務的預設起點。順帶提醒,預訓練模型通常會帶有專屬的 transforms(例如 ResNet18 預期 224x224 的影像、特定的平均值與標準差),用對應的 weights.transforms() 就能拿到正確的前處理流程,省去自己對照說明手冊的麻煩。

結語

這一篇從概念到程式碼,把遷移學習的兩種常見策略都示範了一遍。當你自己的影像資料不多、訓練資源有限時,拿 torchvision 內建的 ResNet18、EfficientNet 等預訓練模型當起點,往往比從頭訓練划算非常多。下次面對一個新影像任務時,先想想有沒有現成的預訓練權重可以借力,會省下大量時間。

明天,我們會進入序列資料的世界,看看 RNN 與 LSTM 是怎麼處理文字與時間序列這類帶有先後順序的資料。

留言

這個網誌中的熱門文章

Day 2 變數與資料型別

Day 2 變數與資料型別 引言 寫程式的過程中,變數與資料型別是處理資料的基礎。變數是存放資料的容器,資料型別則決定這筆資料有哪些特性、可以進行哪些操作。學會定義變數、認識各種資料型別,是學好 Python 的關鍵一步。 這篇文章會帶你了解 Python 中變數的觀念、如何定義變數,以及常見的資料型別,包括整數、浮點數、字串、布林值,還有串列、元組、字典與集合等容器型別。我們也會介紹變數的命名規則與撰寫風格建議,以及如何用 type() 檢查資料型別。 什麼是變數?如何在 Python 中定義變數 變數是在程式執行時用來存放資料的名稱。透過定義變數,我們可以給一筆資料一個名字,並在程式的其他地方用這個名字取用該筆資料。在 Python 中,變數不需要事先宣告型別,因為 Python 是動態型別語言,變數的型別由指定給它的值決定。 定義變數的基本語法 在 Python 中定義變數非常簡單,只要用賦值符號 = 把值指定給變數即可。例如: x = 5 # 定義變數 x,並把整數 5 賦值給它 name = "Alice" # 定義變數 name,並把字串 "Alice" 賦值給它 在這裡,x 是一個變數,被賦予整數 5;name 是另一個變數,被賦予字串 "Alice"。 變數的更新與覆寫 變數的值可以修改,也就是說,我們可以在程式的不同地方給同一個變數新的值。例如: x = 10 # x 最初被賦予 10 x = 15 # x 的值現在被更新為 15 這樣就能依照需求,在程式執行過程中靈活調整變數的值。 Python 的動態型別系統 Python 和某些靜態型別語言不同,定義變數時不需要宣告型別。賦值時,Python 會根據值自動判斷變數的型別。例如: x = 5 # x 是整數 x = 3.14 # x 變成浮點數 x = "Hi" # x 變成字串 同一個變數在程式執行過程中可以存放不同型別的值,這是 Python 的彈性之一。 常見資料型別 在 Python 中,資料型別決定我們可以對變數進行哪些操作...

Day 1 Python 簡介與環境設定

Day 1 Python 簡介與環境設定 引言 在現在的科技環境裡,程式設計已經是一項重要技能。無論你是對資料科學有興趣、想成為開發者,或是想踏入人工智慧(AI)領域,學會寫程式都能明顯提升你的競爭力。在眾多程式語言中,Python 因為語法簡單、功能強大、應用範圍廣泛,成為許多人進入程式世界的第一選擇。這篇文章會帶你認識 Python 的背景與優勢,並一步步教你在不同系統上安裝與設定 Python 開發環境,最後寫出第一支 Python 程式。 為什麼選擇 Python? Python 是一種高階程式語言,由 Guido van Rossum 在 1991 年發布。Python 的設計哲學強調程式碼的可讀性,並用縮排來定義程式區塊,這點和許多使用大括號的語言不同。簡潔的語法讓它成為初學者的理想選擇;就算是經驗豐富的開發者,也能用它完成複雜的專案。 Python 的優勢如下: 簡單易學 :Python 的語法清楚、結構簡潔,初學者很快就能上手。和其他語言相比,學習曲線相對平緩,不需要先弄懂一堆複雜觀念,就能開始寫程式。 應用範圍廣泛 :從資料科學、網頁開發、人工智慧、機器學習、自動化測試到網路爬蟲,Python 都有大量開源函式庫與工具支援,而且在這些領域都扮演關鍵角色。 豐富的函式庫與框架 :Python 的函式庫生態系非常龐大。做資料分析有 NumPy、Pandas;開發網站有 Django、Flask;做深度學習有 TensorFlow、PyTorch。各種需求幾乎都能找到對應的套件,讓開發更有效率。 跨平台支援 :Python 支援 Windows、macOS、Linux 等作業系統,程式通常不需要太多修改就能跨平台執行,讓開發與部署更有彈性。 活躍的社群 :Python 擁有龐大的開發者社群。學習或開發上遇到問題,幾乎都能在社群與論壇(例如 Stack Overflow)找到答案,對初學者來說是很強的後盾,也能減少卡關時的挫折感。 Python 的應用領域 Python 的流行與強大功能,讓許多領域都開始大量使用它。以下是幾個常見的應用方向: 資料科學 :隨著大數據與人工智慧興起,資料科學大量使用 Python。NumPy、Pandas 與 Matplotlib 等工具能處理和分析龐...

Python 從入門到 PyTorch 深度學習:開啟 AI 世界的大門

Python 從入門到 PyTorch 深度學習:開啟 AI 世界的大門 隨著人工智慧(AI)與深度學習(Deep Learning)快速發展,越來越多人對這些技術產生興趣。不論你是想踏入 AI 領域的初學者,還是已經有程式基礎的開發者,學好 Python 與深度學習框架(例如 PyTorch),都能為你打開更多可能。 為什麼選擇 Python? Python 已經是資料科學與人工智慧領域的首選語言。它的語法簡潔、容易上手,而且擁有龐大的生態系與大量開源函式庫。無論是資料處理、資料視覺化,還是建立機器學習與深度學習模型,Python 都能勝任。對想進入 AI 或資料科學領域的人來說,它幾乎是必備工具。 PyTorch 是什麼? PyTorch 是由 Meta(原 Facebook)AI 研究團隊開發的開源深度學習框架,以易用、靈活和動態計算圖著稱,是許多 AI 研究人員與開發者的首選。相較於其他框架,PyTorch 的寫法更貼近原生 Python,對初學者相對友善。無論是簡單的實驗,還是複雜的深度學習模型,PyTorch 都能提供強大的支援。 這個系列能帶給你什麼? 這個系列會從 Python 的基礎開始,帶你一步一步學習,最後能自己用 PyTorch 建立深度學習模型。即使你完全沒有寫過程式,也能跟著文章的節奏累積技能,理解 AI 與深度學習的核心觀念。 本系列涵蓋的主題 Python 基礎:從變數、條件判斷到函式與模組。 資料處理工具:用 NumPy 與 Pandas 有效率地操作資料。 資料視覺化:用 Matplotlib 與 Seaborn 把資料畫成圖表。 深度學習的數學基礎:線性代數、微積分與機率。 PyTorch 入門:理解張量、模型建構與 GPU 加速。 基礎深度學習模型:CNN 與 RNN 的實作應用。 深度學習專案實戰:從資料前處理到模型部署的端到端流程。 誰適合這個系列? 程式初學者 :如果你對 AI 充滿好奇,卻還沒寫過程式,系列的第一部分會帶你快速上手 Python,並幫助你理解深度學習的基本觀念。 資料科學愛好者 :如果你已經熟悉一些資料處理方法,進階部分會教你如何用 PyTorch 建構深度學習模型。 開發者與研究人員 :想更深入了...