關係網路 是什麼?
Relation Networks:關係網路 的完整解釋
學習查詢樣本與支持樣本之間的相似度函數,透過神經網路端到端預測樣本間的關係。
關係網路(Relation Networks)的提出基於一個重要的認識:在不同的任務、特徵空間、甚至不同的模態中,「什麼構成相似性」是不同的。固定的距離度量(如歐氏距離或餘弦相似度)往往無法捕捉這種變化的語義相似性。關係網路透過訓練一個神經網路來學習適應當前任務和特徵空間的相似度判斷。
關係網路的工作原理
關係網路的分類流程包含三個主要步驟:特徵提取、對偶拼接與關係推理。
第一步:特徵提取
使用編碼器網路 f_φ(如 CNN)分別提取查詢樣本和支持樣本的特徵:
- 查詢特徵:f_φ(q)
- 支持特徵:f_φ(s_{i,j})
第二步:對偶拼接(Pair-wise Concatenation)
將查詢樣本的特徵與每個支持樣本的特徵拼接:
z_{i,j} = concat(f_φ(q), f_φ(s_{i,j}))
拼接操作保留了查詢樣本和支持樣本的分別資訊,便於下一步的關係學習。
第三步:關係推理(Relation Module)
一個神經網路 g_ψ(通常為簡單的 MLP 或小型 CNN)將拼接的特徵映射為 [0, 1] 範圍內的關係分數:
r_{i,j} = g_ψ(z_{i,j})
其中 r_{i,j} ∈ [0, 1] 代表查詢樣本與支持樣本 s_{i,j} 的關係強度(相似度)。
分類決策
對於 N-way K-shot 任務,查詢樣本屬於類別 c 的概率為:
P(y=c|q) = Σ_{j=1}^K r_{c,j} / K
即該類別所有支持樣本的平均關係分數。
若使用更靈活的評分方式,可直接選擇關係分數最高的支持樣本所在類別:
class(q) = argmax_c max_j r_{c,j}
與原型網路和匹配網路的比較
| 方法 | 距離/相似度計算 | 靈活性 | 計算複雜度 | 參數量 |
|---|---|---|---|---|
| 原型網路 | 固定(歐氏或餘弦) | 低 | O(Nd) | 僅編碼器 |
| 匹配網路 | 簡單相似度 + 注意力 | 中 | O(NKd) | 編碼器 + 相似度函數 |
| 關係網路 | 學習的神經網路 | 高 | O(NKd') + 關係網路 | 編碼器 + 關係網路 |
其中 d 為特徵維度,d' 為拼接後維度,K 為 shot 數,N 為 way 數。
關係網路的訓練
訓練目標是最小化查詢集上的二元交叉熵損失(針對關係分數)或交叉熵損失(針對類別概率):
若使用二元損失(每個支持樣本與查詢樣本的關係獨立推斷): L = Σ_{i,j} [y_{i,j} × log(r_{i,j}) + (1-y_{i,j}) × log(1-r_{i,j})]
其中 y_{i,j} = 1 if s_{i,j} 與 q 同類,else 0。
梯度透過關係網路 g_ψ 和編碼器 f_φ 反向傳播。與原型網路不同,關係網路的訓練信號直接監督「什麼樣的拼接特徵對應高關係」,這使模型能學習到更任務特定的相似度判斷。
關係網路的優勢
高度可學習的相似度:神經網路相似度函數可捕捉複雜的、非線性的相似度關係,遠超線性度量。
特徵空間與相似度的聯合優化:編碼器和關係網路同時訓練,能找到對當前任務最優的特徵表示與相似度判斷。
在複雜分布上的性能優勢:當類別內部存在複雜的非線性結構時,關係網路的學習能力使其明顯優於原型網路和匹配網路。
一般化能力:關係網路的框架不僅適用於分類,也可用於回歸、排序等任務,只需改變目標變數的定義。
關係網路的限制與挑戰
過擬合風險:關係網路有更多可學習參數,在少樣本設定中容易過擬合。為緩解,需要精心設計正則化、資料增強或元訓練策略。
計算和記憶成本:相比原型網路的簡單距離計算,關係網路需要運行神經網路 N×K 次(對每個支持樣本),計算複雜度較高。
特徵空間與相似度的競爭:編碼器和關係網路都在學習相似度的某個方面,可能導致目標衝突。某些設計需要小心平衡兩者的訓練動力。
改進與變體
圖神經網路(GNN)變體:將支持樣本和查詢樣本看作圖的節點,使用圖卷積或其他圖操作來學習關係,相比簡單的拼接更有結構性。
多層關係:在多個特徵階層上計算關係,聚合得出最終判斷,提升對多尺度語義的捕捉。
條件關係網路:關係網路的參數隨著支持集的改變而改變(條件化),進一步增強適應性。
蒸餾與輕量化:將已訓練的關係網路蒸餾為更簡單的模型(如原型網路),保留性能同時降低推論成本。