对比搜索在 Transformers 中实现人类水平的文本生成
Hugging Face 推出了 Contrastive Search,这是一种用于神经文本生成的最先进解码方法,现已在 transformers 库中提供,支持 PyTorch 和 TensorFlow。Contrastive Search 解决了确定性与随机性解码之间的关键权衡,使得即插即用的语言模型能够生成达到人类水平、语义连贯的文本,避免了贪婪搜索常见的重复以及核采样导致的不连贯。
问题:模型退化 vs. 语义不连贯
现有的解码方法大致分为两类,它们都存在固有缺陷,导致文本质量下降:
确定性方法(贪婪搜索和束搜索)
确定性方法会选择概率最高的续写。这常常导致 模型退化,即生成的文本显得不自然并出现不希望的重复。例如,使用 GPT-2 Large 进行贪婪搜索时,模型往往会多次重复相同的句子。
随机性方法(Top‑k 与核采样)
随机性方法通过引入随机性来避免重复。虽然核采样(top‑p)可以消除重复,但往往难以保持 语义连贯性,导致生成的短语在逻辑上与前缀文本脱节。降低 temperature 虽能在一定程度上缓解,但会使模型趋向贪婪搜索,从而在重复与不连贯之间形成难以平衡的权衡。
对比搜索的工作原理
对比搜索通过同时考虑模型置信度和退化惩罚来优化令牌选择,确保输出既具有高概率,又与先前上下文保持区别。
解码目标
在给定前缀 $x_{<t}$ 时选择下一个令牌 $x_{t}$,该方法会评估一组 top‑k 预测 ($V^{k}$)。选择依据两个主要组成部分:
- 模型置信度:语言模型对候选令牌 $v$ 的概率预测。
- 退化惩罚:衡量候选令牌 $v$ 相对于先前上下文 $x_{<t}$ 的区分度。其计算方式为令牌 $v$ 的表示(模型在前缀与 $v$ 拼接后得到的表示)与上下文中所有已有令牌表示之间的最大余弦相似度。
这两个组成部分通过超参数 $\alpha$ 进行平衡。若 $\alpha = 0$,方法退化为普通的贪婪搜索。较大的 $\alpha$ 会对与已有上下文过于相似的令牌施加更高惩罚,从而防止重复。
性能与可视化证据
对比搜索生成的文本在语法上流畅、语义上连贯且事实可靠。在使用 GPT-2 Large 和 Meta 的 OPT-1.3b 进行测试时,该方法能够生成长篇文档(最长 512 个令牌),保持一致的叙事结构,避免了贪婪搜索中出现的重复循环。
令牌相似度可视化
- 贪婪搜索:在非对角线位置显示高相似度分数,表明存在重复的令牌和模式。
- 对比搜索:高相似度分数主要出现在对角线位置,验证了退化问题已成功得到缓解。
在 Transformers 中的实现
对比搜索已集成到 transformers 库(版本 4.24.0 及以上)。用户可以通过在 generate 方法中指定以下超参数来使用它:
penalty_alpha:控制退化惩罚重要性的 $\alpha$ 超参数。top_k:搜索过程中考虑的 top 预测数量。
示例实现(针对 GPT-2 模型):
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 种不同语言的现成模型上的有效性。