對比搜尋在 Transformers 中的人類水平文本生成
Hugging Face 已推出 Contrastive Search,這是一種最先進的神經文本生成解碼方法,現在已在 transformers 套件中支援 PyTorch 與 TensorFlow。對比搜尋解決了確定性與隨機性解碼之間的關鍵取捨,使即用的語言模型能產生人類水平、語義連貫的文本,且不會出現貪婪搜尋常見的重複或核取樣所帶來的不連貫。
問題:模型退化 vs. 語義不連貫
現有的解碼方法大致可分為兩類,兩者皆存在固有缺陷,會降低文本品質:
確定性方法(貪婪與束搜索)
確定性方法會選擇機率最高的延續。這常導致 model degeneration,使生成的文本變得不自然且出現不必要的重複。例如,使用 GPT-2 Large 進行貪婪搜尋時,模型往往會重複相同的句子多次。
隨機性方法(Top-k 與核取樣)
隨機性方法引入隨機性以避免重複。雖然核取樣(top-p)能消除重複,但常無法維持 semantic coherence。這會導致生成的片語在邏輯上與前綴文本斷裂。降低 temperature 雖能緩解此問題,卻會使模型回到貪婪搜尋的狀態,形成在重複與不連貫之間的困難取捨。
對比搜尋的運作方式
對比搜尋透過同時考量模型信心與退化懲罰來最佳化 token 選擇,確保輸出既具高機率又與先前上下文不同。
解碼目標
在給定前綴 $x_{<t}$ 時選擇下一個 token $x_{t}$,此方法會評估一組 top‑k 預測 ($V^{k}$)。選擇依據兩個主要組件:
- Model Confidence:語言模型對候選 token $v$ 的機率預測。
- Degeneration Penalty:衡量候選 $v$ 相對於先前上下文 $x_{<t}$ 的區辨性。其計算方式為 token $v$ 的表示(模型在前綴與 $v$ 串接後所計算)與已存在於上下文中所有 token 表示之最大餘弦相似度。
這些組件透過超參數 $\alpha$ 進行平衡。若 $\alpha = 0$,方法會退回至普通的貪婪搜尋。較高的 $\alpha$ 會提升對過於相似於現有上下文的 token 的懲罰,從而防止重複。
效能與視覺證據
對比搜尋產生的文本在語法上流暢、語義上連貫且事實上有根據。在使用 GPT-2 Large 與 Meta 的 OPT-1.3b 進行測試時,該方法生成了長篇文件(最多 512 個 token),保持一致的敘事,且避免了貪婪搜尋中常見的重複循環。
Token 相似度視覺化
對 token 相似度矩陣的視覺分析揭示了退化懲罰的技術成功:
- Greedy Search:在非對角線項目中顯示高相似度分數,表明存在重複的 token 與模式。
- Contrastive Search:高相似度分數主要出現在對角線項目,驗證退化問題已成功緩解。
在 Transformers 中的實作
對比搜尋已整合至 transformers 套件(版本 4.24.0 及以上)。使用者可透過 generate 方法並指定以下超參數來實作:
penalty_alpha:調節退化懲罰重要性的 $\alpha$ 超參數。top_k:搜尋過程中考慮的 top 預測數量。
output = model.generate(input_ids, penalty_alpha=0.6, top_k=4, max_length=512)
研究基礎
此實作基於兩篇主要研究論文:
- A Contrastive Framework for Neural Text Generation(NeurIPS 2022),提出了原始框架。
- Contrastive Search Is What You Need For Neural Text Generation(2022),展示了該方法在 16 種不同語言的即用模型上的效能。