模組 09 · 第 6 課

生成文字和 KV 快取

訓練好的模型只會給下一個字打分,怎麼從分數變成一首詩?這一課用自己的 GPT 試溫度、top-k、續寫和藏頭詩,再實現 KV 快取,量一量它省下了多少重複的計算。

  • 約 40 分鐘
  • 難度:深入
  • 實測:2026-09-15 torch 2.14,Apple M4 CPU,固定隨機種子

程式碼和執行結果保留原樣(簡體中文),與實際執行時完全一致。

第 01 模組第 3 課用 DeepSeek 做過溫度實驗,那時模型是一個黑箱。現在模型是我們自己訓練的,可以開啟看看它每一步到底在做什麼。

python generate.py

模型給出的只是機率

模型讀完一段文字,在最後一個位置給詞表裡的 6289 個字各打一個分,softmax 之後變成機率:

== 1. 读完一段文字,模型给下一个字的概率(前 8 名)
  「白日依山尽,黄河入海」→ 无 6.9%  间 4.5%  多 4.3%  深 3.7%  边 3.0%  空 2.6%  难 2.3%  遥 2.1%
  「床前明月光,疑是地上」→ 天 3.2%  花 2.9%  楼 2.4%  清 1.8%  看 1.7%  游 1.7%  人 1.6%  枝 1.3%
  「春风」→ 。 14.2%  , 5.1%  吹 3.3%  满 2.9%  雨 1.9%  起 1.9%  尽 1.4%  落 1.3%

"黃河入海"後面,原詩是"流",可"流"不在前 8 名裡。"疑是地上"後面,"霜"也不在。上一課檢查過,這個模型不背誦,它沒有記住這兩首名詩,只是按學到的規律給出"合理"的字。

順便說一個有意思的發現:《全唐詩》裡收錄的《靜夜思》是"床前看月光,疑是地上霜。舉頭望山月,低頭思故鄉",和課本上的"床前明月光……舉頭望明月"不一樣。一般認為,我們熟悉的版本出自明清時期的選本。所以就算模型背了詩,它背的也是"看月光"。

"春風"後面,機率最高的是句號(14.2%)。我數了一下訓練資料:"春風"一共出現 797 次,其中 111 次後面緊跟著句號,佔 13.9%,和模型給出的機率幾乎一樣。模型學到的,就是這些字在唐詩裡的出現規律。

生成文字,就是反覆做同一件事:拿到機率,選一個字,接到後面,再算下一個字的機率。怎麼"選",決定了寫出來的是什麼。

溫度

第 01 模組講過溫度:把分數除以溫度,再 softmax。溫度小於 1,高分和低分的差距被拉大,模型更"保守";大於 1,差距被縮小,模型更"大膽"。

logits = logits[:, -1, :] / temperature
...
next_id = torch.multinomial(F.softmax(logits, dim=-1), num_samples=1)

同樣的隨機種子,三個溫度各寫 3 首:

  温度 0.3:
    白云无限意,白发不如何。何事无人在,无人不可知。
    一望东山路,千行万里愁。夜深人不见,秋色月无声。
    江南春水上,江上月明时。独有清风起,何人别有期。
  温度 1.0:
    白露起春光,知君得舞酲。霜时飏先没,桐死势悠扬。
    白发摩牛放,真关万趾兴。夜寒休绕郡,秋杀杜娘宫。玉镜收新镜,金尊坐晚麕。何言吊鱼道,双锡空余年。
    戍楼映江岸,淼危棹度斜。欲分凝炯娩,归盖乱精冰。
  温度 1.5:
    淡露依春似映兹,襞舞凋年雪路揉筝衬情。死势悠扬双觉处,摩牛放将溉噭胞。……

溫度 0.3 寫得最通順,"一望東山路,千行萬里愁。夜深人不見,秋色月無聲"幾乎可以亂真。代價是用字很常見,重複也多:"無人"出現了兩次,"江"出現了兩次,讀多了會覺得都差不多。

溫度 1.0 就是按模型給出的原始機率選字。用字豐富了,但開始出現"麕""酲""炯娩"這些生僻字,意思也更難接上。

溫度 1.5 徹底亂了:格式崩了("襞舞凋年雪路揉箏襯情"一句九個字),滿篇生僻字。原因是溫度高時,機率分佈被壓平,6000 多個字裡那些本來機率極低的字,加起來也分走了可觀的機率。一旦選中一個奇怪的字,後面的上下文就變得奇怪,錯誤越積越多。

top-k:只在靠前的字裡選

溫度高時的問題是"長尾":幾千個低機率的字加起來佔了太多份額。top-k 的辦法很直接:只保留機率最高的 k 個字,其餘的全部去掉,在這 k 個裡按機率選。

if top_k is not None:
    kth = torch.topk(logits, top_k).values[:, -1:]
    logits = logits.masked_fill(logits < kth, float("-inf"))

溫度同樣是 1.5,加上 top-k=20:

== 3. 温度 1.5,但只从概率最高的 20 个字里选(top-k=20)
    白头白马在西陵,自是西林下路尘。今日不辞青史士,可堪回首白苹州。
    万国风尘里,千金雨色深。不知归未到,应似到人中。
    东北山川尽不愁,三千春尽一枝愁。春寒未入千门色,雨尽空闻九夜寒。……

格式恢復了,生僻字也不見了,但比溫度 0.3 更有變化。高溫度負責"大膽",top-k 負責"別太離譜",兩個一起用效果最好。

另一個常用的辦法叫 top-p(核取樣):不固定個數,而是從最高的開始往下加,加到累計機率達到 p(比如 0.9)為止。機率集中時只剩幾個字,分散時多留一些。第 01 模組用 API 時見過這個參數。

續寫和藏頭詩

生成的開頭不一定是空的。給出前半句,模型就接著寫:

== 4. 给出开头,让它续写
    床前明月光,已有四邻情。知君得舞袖,何奈在山期。
    大漠孤烟直,高楼楚水清。雁雕摩露下,鸿劒鼓歌中。夜雨初分郡,秋风又过巴。此才应不遇,空与钓鱼舟。
    人工智能好,言蠹不见人。徒思一岁月,独在岘生春。

"大漠孤煙直"續成了一首完整的五言律詩,八句。"人工智慧"這個唐朝沒有的詞,它也照樣接得下去,因為這四個字各自都在唐詩裡出現過。

更進一步,可以在生成過程中插手。藏頭詩要求每句的第一個字是指定的:先把指定的字放進去,再讓模型一個字一個字地寫完這一句,寫的時候不許它輸出標點和換行,最後由我們加上標點:

@torch.no_grad()
def acrostic(heads, n_char=7, temperature=0.8):
    ids = encode("\n")
    for i, head in enumerate(heads):
        ids = torch.cat([ids, encode(head)], dim=1)  # 先把指定的字放进去
        for _ in range(n_char - 1):  # 再让模型一个字一个字地写完这一句
            logits = model(ids)[0][0, -1] / temperature
            logits[BANNED] = float("-inf")
            next_id = torch.multinomial(F.softmax(logits, dim=-1), 1)
            ids = torch.cat([ids, next_id[None]], dim=1)
        ids = torch.cat([ids, encode(",。"[i % 2])], dim=1)  # 标点由我们加,保证格式
    return decode(ids[0].tolist()).strip()
== 5. 藏头诗:每句开头的字由我们指定,其余由模型写
    春色欲依春似春,眠花先得舞花年。不知诗句先酬情,觉后还随楚水人。
    学得重雕出,无情忘却回。止行空向兴,境静更闻闻。

每句的開頭連起來是"春眠不覺"和"學無止境"。

這個技巧有一個正式的名字:約束解碼。在選字之前,把不符合要求的選項的分數設成負無窮。第 02 模組第 4 課讓 DeepSeek 輸出 JSON 時,服務端保證輸出格式正確,用的就是同一個原理:在每一步,把會讓 JSON 不合法的詞元都排除掉。

重複的計算

現在看生成的效率。回到 generate 方法,最簡單的寫法是每一步都把整段文字送進模型:

logits, _ = self(idx)  # 不用缓存:每一步都把整段序列从头算一遍

寫第 50 個字時,模型要處理前面全部 49 個字;寫第 51 個字時,又要把這 49 個字重新處理一遍,再加上第 50 個字。

但仔細想想,前 49 個字的計算結果並沒有變。因果掩碼保證每個位置只看前面,後面加了新字,前面位置的計算不受任何影響。真正新的,只有最後一個位置。

最後一個位置要算注意力,需要什麼?它自己的 q,以及前面所有位置的 k 和 v。所以只要把每一層的 k 和 v 存起來,每一步只算新字的 q、k、v,把新的 k、v 接到存好的後面,就夠了。這就是 KV 快取

實現

gpt.py 的注意力裡,快取只多了幾行:

if cache is not None:  # KV 缓存:把之前算过的 K、V 接在前面,不用重算
    if "k" in cache:
        k = torch.cat([cache["k"], k], dim=2)
        v = torch.cat([cache["v"], v], dim=2)
    cache["k"], cache["v"] = k, v

有一個細節要注意:因果掩碼。不用快取時,q 和 k 的長度一樣,掩碼是一個正方形的下三角。用快取時,q 只有新來的 1 個字,k 卻有全部的字,而且新字排在最後,它可以看到所有的 k。所以掩碼要錯開:

total = k.size(2)
mask = torch.ones(T, total, dtype=torch.bool, device=x.device).tril(diagonal=total - T)

位置嵌入也要跟著調整:新字不是第 0 個位置,而是第 start 個。這就是 forwardstart 參數的作用。

生成時,第一步把整段提示詞送進去,把所有位置的 k、v 存進快取;之後每一步只送最新的一個字:

if use_cache:
    # 第一步把整段提示词送进去;之后每步只送最新的一个词元,前面的 K、V 从缓存里取
    logits, _ = self(idx[:, start:], caches=caches, start=start)
    start = idx.size(1)

量一量

== 6. KV 缓存:结果一样吗?快了多少?
  同样的随机种子,写 120 个字:用缓存和不用缓存的结果完全一样
  写  32 个字:不用缓存     23 毫秒,用缓存    12 毫秒,快了 1.8 倍
  写  64 个字:不用缓存     52 毫秒,用缓存    26 毫秒,快了 2.0 倍
  写 120 个字:不用缓存    130 毫秒,用缓存    49 毫秒,快了 2.6 倍
  这个模型每个词元的缓存:2 × 4 层 × 128 维 × 4 字节 = 4 KB,128 个词元共 512 KB

先確認正確性:同樣的隨機種子,用不用快取,寫出的 120 個字完全一樣。快取只是省掉了重複計算,結果不變。

再看速度。寫的字越多,省下的越多:32 個字快 1.8 倍,120 個字快 2.6 倍。不用快取時,寫第 n 個字要處理 n 個字,總的計算量隨長度按平方增長;用快取後,每一步只處理 1 個字。

在我們這個模型上,加速只有兩倍多,並不驚人。原因是模型太小、文字太短:每一步真正的計算只要零點幾毫秒,Python 本身呼叫函式、拼接張量的開銷佔了很大比例,這部分開銷快取省不掉。在大模型上,每一步的計算量大得多,文字動輒幾千上萬個詞元,KV 快取省下的就是數量級的差距。

快取的代價:記憶體

KV 快取用記憶體換時間。每個詞元,每一層都要存一個 k 和一個 v。我們的模型每個詞元只要 4 KB,128 個詞元 512 KB,微不足道。

大模型就不一樣了。一個幾十層、每層幾千維的模型,每個詞元的快取可能要幾百 KB 到 1 MB 以上。一個十幾萬詞元的長對話,快取就要佔到幾十 GB 的視訊記憶體,常常比模型本身還大。同時服務很多使用者時,這就是最主要的視訊記憶體開銷。

所以上一課末尾提到的那些改進,很多都是衝著 KV 快取來的:多個頭共用 k 和 v(分組查詢注意力),或者把 k 和 v 壓縮成更小的向量再存(DeepSeek 的多頭潛在注意力,MLA)。第 10 模組講 vLLM 時,還會看到服務端怎麼管理這些快取。

第 06 模組第 4 課講的"快取命中的輸入更便宜",是另一層的快取:服務商把常用字首(比如固定的系統提示詞)的 KV 快取儲存下來,下一個請求用到相同的字首時,直接取出來,不用重新計算,所以能便宜很多。

練習

  1. 實現 top-p 取樣:把機率從大到小排序,保留累計機率達到 p 的那些字。用溫度 1.5、p=0.9 寫幾首,和 top-k=20 比較。
  2. 寫一個"固定格式"的約束解碼:只允許生成七言絕句,第 8、16、24 個字元只能是逗號或句號,其他位置不許出現標點。
  3. gpt.pyblock_size 改成 512 並重新訓練(或者只是隨機初始化),測一下寫 500 個字時,用快取和不用快取各要多久。

自測

1. 溫度和 top-k 分別改變了什麼?為什麼常常一起用?

溫度在 softmax 之前把分數整體縮放,改變機率分佈的"尖銳"程度:溫度低時集中在少數高分的字上,溫度高時更平均。top-k 只保留機率最高的 k 個字,把其餘的去掉。溫度高能帶來多樣性,但會讓大量低機率的字分走份額,top-k 把這些長尾去掉,兩者一起用既有變化又不至於太離譜。

2. KV 快取為什麼不會改變生成的結果?

因為有因果掩碼,每個位置的計算只依賴它自己和前面的位置,後面加上新字不會改變前面位置的 k 和 v。快取只是把這些不會變的結果存起來重複使用,省掉了重複計算,數學上完全等價。

3. KV 快取的代價是什麼?為什麼在大模型上這個代價很重要?

代價是記憶體:每個詞元在每一層都要存一份 k 和 v。大模型層數多、維度大,每個詞元的快取很大,長上下文和多使用者同時使用時,快取佔的視訊記憶體可能超過模型本身,所以很多模型結構上的改進(如分組查詢注意力、多頭潛在注意力)都是為了減小它。

提問與討論

這一課沒看懂的地方,在這裡問。看到別人的問題,也歡迎你來回答。

提問 +3 點,回答別人 +6 點。內容經審核後公開。

正在載入討論…