一字一诗
LESSON 06 / 26

给猜测打分

这一课完成什么

模型会先给每个候选字一个分数,再将这些分数变成概率,最后根据正确答案算出一个误差。代码中原始分数叫 logits,误差叫 loss。这节让两者都有可手算的含义。

分数不是概率,先把它们变成比例

模型对每个候选字给一个分数。先做指数变换,再除以总和,得到加起来为 1 的比例,这一步叫 softmax。

分数不是概率,先把它们变成比例
用文字逐步读这张图
  1. 候选字:月 / 山 / 水;正确答案是“月”
  2. 原始分数:2 / 1 / 0;还不是百分比
  3. 变成正数:e² / e¹ / e⁰;约 7.39 / 2.72 / 1
  4. 除以总和:每项 ÷ 11.11;全部加起来为 1
  5. 得到概率:66.5% / 24.5% / 9%;“月”最有可能
  6. 计算误差:−ln(0.665);约 0.408,越小越好
把正确答案的概率提高,误差怎样变化?

会降低。负对数把接近 1 的正确概率变为接近 0 的误差。

从分数到概率

候选字为“月、山、水”,原始分数是 [2,1,0]。这些分数可以是负数,也不要求和为 1。softmax 的做法是先取指数,再除以指数之和。为了避免数字过大,实际计算通常先把所有分数减去最大值,最终概率不变。

text
减去最大值: [0, -1, -2]
取指数:     [1, 0.3679, 0.1353]
除以总和:   [0.6652, 0.2447, 0.0900]

三个概率均大于 0,总和为 1。模型认为“月”最可能,但不代表它一定正确,也不代表分数已经等同于人的文学评价。

给一次预测打分

若正确答案为“月”,只需看它被分到的概率 0.6652。误差定义为 -ln(0.6652),约为 0.4076。ln 是自然对数。当正确字概率趋近 1,误差趋近 0;正确字概率很小,误差很大。

正确字概率误差约为直观意义
0.90.105正确字获得较高信心
0.50.693正确字只占一半概率
0.12.303正确字被低估

对许多字位置重复此计算并平均,就是这里使用的交叉熵误差。它评估“预测原文下一个字”的能力,不能直接等同于一首新诗的好坏。

用代码验证手算

bash
./poet lesson 6
python
logits = torch.tensor([[2., 1., 0.]])
target = torch.tensor([0])
loss = torch.nn.functional.cross_entropy(logits, target)

目标 0 表示正确字在三个候选中的第一个位置。这个函数接收原始分数,内部会完成稳定的概率计算,不要先 softmax 后再传进去。为了展示,你可以另算 logits.softmax(-1) 并与手算结果对照。

评分范围必须明确

正式训练只给正文和结束标记评分。已知的关键词前缀、补齐到统一长度的空位置设为 -100,不纳入分母。若大量空位置也参与平均,一个只会预测空白的程序也可能看起来成绩不错。验证函数会按实际有效字数加权,避免长度不同的批次产生偏差。

验收与排错

小练习与答案

如果给三个分数同时加上 100,概率会改变吗?如果只给“山”加上 100 呢?

查看答案

同时加同一个常数不会改变 softmax,因为指数中的公共因子会在分子分母约掉。只给“山”大幅加分会使概率集中到“山”;如果正确字仍是“月”,误差会明显增加。可以分别修改 logits 验证。

动手补全一小段

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

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

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

展开本课补全练习的完整参考答案
python
"""第 06 课补全练习:给正确字评分。修改 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, targets):
    return F.cross_entropy(logits, targets)

assert abs(solve(torch.tensor([[2.,1.,0.]]),torch.tensor([0])).item()-0.40760595)<1e-6
print("本课补全练习通过。")
展开本课完整、可独立运行的实验代码
python
"""本课独立实验;在仓库根目录执行 .venv/bin/python lessons/06/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]])
p = logits.softmax(-1)
show('候选 月 山 水', p.tolist())
show('正确字为 月 的误差', F.cross_entropy(logits, torch.tensor([0])).item())
show('手算 -ln(正确字概率)', -math.log(p[0, 0].item()))
下载本课实验