First Token Cutoff LLM sampling

Salvatore Sanfilippo

首個 Token 截斷式 LLM 取樣

原文由 Salvatore Sanfilippo 發布,訂閱此部落格

從理論上來說,LLM 能給出的最佳回覆,就是每次都挑選機率最高的那個 token。這種做法會讓 LLM 的輸出變得完全確定(deterministic),而這對許多應用來說並不是好事。也因此,為了在保持貼合上下文的同時兼顧 LLM 的創造力,近年來人們提出了各種不同的取樣演算法。

如今最常被使用、幾乎可說是預設值的一種,叫做 top-p:它是一種核取樣(nucleus sampling),作法是把機率最高的 token 依序累加,直到總機率達到 p 為止,然後在這些 token 中進行加權隨機取樣。

在這篇部落格文章中,我想談談為什麼我認為核取樣未必是最佳做法,並提出一個簡單、易於理解的替代方案來避開核取樣的問題。這個演算法目前仍在開發階段,但我現在就把它發表出來,希望能激起一些討論與動手嘗試。

logits 裡藏著黃金

儘管 LLM 的 logits 是少數幾個完全可以被理解的內部運作環節,我卻很少看到有人對研究它們的特性、探索更進階的取樣方法、偵測並向使用者提示不確定性與可能的幻覺感興趣。把連續 token 的機率分布視覺化,是一個簡單又實用的練習,能幫助我們獲得一些洞見:

在上圖中可以看到機率最高的 32 個候選 token,依機率以顏色標示(白色 = 0,藍色 = 1),以及被選中的 token 及其排名(機率最高者為 0,前一個為 1,依此類推)。

在上面的例子中,Mistral 基礎模型知道 Umberto Eco 的出生與逝世日期,因此它很有信心地把絕大多數的總機率都集中在最有可能的那個 token 上。而在其他時候,模型則會比較困惑,原因可能是接續文本有多種表達方式,或是它對某些事實並不確定。詢問一個在訓練過程中未曾學過生日的姓名日期時,就會產生截然不同的 token 機率分布。

然而,一般來說我們想避免的是選到次佳的 token,以免讓 LLM 的生成偏離最低困惑度(perplexity)的路徑,進而產生幻覺。核取樣在這點上做得不好,因為累加 token 直到達到 p 的做法,取決於分布的形狀,可能會把遠遠弱於首選的 token 也納入。舉例來說,看看上圖中的「writer」這個 token。Umberto Eco 的確是作家,與 writer 對應的 token 雖然分數不算特別高,但仍遠遠高於第二順位的選擇。然而,若 p 設為 0.3,就有可能把第二個僅有 0.01 這種低分的 token 也累加進來,並以幾個百分比的機率抽中它,進而有讓生成走上錯誤路徑的風險。

相反地,在下一個 token「who」的情況中,則有多個分數相對接近的選項,這些都是可以用來產生不同生成結果的候選。

利用雪崩效應

這裡的一個關鍵觀察是:為了追求多樣性而去挑選過弱的 token 並不划算:即使我們只在有少數幾個不錯候選者的生成時機來讓輸出多樣化,雪崩效應也會幫上忙:一旦輸入的上下文改變,LLM 的輸出就會受到擾動,進而更有可能產生另一種版本的文字。

First Token Cutoff(FTC)演算法

免責聲明:此處描述的演算法尚未經過任何科學驗證。我花了幾天時間實驗了幾種具有「最差 token 有界」特性的取樣演算法,而這個看起來在適用性、效果與可理解性之間取得了最佳平衡。

這裡提出的演算法可以用比較不正式的方式表述如下:

  • 當 LLM 強烈偏向某個候選時,就選它。
  • 當存在多個可行的候選時,就產生不同的選擇。
  • 最差可能 token 的選取應該被限制在一定的範圍內。

過去的研究,例如 Tail Free Sampling,也指出應該在一小組由 LLM 產生的高品質 token 中進行選擇。不過在 TFS 中,這樣的集合是透過計算微分來找出對應的群集,也就是那些在 token 品質急遽下降的陡峭曲線之前的 token。

而在這裡提出的演算法中,我們希望取樣的選擇能遵循一個相對於最高分 token T0 更有界、也更易於理解的截斷標準,並賦予 T0 相對於其他所有 token 的特殊意義:也就是 LLM 在產生該 token 時的確定性程度(基本上可視為困惑度的一個代理指標)。因此,這個演算法會拒絕所有相較於 T0 差於某個百分比的 token。

這個截斷百分比的取值範圍是 0 到 1,我們稱之為「co」。一個可行且具代表性的 co 值例如是 0.5。

演算法的運作方式如下:

  1. 對 logits 計算 softmax()。
  2. 依機率對 token 進行排序。
  3. 給定最佳 token 的機率 T0,對所有其他 token 計算比值:r = 1 - (T[i] / T0)
  4. 只選取 r <= co 的 token
  5. 在被選中的 token 中進行加權隨機抽選。

請注意,透過這種方式,無論 token 的數值是否呈現平滑的單調遞減,都會有一個硬性上限,限制我們能納入可能集合的 token。相對地,其他試圖識別高分群集的方法則沒有這種限制。

實際範例

核取樣 top-p 在實務上看起來沒有災難性失效的原因之一,是很多時候第一個 token 的機率非常高,也就是困惑度很低的時候,而且很多時候我們碰巧也不會收集並選到低品質的 token,因此生成仍會沿著合理的路徑繼續。

當出現像下面這樣連續 token 的機率時,問題就比較大了:

0.25, 0.14, 0.01

若 p=0.4,我們就有可能把第三個低品質的 token 也收集進來,並給予它約 3% 的機率。

現在來看看 co 值為 0.5 的 First Token Cutoff(意即 token 最多只能比第一個差 50%):

第二個 token 的 r 值為:

r[t1] = 1-(0.14/0.25) = 0.44 # 0.44 <= 0.5,此 token 被接受

r[t2] = 1-(0.01/0.25) = 0.96 # 0.96 > 0.5,此 token 被拒絕

輸出範例

以下是在提示詞「Sorted sets are」下,使用 co=0.7 的 Mistral 基礎模型(非 instruct 版)的輸出。以下三個都是成功的輸出,並未為了品質而刻意挑選。

  1. 有序集合是 Redis 中一種強大的資料結構。它們可以用來儲存有序列表、儲存唯一值、儲存用於排名的分數,以及儲存由有序集合組成的有序列表。
  2. 有序集合是一種非常強大的資料結構。它讓你能以一種易於找出集合中最大值或最小值的方式來儲存資料,也讓你能對資料進行排序。這在許多不同的任務中都很有用,例如依分數為使用者排名,或是在資料庫中找出最受歡迎的項目。
  3. 有序集合是一種非常強大的資料結構,可用來解決許多不同的問題。最常見的使用情境是儲存一份唯一元素的清單,其中每個元素都有一個關聯的數值。舉例來說,你可以用一個有序集合來儲存家中所有人的姓名,並以他們的年齡作為關聯數值。

為什麼可理解性很重要

取樣參數是 LLM 終端使用者或 API 使用者少數必須調整的東西之一。往往需要透過反覆嘗試來調校,然而,若能有一個可直接對應到現實描述、並能直覺理解的單一可調參數,會有很大的幫助。此外,「co」是一個線性參數,因此特別容易推論,相較之下,像 temperature 這類參數,或甚至 top_p 的「p」雖然也是線性的,卻強烈依賴於 logits 的分布形狀。

未來的工作

我的研究才剛起步,接下來會更深入地研究並評估這個演算法。比起其他,更重要的是,我非常希望看到大家對取樣演算法有更多的關注,並有更多興趣去超越類似 top-p 的做法向前邁進。

從 logits 的分布中或許還能收集到一些有趣的資訊。舉例來說,線性探針(linear probe)很可能學會判斷 LLM 的隱藏層何時正在處理某些事實性資訊。這一點若結合 token 的困惑度,就可以用來向 LLM 使用者提示輸出的某個部分很可能是錯的。一般來說,將 token 的機率分布視覺化是非常有資訊量的,在某種程度上也能讓人親手觸摸到 LLM 是如何運作、以及在每一步有哪些候選。

參考實作

logits = mx.softmax(logits)
np_logits = np.array(logits) # MX -> NumPy
np_logits = np_logits.flatten()
sorted_indices = np.argsort(np_logits)
sorted_indices = sorted_indices[::-1]

co = 0.7
j = 1
t0 = np_logits[sorted_indices[0]]
while 1 - (np_logits[sorted_indices[j]] / t0) < co and j < len(np_logits):
    j += 1
accepted_logits = []
for i in range(0,j):
    accepted_logits.append(float(np_logits[sorted_indices[j]]))
accepted_logits = mx.array(accepted_logits)

idx = mx.random.categorical(accepted_logits)
idx = int(np.array(idx)) # Convert zero-dim array to scalar
token_id = sorted_indices[idx]

致謝

這些實驗之所以能如此輕鬆地完成,要感謝來自 Apple 的 MLX 函式庫,以及那些持續不懈、又酷又聰明的開發者們。MLX 非常容易上手,本就該如此:畢竟 LLM 本身是難以捉摸的,但推論本身卻是一個簡單的過程。

本文章由 muse-spark-1.2-contributor 進行翻譯

留言