搜尋意圖: 如果你在找「原型網路 是什麼」或「原型網路 和相近概念差在哪」,先看這頁的短定義、完整說明與延伸比較。
TL;DR: 基於度量的少樣本學習方法,透過計算支持樣本的類別原型(均值),以距離判斷查詢樣本的類別。
實用情境: 適合用在閱讀 AI 文章、產品文件或和同事討論時,先用一頁快速對齊概念。
下一步: 先讀完定義,再往下看延伸比較與對應工具,把概念轉成實際應用。
基於度量的少樣本學習方法,透過計算支持樣本的類別原型(均值),以距離判斷查詢樣本的類別。
原型網路(Prototypical Networks)的提出源於一個簡單但深刻的觀察:在特徵空間中,同一類別的樣本往往聚集在一起,不同類別的樣本相離散。若能學到一個好的特徵表示空間,那麼分類就退化為簡單的距離計算:找出查詢樣本最近的原型(Prototype)即可判斷其類別。這個思想既直觀又高效,在少樣本學習中展現了強大的性能。
原型網路的工作原理
原型網路的運作分為兩個階段:特徵提取和相似度匹配。
第一階段:特徵提取(Embedding)
使用一個參數化的特徵提取器(如 CNN)f_φ,將原始輸入映射到一個 d 維的特徵空間。此階段的目標是學到一個使得同類樣本緊集、不同類樣本相離散的特徵空間。
第二階段:原型計算與分類
給定 N-way K-shot 的任務(N 個類別、每類 K 個支持樣本):
對每個類別 c,計算其原型 p_c : 即該類別所有支持樣本的特徵平均值: p_c = (1/K) Σ_{i=1}^K f_φ(s_{c,i}) 其中 s_{c,i} 是類別 c 的第 i 個支持樣本。
對於查詢樣本 q,計算其特徵 f_φ(q)。
計算查詢樣本與各類別原型的距離(通常使用歐氏距離): d_c = ||f_φ(q) - p_c||²
將查詢樣本分類為最近的原型所代表的類別: class(q) = argmin_c d_c
或使用軟分類,以距離的相反數作為相似度,通過 softmax 計算概率分布: P(y=c|q) = exp(-d_c) / Σ_{c'} exp(-d_{c'})
原型網路的訓練目標
在元訓練階段,對每個任務採樣,原型網路最小化查詢集上的分類交叉熵損失:
L = -Σ_q log P(y=y_q | q)
梯度透過特徵提取器 f_φ 反向傳播。訓練目標是調整 f_φ,使得在任何新任務上,同類樣本的特徵更接近,不同類樣本的特徵更遠離。
原型網路的優勢
簡潔高效:無需在每個任務上優化分類器參數,特徵提取後直接計算距離分類,推論速度快。
可解釋性強:原型的概念直觀,可視化原型能理解模型的決策邊界。
少樣本適配性好:只需計算平均值和距離,對小樣本不敏感,在 K 很小時穩定性好。
靈活的距離度量:可使用不同的距離函數(歐氏距離、餘弦相似度、學習的距離函數等),便於模型定制。
原型網路的限制與改進
- 類別內分布假設:原型網路假設每個類別的樣本在特徵空間中均勻分布在原型周圍,但現實中類別往往呈現多峰分布(Multimodal)。例如,「狗」的類別可能包括不同大小、顏色、姿態的狗,用單一原型可能不夠。
改進方案:
- 高斯原型網路(Gaussian Prototypical Networks):為每個類別建模一個高斯分布而非單點,能捕捉類別內的變異。
- 多原型方法:為每個類別學習多個原型,用聚類或自適應方式決定每個樣本屬於哪個亞類別。
- 類別間分布的非均衡:在不同的特徵空間中,不同類別的特徵分布密度可能不同,簡單的距離閾值不適用於所有類別。
改進方案:
- 自適應距離度量:學習一個非歐氏的、類別特定的度量函數。
- 元選擇器(Meta-selector):為每個任務學習一個自適應的距離度量。
- 噪音支持樣本的敏感性:若支持集中包含錯誤標注或異常樣本,原型計算會被污染,影響後續分類。
改進方案:
- 魯棒原型網路:使用中位數或加權平均而非簡單平均,增強對離群值的抵抗力。
- 注意力機制:為支持樣本分配不同權重,自動降低異常樣本的影響。
原型網路與其他度量學習方法的比較
原型網路 vs. 匹配網路(Matching Networks):原型網路用簡單的平均計算原型,匹配網路則使用注意力機制為每個支持樣本分配權重。在 shot 數量充足時兩者性能接近,但原型網路計算更快。
原型網路 vs. 關係網路(Relation Networks):原型網路固定距離度量(通常歐氏距離),關係網路學習一個神經網路來預測查詢樣本和原型的相似度,靈活性更高但計算複雜度也更高。
原型網路 vs. MAML:原型網路無需梯度優化新分類器,推論快且對小樣本穩定;MAML 需要梯度步,計算複雜但在資料有限時可能泛化更好。兩者往往相輔相成,有研究結合兩者優點。
實務應用
原型網路在以下場景中得到廣泛應用:
- 零樣本學習(Zero-shot Learning):利用類別語義信息生成偽原型。
- 開集識別(Open-set Recognition):未見過的類別可視為距離所有已知原型都很遠。
- 持續學習:新類別的原型可逐漸加入而無需重訓練。
- 跨模態檢索:學習共用的特徵空間,使不同模態(文字、影像)的原型相近。
常見問題
為什麼原型網路計算類別原型時用平均值而不是其他統計量?
平均值作為原型的選擇既簡潔又在數學上有理。從概率角度,若假設每個類別的樣本獨立同分布(i.i.d.)來自某個分佈,那麼樣本均值是該分佈均值的最大似然估計。從最小化方差的角度,樣本均值也是使類別內點到該點距離平方和最小的點。當支持集樣本數較多時,平均值是穩定的統計量;但當 K 很小(1-5 shot)時,單個噪音樣本的影響會相對放大。這也是為什麼後續研究提出了使用中位數、加權平均或多原型的改進方法。選擇平均值的深層原因是在少樣本設定中尋求計算簡潔性與統計穩定性的平衡。
原型網路能處理多標籤分類或不平衡的類別分布嗎?
標準原型網路設計用於單標籤分類,每個樣本只屬於一個類別。對於多標籤問題,需要修改:可為每個標籤單獨訓練原型網路(Binary Relevant 方法),或設計聯合特徵空間使多個原型同時激活。至於類別不平衡,原型網路本身對不平衡不敏感,因為分類決策基於距離而非概率校準。若支持集中不同類別的樣本數不等(K_1 ≠ K_2),簡單平均可能被樣本多的類別主導;解決方案是使用加權平均或損失函數中加入類別權重。實務上,若類別極度不平衡,建議搭配類別權重或採樣策略。
原型網路在實際部署中相比基於梯度的元學習有什麼優勢?
原型網路最大的優勢是推論速度快且無需梯度計算。給定預訓練的特徵提取器,對新任務的適應只需提取支持集的特徵、計算均值、再計算距離:整個過程是向前傳播 + 簡單數值計算,無需反向傳播。相比 MAML 需要在新任務上執行多步梯度更新,原型網路的推論時間通常快 100 倍以上。另一個優勢是它對噪音標註相對魯棒,因為壞的支持樣本只是拉動原型位置一點點,而不像基於梯度的方法會產生過大的梯度。然而,若新任務與訓練任務分布有較大偏移,梯度優化方法(MAML)可能適應能力更強。實務選擇往往取決於『對推論延遲的容忍度』和『任務分布穩定性』。