一字一诗
LESSON 20 / 26

从概率中选择下一个字

这一课完成什么

理解模型如何把一张候选分数表变成实际文字。生成时不再计算“正确答案误差”或更新参数,而是反复取最后位置的分数,选一个字,把这个字加回输入,再计算下一次。

生成时,把分数变成一次选择

每生成一个字,都把它接到前文末尾,再重新计算下一字。温度改变分布的尖锐程度,top-k 限制参与选择的候选数量。

生成时,把分数变成一次选择
用文字逐步读这张图
  1. 给定前文:例如“春江”;先转为编号
  2. 模型打分:所有候选汉字;只取最后位置输出
  3. 调整变化程度:分数 ÷ temperature;小则更偏向高分字
  4. 保留候选:选分数最高的 k 个;top-k
  5. 按比例抽一个:并非总取最高分;随机种子控制抽样
  6. 接回前文:春江 → 春江花;继续算下一字
设定随机种子,是否跨所有设备逐字相同?

不保证。设备和计算实现的差异也会影响结果,应在同一环境比较。

一次生成的循环

text
条件 + 春  → 对所有候选字打分 → 选“风”
条件 + 春风 → 再打分          → 选“吹”
条件 + 春风吹 → 再打分        → ……

上面只说明循环,不是一次模型实测输出。每次生成的选择都会成为下一次的前文,因此早期选到不自然的字,也可能影响后面几句。模型并没有先写好整首诗再一次吐出来。

选择最可能的字,还是保留变化

每次都选最大分数叫贪心选择,通常很稳定,也可能陷入重复。采样则按照概率抽选,低概率字也有机会出现。为了避免从非常低分的生僻字里乱选,项目先保留分数最高的 K 个字,再在这些字之间采样。

top_k=40 表示只保留当前最高分的 40 个允许汉字。它不是只用字表前 40 个字,也不是整首诗只能使用 40 种字。每个位置的候选名单会重新计算。

temperature 控制分布的尖锐程度

把原始分数除以一个正数,再计算 softmax。这个数叫 temperature,本课程称“变化程度”:小于 1 往往使高分项更集中,大于 1 则让分布更平。它不直接衡量“创造力”,值过高常会让句子变得不自然。

bash
./poet lesson 20
./poet write --start 春 --temperature 0.6 --seed 42
./poet write --start 春 --temperature 1.1 --seed 42

对照实验保持开头、候选数、模型文件和随机种子一致,只改变一个设置。固定种子可以帮助在同一环境重现采样,但不同软件版本、设备和计算路径可能不保证逐字完全相同。

与注意力中的概率区别

注意力权重分配的是“前文各位置贡献多少”;本课的概率分配的是“下一个字选谁”。它们都可能使用 softmax,却在不同位置处理不同对象。正式模型中注意力缩放使用每头宽度的平方根,CLI 的 temperature 只调整最后选字的分布。

验收与排错

小练习与答案

temperature 从 1 改为 0.5,候选排序会变化吗?概率会变化吗?

查看答案

正数缩放不会改变原始分数排序,但会改变相对差距经过指数后的比例,所以概率会更集中在高分项上。若使用相同的 top-k,保留的候选集合不变,集合内部的采样概率改变。

动手补全一小段

先运行上面的完整实验,再复制本课起始文件为自己的练习。starter 有意留空 solve 函数;补全后运行文件,底部检查会告诉你是否符合本课要求。卡住时打开参考答案,比较每一步。

bash
cp lessons/20/starter.py lessons/20/my_exercise.py
# 编辑 my_exercise.py 中的 solve 函数,然后运行:
.venv/bin/python lessons/20/my_exercise.py
# 对照完整答案:
.venv/bin/python lessons/20/solution.py

下载起始代码 · 下载参考答案

展开本课补全练习的完整参考答案
python
"""第 20 课补全练习:调节候选分布。修改 solve,保持下方检查不变。"""
import math, io, argparse
import torch
from torch import nn
from torch.nn import functional as F
torch.set_num_threads(2)
torch.manual_seed(26)

def solve(logits, temperature):
    if temperature <= 0: raise ValueError('必须大于0')
    return (logits/temperature).softmax(-1)

x=torch.tensor([2.,1.,0.])
a=solve(x,.5);b=solve(x,1.)
assert a[0]>b[0]
assert abs(a.sum().item()-1)<1e-6
print("本课补全练习通过。")
展开本课完整、可独立运行的实验代码
python
"""本课独立实验;在仓库根目录执行 .venv/bin/python lessons/20/experiment.py。"""
import sys, json, math
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
import torch
from torch import nn
from torch.nn import functional as F
from poetry_gpt.common import ROOT, DEFAULT_DATA, read_json, read_jsonl
from poetry_gpt.model import ModelConfig, PoetryGPT, CausalAttention, Block
from poetry_gpt.data import Tokenizer, SPECIAL, clean_record, keywords_for
from poetry_gpt.labs import show
torch.set_num_threads(2)
torch.manual_seed(26)

logits = torch.tensor([2.0, 1.0, 0.0])
for temperature in [0.5, 0.8, 1.2]:
    show(f'变化程度 {temperature}', (logits / temperature).softmax(-1).tolist())
下载本课实验