匹配網路 是什麼?

Matching Networks:匹配網路 的完整解釋

結合注意力機制的度量學習方法,透過為支持樣本分配動態權重來預測查詢樣本的類別。

匹配網路(Matching Networks for One Shot Learning)的提出標誌著注意力機制在少樣本學習中的應用開端。與同時期的原型網路相比,匹配網路提供了一個更靈活的框架,能夠根據查詢樣本的特性動態調整支持樣本的權重,從而在許多場景下達到更優的性能。

匹配網路的核心思想

匹配網路的核心洞察是:在進行少樣本分類時,並非所有支持樣本對查詢樣本的貢獻都相等。例如,在狗品種識別中,如果查詢樣本是一隻特定姿態的棕色狗,那麼支持集中那些棕色、相似姿態的狗應該給予更高的權重,而不同顏色或姿態的狗應降低權重。匹配網路透過學習一個注意力函數來實現這種自適應的權重分配。

匹配網路的數學框架

設支持集包含 N 個類別,每類 K 個樣本:S = {(s_{i,j}, y_{i,j})},查詢樣本為 q。

匹配網路的分類概率定義為:

P(y|q, S) = Σ_{i,j} a(q, s_{i,j}) × p(y|s_{i,j})

其中:

  • a(q, s_{i,j}) 是注意力權重,衡量查詢樣本 q 與支持樣本 s_{i,j} 的匹配程度。
  • p(y|s_{i,j}) = 1 if s_{i,j} 屬於類別 y,else 0(one-hot 編碼)。

注意力權重透過指數相似度歸一化計算:

a(q, s_{i,j}) = exp(c(q, s_{i,j})) / Σ_{i',j'} exp(c(q, s_{i',j'}))

其中 c(q, s) 是相似度函數(Cosine Distance 或其他)。

匹配網路與原型網路的關鍵區別

原型網路先計算原型(平均值),再計算查詢樣本與原型的距離:

P(y|q) = softmax(-||f(q) - p_y||²)

其中 p_y = (1/K) Σ_j f(s_{y,j})

匹配網路則直接對每個支持樣本計算相似度,並加權平均:

P(y|q) = Σ_j a(q, s_{y,j})

直觀上,原型網路的權重是均勻的(每個支持樣本 1/K),而匹配網路的權重是動態的(a(q, s_{i,j}))。這個區別看似微妙,但在實踐中能帶來顯著的性能差異,尤其是在支持樣本分布異質或包含異常值時。

相似度函數的設計

匹配網路中的相似度函數 c(q, s) 有多種選擇:

  1. 簡單相似度

    • 餘弦相似度:c(q, s) = f(q)ᵀ f(s) / (||f(q)|| × ||f(s)||)
    • 歐氏距離:c(q, s) = -||f(q) - f(s)||²
  2. 學習的相似度

    • 使用神經網路:c(q, s) = g_φ(f(q), f(s)),其中 g_φ 是參數化的函數。
    • 雙尖度評分(Bilinear Scoring):c(q, s) = f(q)ᵀ W f(s),其中 W 是可學習的權重矩陣。

記憶增強與訓練-測試一致性

匹配網路的一個獨特設計是使用外部記憶和閱讀機制(Memory-Augmented)。支持集被儲存在一個可讀的記憶中,在推論時動態讀取。這個設計的目的是在訓練和測試中保持一致性:訓練時也把支持集看作外部記憶動態讀取,而非預先編碼成固定的特徵。這種一致性幫助模型在訓練中適應「變化的支持集」的情況。

損失函數與訓練目標

匹配網路的訓練目標是最小化查詢集上的交叉熵損失:

L = -Σ_q log P(y_q | q, S)

即使 P(y_q|q, S) 是注意力加權和而非直接的神經網路輸出,梯度仍可透過注意力權重反向傳播至特徵提取器和相似度函數。

匹配網路的優勢

  1. 動態權重機制:注意力機制賦予模型根據具體查詢樣本調整權重的能力,相比固定的平均值更靈活。

  2. 在多模態分布上的魯棒性:當類別內存在多個子類群(如「狗」分為小型狗、大型狗等),匹配網路能自動尋找最相似的支持樣本,而原型網路的單一原型可能不在這些子類群的任何一個上。

  3. 端到端可微分:整個管道(注意力計算、特徵提取、相似度函數)都可微分,允許端到端訓練。

匹配網路的限制

  1. 計算複雜度較高:需要對每個查詢樣本計算與所有支持樣本的相似度,時間複雜度為 O(N×K)。在大規模推論中可能成為瓶頸。

  2. 對支持集大小敏感:K 增加時,計算量線性增長。相比之下,原型網路計算時間與 K 無關(只需平均值)。

  3. 噪音支持樣本的影響:若支持集中有誤標記,注意力機制會自動給其高權重(如果它恰好相似),導致分類錯誤。原型網路的「稀釋」效應(誤樣本的影響被平均掉)可能更魯棒。

匹配網路的改進與變體

  1. 匹配網路與元學習的融合:後續研究發現,直接在元訓練目標上優化匹配網路的注意力機制,相比監督微調能進一步提升性能。

  2. 層次化匹配:在多層特徵上計算相似度,而非單一層的特徵,提升匹配精度。

  3. 硬注意力(Hard Attention):改進軟注意力的計算成本,對每個查詢樣本選擇 top-K 相似的支持樣本,而非全部考慮。

匹配網路 vs. 相關方法的實證比較

在 omniglot 和 miniImageNet 等基準上的典型結果:

  • Matching Networks:相比隨機初始化提升 50-60%,是當時的 SOTA。
  • Prototypical Networks(隨後提出):性能相當,但計算更快。
  • MAML:訓練和推論的準確度都更高,但計算成本更大。
  • Relation Networks:性能優於匹配網路和原型網路,但模型複雜度更高。

常見問題