对比搜索在 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}$)。选择依据两个主要组成部分:

  1. 模型置信度:语言模型对候选令牌 $v$ 的概率预测。
  2. 退化惩罚:衡量候选令牌 $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)

研究基础

实现基于两篇主要的研究论文:

  1. A Contrastive Framework for Neural Text Generation (NeurIPS 2022),提出了最初的框架。
  2. Contrastive Search Is What You Need For Neural Text Generation (2022),展示了该方法在 16 种不同语言的现成模型上的有效性。

Sources