案例 12:《爱丽丝梦游仙境》字符语言模型:RNN、LSTM 与 GRU#

学习目标与记号#

  本案例以字符级自回归语言模型为主线,说明循环神经网络怎样利用前文预测下一个字符。主要讨论以下内容:

  1. 从原始文本建立训练集词表,并把序列变成 (批量, 时间) 张量;

  2. 解释“用前文预测下一个字符”的因果信息边界;

  3. 理解隐藏状态的递归更新、不同时间步之间的参数共享与随时间后向传播(backpropagation through time,BPTT);

  4. 比较普通 RNN、LSTM 和 GRU 在相同预算下的困惑度;

  5. 使用梯度裁剪并监控裁剪前梯度范数;

  6. 区分生成样例、验证困惑度和测试困惑度的职责。

  记第 \(t\) 个字符和对应隐藏状态分别为 \(x_t\)\(\boldsymbol{h}_t\)序列长度为 \(T\)词表大小为 \(V\)模型只学习文学文本中的字符变化规律,不具备事实问答能力。

数据、版权与缓存#

  • 文本:Lewis Carroll,Alice's Adventures in Wonderland

  • 说明页:https://www.gutenberg.org/ebooks/11;

  • 匿名 HTTPS 下载地址:https://www.gutenberg.org/cache/epub/11/pg11.txt;

  • 权利说明:Project Gutenberg 标注该作品在美国属于公版;美国以外的使用者需要自行核验当地法律。

  文本文件较小,但程序仍遵循统一缓存规则。我们按照原文顺序,把前 \(80\%\)随后 \(10\%\) 和最后 \(10\%\) 分别作为训练集、验证集和测试集;滑动窗口不会跨越不同数据集之间的边界。词表只根据训练文本建立,训练文本中没有出现过的字符映射为 <unk>

from __future__ import annotations

import math
import os
import random
import re
import urllib.request
from collections import Counter
from pathlib import Path

import numpy as np
import pandas as pd
import matplotlib.pyplot as plt

SEED = 20260812
FAST_MODE = os.getenv("FAST_MODE", "1") != "0"
random.seed(SEED)
np.random.seed(SEED)

try:
    import torch
    from torch import nn
    from torch.utils.data import DataLoader, Dataset
except ImportError as exc:
    raise ImportError("本案例需要 PyTorch。") from exc

torch.manual_seed(SEED)
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
CACHE_ROOT = Path(os.getenv("AI_COURSE_DATA_DIR", Path.home() / ".cache" / "ai-course"))
CACHE_DIR = CACHE_ROOT / "gutenberg_alice"
CACHE_DIR.mkdir(parents=True, exist_ok=True)
print({"快速模式": FAST_MODE, "计算设备": str(DEVICE), "缓存目录": str(CACHE_DIR)})
{'快速模式': True, '计算设备': 'cpu', '缓存目录': '/private/tmp/ai-course-case-data/gutenberg_alice'}
URL = "https://www.gutenberg.org/cache/epub/11/pg11.txt"
TEXT_PATH = CACHE_DIR / "pg11.txt"
if not TEXT_PATH.exists():
    urllib.request.urlretrieve(URL, TEXT_PATH)
assert TEXT_PATH.stat().st_size > 100_000

raw_text = TEXT_PATH.read_text(encoding="utf-8-sig")
start_marker = "*** START OF THE PROJECT GUTENBERG EBOOK"
end_marker = "*** END OF THE PROJECT GUTENBERG EBOOK"
start = raw_text.find(start_marker)
end = raw_text.find(end_marker)
assert start >= 0 and end > start
body = raw_text[raw_text.find("\n", start) + 1:end]
body = body.replace("\r\n", "\n").replace("\r", "\n")
body = re.sub(r"\n{3,}", "\n\n", body)
assert len(body) > 100_000
print(f"清理后字符数:{len(body):,}")
清理后字符数:144,532
counts = Counter(body)
print("不同字符数:", len(counts))
print("最常见字符:", counts.most_common(15))
print("换行比例:", counts["\n"] / len(body))

lengths = [len(line) for line in body.splitlines() if line]
line_summary = pd.Series(lengths).describe(percentiles=[0.5, 0.9, 0.99])
line_summary = line_summary.rename({
    "count": "行数", "mean": "均值", "std": "标准差", "min": "最小值",
    "50%": "中位数", "90%": "百分之九十分位数", "99%": "百分之九十九分位数", "max": "最大值",
})
line_summary.index.name = "统计量"
display(line_summary.to_frame("每行字符数"))
不同字符数: 75
最常见字符: [(' ', 24617), ('e', 13552), ('t', 10345), ('a', 8239), ('o', 8073), ('h', 7176), ('n', 6955), ('i', 6856), ('s', 6361), ('r', 5376), ('d', 4779), ('l', 4681), ('u', 3461), ('\n', 3312), ('w', 2468)]
换行比例: 0.022915340547422024
每行字符数
统计量
行数 2494.000000
均值 56.623897
标准差 18.761633
最小值 4.000000
中位数 67.000000
百分之九十分位数 71.000000
百分之九十九分位数 71.000000
最大值 71.000000
n = len(body)
train_text = body[:int(0.80 * n)]
valid_text = body[int(0.80 * n):int(0.90 * n)]
test_text = body[int(0.90 * n):]

# 词表仅看训练段;特殊词元占固定位置。
train_chars = sorted(set(train_text))
itos = ["<unk>"] + train_chars
stoi = {ch: i for i, ch in enumerate(itos)}
UNK = stoi["<unk>"]

def encode(text):
    return np.asarray([stoi.get(ch, UNK) for ch in text], dtype=np.int64)

train_ids, valid_ids, test_ids = map(encode, (train_text, valid_text, test_text))
assert len(train_ids) + len(valid_ids) + len(test_ids) == len(body)
assert set(stoi) == set(itos)
print({
    "训练集字符数": len(train_ids), "验证集字符数": len(valid_ids),
    "测试集字符数": len(test_ids), "词表大小": len(itos),
    "验证集未知字符率": float(np.mean(valid_ids == UNK)),
    "测试集未知字符率": float(np.mean(test_ids == UNK)),
})
{'训练集字符数': 115625, '验证集字符数': 14453, '测试集字符数': 14454, '词表大小': 76, '验证集未知字符率': 0.0, '测试集未知字符率': 0.0}
SEQ_LEN = 64

class NextCharacterDataset(Dataset):
    def __init__(self, ids, seq_len, max_sequences=None):
        self.ids = np.asarray(ids, dtype=np.int64)
        self.seq_len = seq_len
        self.starts = np.arange(0, len(self.ids) - seq_len - 1, seq_len)
        if max_sequences is not None:
            self.starts = self.starts[:max_sequences]
        assert len(self.starts) > 0
    def __len__(self):
        return len(self.starts)
    def __getitem__(self, idx):
        s = int(self.starts[idx])
        x = torch.from_numpy(self.ids[s:s+self.seq_len])
        y = torch.from_numpy(self.ids[s+1:s+self.seq_len+1])
        return x, y

train_ds = NextCharacterDataset(train_ids, SEQ_LEN, 900 if FAST_MODE else None)
valid_ds = NextCharacterDataset(valid_ids, SEQ_LEN, 200 if FAST_MODE else None)
test_ds = NextCharacterDataset(test_ids, SEQ_LEN, 200 if FAST_MODE else None)
generator = torch.Generator().manual_seed(SEED)
train_loader = DataLoader(train_ds, batch_size=64, shuffle=True, generator=generator)
valid_loader = DataLoader(valid_ds, batch_size=128, shuffle=False)
test_loader = DataLoader(test_ds, batch_size=128, shuffle=False)

xb, yb = next(iter(train_loader))
assert xb.shape == yb.shape == (64, SEQ_LEN)
assert torch.equal(xb[:, 1:], yb[:, :-1])
print("一个批次输入/标签维度:", xb.shape, yb.shape)
一个批次输入/标签维度: torch.Size([64, 64]) torch.Size([64, 64])
# 零阶基线:每个位置都按训练字符边际分布预测,不利用上下文。
train_freq = np.bincount(train_ids, minlength=len(itos)).astype(np.float64)
unigram_prob = (train_freq + 1.0) / (train_freq.sum() + len(itos))

def unigram_nll(ids):
    return float(-np.log(unigram_prob[ids]).mean())

baseline_nll = unigram_nll(test_ids)
baseline_ppl = math.exp(baseline_nll)
print({"单字符基线测试负对数似然": baseline_nll, "单字符基线测试困惑度": baseline_ppl})
assert np.isfinite(baseline_ppl) and baseline_ppl > 1
{'单字符基线测试负对数似然': 3.1716585002752344, '单字符基线测试困惑度': 23.84700183660418}

自回归、循环状态与随时间后向传播#

  给定字符序列 \((x_1,\ldots,x_T)\)字符语言模型把联合概率写为

\[ p(x_1,\ldots,x_T)=\prod_{t=1}^{T}p(x_t\mid x_{<t}). \]

其中,\(x_{<t}=(x_1,\ldots,x_{t-1})\) 表示第 \(t\) 个字符之前的全部字符。普通循环神经网络(recurrent neural network,RNN)使用 \(\boldsymbol{h}_t=\tanh(\boldsymbol{W}_x\boldsymbol{x}_t+\boldsymbol{W}_h\boldsymbol{h}_{t-1}+\boldsymbol{b})\) 逐步更新隐藏状态,并在所有时间步共享同一组参数。长短期记忆网络(long short-term memory,LSTM)增加细胞状态、输入门、遗忘门和输出门;门控循环单元(gated recurrent unit,GRU)使用更新门和重置门简化门控结构。训练时将各时间位置的交叉熵汇总,梯度沿展开的时间计算过程向前面的时间步传播,这一过程称为随时间后向传播。

  本案例在一个长度为 64 的序列片段内部连续更新隐藏状态,但在不同片段之间重新初始化状态。这样可以使程序更简单,并保证三种模型使用相同的比较方法;相应地,模型最多直接利用 64 个字符的上下文。

class CharacterModel(nn.Module):
    def __init__(self, vocab_size, cell_type="rnn", embedding_dim=48, hidden_dim=96):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embedding_dim)
        cells = {"rnn": nn.RNN, "lstm": nn.LSTM, "gru": nn.GRU}
        if cell_type not in cells:
            raise ValueError(f"未知 cell_type={cell_type}")
        self.recurrent = cells[cell_type](
            embedding_dim, hidden_dim, batch_first=True
        )
        self.output = nn.Linear(hidden_dim, vocab_size)
        self.cell_type = cell_type
    def forward(self, token_ids, state=None):
        assert token_ids.ndim == 2
        embedded = self.embedding(token_ids)           # (B,T,E)
        hidden, state = self.recurrent(embedded, state) # (B,T,H)
        logits = self.output(hidden)                    # (B,T,V)
        assert logits.shape[:2] == token_ids.shape
        return logits, state

model_labels = {"rnn": "普通循环神经网络", "lstm": "长短期记忆网络", "gru": "门控循环单元网络"}
for kind in ("rnn", "lstm", "gru"):
    logits, state = CharacterModel(len(itos), kind)(torch.zeros(2, 7, dtype=torch.long))
    assert logits.shape == (2, 7, len(itos))
    print(model_labels[kind], "参数量:", sum(p.numel() for p in CharacterModel(len(itos), kind).parameters()))
普通循环神经网络 参数量: 25036
长短期记忆网络 参数量: 67084
门控循环单元网络 参数量: 53068
def evaluate_nll(model, loader):
    model.eval()
    total_loss = total_tokens = 0
    with torch.no_grad():
        for xb, yb in loader:
            xb, yb = xb.to(DEVICE), yb.to(DEVICE)
            logits, _ = model(xb)
            loss_sum = nn.functional.cross_entropy(
                logits.reshape(-1, len(itos)), yb.reshape(-1), reduction="sum"
            )
            total_loss += float(loss_sum)
            total_tokens += yb.numel()
    return total_loss / total_tokens

def train_model(kind, epochs):
    torch.manual_seed(SEED)
    model = CharacterModel(len(itos), kind).to(DEVICE)
    optimizer = torch.optim.Adam(model.parameters(), lr=2e-3)
    history, grad_norms = [], []
    for epoch in range(epochs):
        model.train()
        losses = []
        for xb, yb in train_loader:
            xb, yb = xb.to(DEVICE), yb.to(DEVICE)
            optimizer.zero_grad()
            logits, _ = model(xb)
            loss = nn.functional.cross_entropy(
                logits.reshape(-1, len(itos)), yb.reshape(-1)
            )
            assert torch.isfinite(loss)
            loss.backward()
            norm = nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
            assert torch.isfinite(norm)
            grad_norms.append(float(norm))
            optimizer.step()
            losses.append(loss.item())
        valid_nll = evaluate_nll(model, valid_loader)
        history.append((float(np.mean(losses)), valid_nll))
    return model, np.asarray(history), np.asarray(grad_norms)

epochs = 2 if FAST_MODE else 12
models, histories, gradient_norms = {}, {}, {}
for kind in ("rnn", "lstm", "gru"):
    models[kind], histories[kind], gradient_norms[kind] = train_model(kind, epochs)
    print(model_labels[kind], "最后一轮训练集/验证集平均负对数似然:", histories[kind][-1],
          "裁剪前梯度范数百分之九十五分位数:", np.percentile(gradient_norms[kind], 95))
普通循环神经网络 最后一轮训练集/验证集平均负对数似然: [2.9369016  2.83011871] 裁剪前梯度范数百分之九十五分位数: 1.0194228768348692
长短期记忆网络 最后一轮训练集/验证集平均负对数似然: [3.09380007 3.03211472] 裁剪前梯度范数百分之九十五分位数: 0.6272629290819168
门控循环单元网络 最后一轮训练集/验证集平均负对数似然: [2.98397778 2.89189171] 裁剪前梯度范数百分之九十五分位数: 0.7260805428028106
for kind, hist in histories.items():
    plt.plot(hist[:, 0], "--", label=f"{model_labels[kind]}-训练集")
    plt.plot(hist[:, 1], label=f"{model_labels[kind]}-验证集")
plt.xlabel("训练轮次")
plt.ylabel("平均词元负对数似然")
plt.title("词表、序列长度与更新次数相同,只改变网络类型进行比较")
plt.legend()
plt.tight_layout()
../../_images/aeac5325c5d06cef2a72376dde3e480045b01395673c9e795ed8e0ef3fb7ff6d.png
rows = [{"model": "unigram", "test_nll": baseline_nll, "test_perplexity": baseline_ppl}]
for kind, model in models.items():
    nll = evaluate_nll(model, test_loader)
    rows.append({
        "model": kind,
        "test_nll": nll,
        "test_perplexity": math.exp(min(nll, 20)),
        "parameters": sum(p.numel() for p in model.parameters()),
        "grad_norm_p95": np.percentile(gradient_norms[kind], 95),
    })
results = pd.DataFrame(rows).set_index("model")
result_labels = {"unigram": "单字符频率基线", **model_labels}
results_display = results.rename(
    index=result_labels,
    columns={"test_nll": "测试集平均负对数似然", "test_perplexity": "测试集困惑度",
             "parameters": "参数量", "grad_norm_p95": "裁剪前梯度范数百分之九十五分位数"},
)
results_display.index.name = "模型"
display(results_display)
assert np.isfinite(results[["test_nll", "test_perplexity"]].to_numpy()).all()
print("困惑度越低表示对真实下一个字符分配的平均概率越高;它不评价事实性或文学质量。")
测试集平均负对数似然 测试集困惑度 参数量 裁剪前梯度范数百分之九十五分位数
模型
单字符频率基线 3.171659 23.847002 NaN NaN
普通循环神经网络 2.813974 16.676055 25036.0 1.019423
长短期记忆网络 3.013302 20.354495 67084.0 0.627263
门控循环单元网络 2.867892 17.599884 53068.0 0.726081
困惑度越低表示对真实下一个字符分配的平均概率越高;它不评价事实性或文学质量。
def sample_text(model, prompt="Alice ", length=180, temperature=0.8):
    if temperature <= 0:
        raise ValueError("temperature 必须为正")
    model.eval()
    ids = [stoi.get(ch, UNK) for ch in prompt]
    state = None
    with torch.no_grad():
        # 先读入提示,建立状态。
        prefix = torch.tensor(ids, dtype=torch.long, device=DEVICE)[None, :]
        logits, state = model(prefix, state)
        current = prefix[:, -1:]
        generated = list(prompt)
        generator = torch.Generator(device=DEVICE).manual_seed(SEED)
        for _ in range(length):
            logits, state = model(current, state)
            prob = torch.softmax(logits[0, -1] / temperature, dim=0)
            next_id = int(torch.multinomial(prob, 1, generator=generator))
            token = itos[next_id]
            char = "�" if token == "<unk>" else token
            generated.append(char)
            current = torch.tensor([[next_id]], device=DEVICE)
    return "".join(generated)

best_kind = results.drop(index="unigram")["test_perplexity"].idxmin()
generated_sample = sample_text(models[best_kind])
print(f"已用{model_labels[best_kind]}生成 {len(generated_sample)} 个字符的定性诊断样例。")
print("生成片段不能替代预先定义的测试指标,也不能评价事实性或文学质量。")
已用普通循环神经网络生成 186 个字符的定性诊断样例。
生成片段不能替代预先定义的测试指标,也不能评价事实性或文学质量。

逐步理解本案例#

1. 怎样把文本变成可预测的序列#

  数据来自 Project Gutenberg 提供的《爱丽丝梦游仙境》英文原文。程序先删除电子书开头和结尾的说明,再建立字符词表,把每个字符映射为一个整数。字符级模型不需要先把英文文本切分为单词,能够直接展示词元、词表和序列张量之间的关系;代价是同一句话会形成更长的序列,生成文字也需要经过更多时间步。

  训练集、验证集和测试集必须按照原文顺序划分,不能随机打乱字符。如果把同一段落中彼此相邻的片段分到不同数据集,模型在测试时会看到与训练内容几乎相同的上下文。词表只根据训练文本建立,验证集或测试集中未出现在训练词表中的字符映射为未知标记。每个样本包含长度为 \(L\) 的输入序列,以及向右移动一个位置的观测标签序列,二者的维度均为 \((批量,L)\)

2. 自回归目标与简单基线#

  自回归语言模型把联合概率分解为每一步的条件概率,预测当前字符时只能使用更早出现的字符。单字符频率基线忽略上下文,始终使用训练文本中的字符频率进行预测,但它仍提供了一个基本的困惑度参考。如果循环网络不能优于该基线,应先检查输入与观测标签是否正确错开一个位置、隐藏状态是否正确传递、交叉熵在哪些维度上计算以及参数是否正常更新。

  输出层线性运算结果的维度为 \((批量,L,词表大小)\)计算交叉熵前,通常把前两个轴合并,使每个时间位置成为一个分类样本,观测标签也采用相同方式展开。损失越小,模型给真实下一个字符分配的平均概率越高。困惑度等于平均负对数似然的指数,数值越低通常越好,但只有在字符词表、数据处理方法和评价文本相同时才能直接比较。

3. 普通循环神经网络怎样共享参数#

  在时间 \(t\)普通 RNN 把当前字符的嵌入向量 \(\boldsymbol{x}_t\) 与前一个隐藏状态 \(\boldsymbol{h}_{t-1}\) 组合,得到新的隐藏状态 \(\boldsymbol{h}_t\)所有时间步使用同一组权重,因此模型可以处理不同长度的序列,并把过去的信息压缩到固定维度的隐藏状态中。后向传播需要沿展开的时间计算过程向前传递梯度,连续进行很多次矩阵乘法时,梯度可能逐渐接近 0,也可能变得很大。

  截断的随时间后向传播只在一个训练片段内计算梯度,可以控制内存和计算量,但也限制了模型能够直接学习的依赖距离。是否在不同批次之间传递隐藏状态,必须与文本片段的排列方式相匹配;本案例把每个窗口独立处理,避免把原文中并不相邻的片段错误连接起来。

4. LSTM 与 GRU 怎样帮助学习长期关系#

  LSTM 通过输入门、遗忘门和输出门控制细胞状态。当遗忘门接近 1 时,信息可以沿较直接的加法路径保留。GRU 使用更新门和重置门合并其中部分机制,参数通常少于相同隐藏维度的 LSTM。门控结构能够改善训练条件,但不能保证模型学到任意长的依赖关系;输入窗口长度、训练数据量和训练目标仍会限制可利用的信息。

  比较普通 RNN、LSTM 和 GRU 时,应保持嵌入维度、隐藏状态维度、数据划分、批量大小和训练轮数相同,并同时报告参数量。如果只固定隐藏状态维度,LSTM 会因为具有更多门而拥有更多参数,这一点必须在结论中说明。

5. 梯度裁剪在程序中的位置#

  先进行后向传播得到所有参数的梯度,再计算整体梯度范数并进行裁剪,最后由优化器更新参数。梯度裁剪可以限制一次参数更新的大小,避免爆炸梯度严重破坏模型参数;它不能修复输入中的无效数值、错误的观测标签或过大的学习率。记录裁剪前梯度范数的分位数,可以判断裁剪阈值是否经常被触发。如果几乎每一步都发生强烈裁剪,应继续查找训练设置不稳定的原因。

  固定随机种子可以控制参数初始化和批次顺序。不同计算设备仍可能使用结果略有差异的运算,因此完整实验报告应记录计算设备、PyTorch 版本、训练时间、词表摘要和数据校验值。

6. 怎样解释生成文字#

  生成时先给模型一段起始文字,模型预测下一个字符,再把抽样得到的字符追加到输入中。温度小于 1 时,概率分布更加集中,生成结果通常更保守;温度较大时,多样性和错误都会增加。生成片段只能用来直观检查模型是否学到拼写、空格和局部语法,不能替代测试集上的平均负对数似然。

  为了避免把记住训练文本误认为生成能力,还应检查较长的生成片段是否直接重复原文。该作品在美国属于公版,但其他司法辖区的使用者仍需自行确认;发布案例时应保留数据来源和权利说明。

7. 阅读结果的顺序#

  先确认三种模型的输入、输出和观测标签维度一致,再查看训练损失与验证损失,然后比较测试困惑度、参数量和梯度范数,最后再观察生成文字。生成文字看似流畅但困惑度较差,可能只是某一次抽样比较幸运;困惑度较低但生成结果大量重复,则可能与抽样方法或过于集中的概率分布有关。

  面向读者的标题、解释、代码注释、运行提示、坐标轴和表格列名均使用中文。英文只保留在原始语料、官方书名、Python 接口、变量名以及 RNN、LSTM 和 GRU 等缩写中。

本案例小结与局限#

  • 一本小说的词汇、写作风格和主题较为单一,测试文本也来自同一作品,因此不能根据本案例判断模型对其他作者作品的表现;

  • 字符级模型不需要解决单词切分问题,但需要较长序列才能表达词语和句子之间的长期关系;

  • 不同序列片段之间重新初始化状态,使模型能够直接使用的上下文长度受到 SEQ_LEN 限制。如果在相邻批次之间传递状态,还必须保证文本片段确实连续,并在适当位置使用 detach 停止跨批次计算梯度;

  • RNN、LSTM 和 GRU 的参数量不同,相同的隐藏状态维度不表示三种模型具有完全相同的容量,因此报告中必须同时列出参数量;

  • 梯度裁剪只能限制过大的梯度,不能修复错误的损失函数、张量维度或学习率;

  • 困惑度只衡量模型预测下一个词元的概率。生成文字看起来合理,不表示模型理解内容、事实正确或没有记住训练文本;

  • 作品是否属于公版取决于所在司法辖区,发布课程时应保留 Project Gutenberg 的说明和数据来源。

综合练习#

  1. 写出一个长度为 4 的普通 RNN 展开计算图,标出参数共享位置,并推导 (\partial L/\partial W_h) 的求和结构。

  2. SEQ_LEN 改为 16、64、128;固定总词元更新数,比较困惑度、时间和峰值内存。

  3. 实现跨相邻批次传递隐藏状态的状态化训练;每次优化后对状态 detach说明原因。

  4. 构造“首字符决定末字符”的合成任务,比较 RNN、LSTM、GRU 随距离增加的准确率。

  5. 在温度 0.5、0.8、1.2 下各用相同随机种子生成文本;解释多样性与局部连贯性的变化,但不要把定性样例当作测试结论。