序列建模预备知识:参考答案

目录

\[ \begin{align}\begin{aligned}\newcommand{\ba}{\boldsymbol{a}} \newcommand{\bb}{\boldsymbol{b}} \newcommand{\be}{\boldsymbol{e}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bw}{\boldsymbol{w}} \newcommand{\bx}{\boldsymbol{x}} \newcommand{\by}{\boldsymbol{y}} \newcommand{\bz}{\boldsymbol{z}} \newcommand{\bd}{\boldsymbol{d}} \newcommand{\bv}{\boldsymbol{v}} \newcommand{\bs}{\boldsymbol{s}}\\\newcommand{\btheta}{\boldsymbol{\theta}} \newcommand{\bbeta}{\boldsymbol{\beta}} \newcommand{\bgamma}{\boldsymbol{\gamma}} \newcommand{\bsigma}{\boldsymbol{\sigma}} \newcommand{\md}{\mbox{d}} \newcommand{\bmu}{\boldsymbol{\mu}} \newcommand{\bone}{\boldsymbol{1}} \newcommand{\bzero}{\boldsymbol{0}} \newcommand{\bepsilon}{\boldsymbol{\epsilon}} \newcommand{\bphi}{\boldsymbol{\phi}} \newcommand{\bh}{\boldsymbol{h}} \newcommand{\bc}{\boldsymbol{c}} \newcommand{\br}{\boldsymbol{r}} \newcommand{\bQ}{\boldsymbol{Q}} \newcommand{\bK}{\boldsymbol{K}} \newcommand{\bV}{\boldsymbol{V}} \newcommand{\bSigma}{\boldsymbol{\Sigma}} \newcommand{\bg}{\boldsymbol{g}} \newcommand{\bxi}{\boldsymbol{\xi}} \newcommand{\bvarepsilon}{\boldsymbol{\varepsilon}} \newcommand{\bdelta}{\boldsymbol{\delta}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bJ}{\boldsymbol{J}} \newcommand{\bp}{\boldsymbol{p}} \newcommand{\bi}{\boldsymbol{i}} \newcommand{\bo}{\boldsymbol{o}} \newcommand{\bE}{\boldsymbol{E}} \newcommand{\bH}{\boldsymbol{H}} \newcommand{\bL}{\boldsymbol{L}} \newcommand{\bu}{\boldsymbol{u}} \newcommand{\bLambda}{\boldsymbol{\Lambda}} \newcommand{\trans}{^{\rm\scriptsize T}} \newcommand{\var}{\mathrm{var}}\\\newcommand{\bA}{\boldsymbol{A}} \newcommand{\bB}{\boldsymbol{B}} \newcommand{\bC}{\boldsymbol{C}} \newcommand{\bD}{\boldsymbol{D}} \newcommand{\bG}{\boldsymbol{G}} \newcommand{\bI}{\boldsymbol{I}} \newcommand{\bM}{\boldsymbol{M}} \newcommand{\bP}{\boldsymbol{P}} \newcommand{\bS}{\boldsymbol{S}} \newcommand{\bU}{\boldsymbol{U}} \newcommand{\bW}{\boldsymbol{W}} \newcommand{\bX}{\boldsymbol{X}} \newcommand{\bY}{\boldsymbol{Y}} \newcommand{\bZ}{\boldsymbol{Z}} \newcommand{\cotp}{\textcolor[RGB]{48,209,88}{TP}} \newcommand{\cotn}{\textcolor[RGB]{100,210,255}{TN}} \newcommand{\cofp}{\textcolor[RGB]{94,92,230}{FP}} \newcommand{\cofn}{\textcolor[RGB]{191,90,242}{FN}}\\\newcommand{\numcotp}{\textcolor[RGB]{48,209,88}{50}} \newcommand{\numcotn}{\textcolor[RGB]{100,210,255}{30}} \newcommand{\numcofp}{\textcolor[RGB]{94,92,230}{10}} \newcommand{\numcofn}{\textcolor[RGB]{191,90,242}{10}} \DeclareMathOperator*{\argmin}{arg\,min}\end{aligned}\end{align} \]

序列建模预备知识:参考答案#

返回正文练习 · 返回答案索引

说明#

以下答案与正文10道题逐题对应。前四题给出主要推导与计算;第 5--7 题给出可以运行的程序和检查方法;第 8--10 题说明比较时需要保持哪些条件相同,以及实验结果能够说明什么、不能说明什么。

  1. 平稳时 \(\mathbb{E}(x_t)=\mathbb{E}(x_{t-1})=\mu\)⁠,所以

    \[\mu=c+\phi\mu, \qquad \mu=\frac{c}{1-\phi}.\]

    \(u_t=x_t-\mu\)⁠,则 \(u_t=\phi u_{t-1}+\varepsilon_t\)⁠。噪声与过去状态不相关,因此

    \[\operatorname{Var}(u_t) =\phi^2\operatorname{Var}(u_{t-1})+\sigma^2, \qquad \operatorname{Var}(x_t)=\frac{\sigma^2}{1-\phi^2}.\]

    \(|\phi|<1\) 保证递推中的 \(\phi^h\)\(h\) 增大而趋于 0,且方差级数 \(\sigma^2\sum_{j=0}^{\infty}\phi^{2j}\) 收敛。该结论还依赖噪声与过去状态不相关;只给出均值和方差并不能说明过程必为高斯。

  2. 逐步代入得到

    \[\widehat x_{t+1\mid t}=1+0.5(6)=4, \qquad \widehat x_{t+2\mid t}=1+0.5(4)=3, \qquad \widehat x_{t+3\mid t}=1+0.5(3)=2.5.\]

    第二步不再使用真实的 \(x_{t+1}\)⁠,而使用第一步的预测 4;第三步又使用第二步的预测 3。因此前一步的参数误差和预测误差会进入后一步输入。本例中影响按 \(0.5\) 缩小;若 \(|\widehat\phi|\) 接近或超过 1,误差可能保留更久或被放大。

  3. 概率乘法公式先给出 \(p(w_1,w_2)=p(w_1)p(w_2\mid w_1)\)⁠,继续展开可得

    \[p(w_{1:T}) =p(w_1)\prod_{t=2}^{T}p(w_t\mid w_{1:t-1}) =\prod_{t=1}^{T}p(w_t\mid w_{<t}).\]

    对有效位置集合 \(\mathcal I\) 取负对数,乘积变为和:

    \[-\log p(w_{1:T}) =-\sum_{t\in\mathcal I}\log p(w_t\mid w_{<t}), \qquad \mathcal J=-\frac1{|\mathcal I|}\sum_{t\in\mathcal I}\log p(w_t\mid w_{<t}).\]

    padding 是为批处理添加的存储内容,不是原序列中的随机变量。把它计入既会训练模型预测人为符号,也会让长、短样本因填充数量不同而得到不同权重。

  4. 有效概率是 \(0.8,0.5,0.1\)⁠,所以

    \[\mathcal J =-\frac{\log0.8+\log0.5+\log0.1}{3} \approx1.0730, \qquad \operatorname{PPL}=\exp(\mathcal J)\approx2.924.\]

    若错误计入第三个位置,则 \(\mathcal J_{\mathrm{wrong}}=-(\log0.8+\log0.5+\log0.25+\log0.1)/4\approx1.1513\)⁠,困惑度约为 \(3.163\)⁠。数值变大还是变小取决于被误计位置的概率,关键问题是指标不再对应真实有效词元。

  5. 一个最小实现骨架如下;实际项目可把频数阈值和特殊词元顺序作为显式参数。

    from collections import Counter
    import numpy as np
    
    def build_vocab(train_tokens, min_freq=1):
        count = Counter(tok for sent in train_tokens for tok in sent)
        vocab = {"[PAD]": 0, "[UNK]": 1}
        for tok in sorted(count):
            if count[tok] >= min_freq:
                vocab[tok] = len(vocab)
        return vocab
    
    def encode(tokens, vocab):
        unk = vocab["[UNK]"]
        return [vocab.get(tok, unk) for tok in tokens]
    
    def decode(ids, vocab):
        inverse = {value: key for key, value in vocab.items()}
        return [inverse.get(int(index), "[UNK]") for index in ids]
    
    def collate_batch(sequences, pad_id):
        if not sequences:
            raise ValueError("批次不能为空")
        lengths = np.asarray([len(seq) for seq in sequences], dtype=np.int64)
        if np.any(lengths == 0):
            raise ValueError("本实现不接收空序列")
        width = int(lengths.max())
        ids = np.full((len(sequences), width), pad_id, dtype=np.int64)
        mask = np.zeros_like(ids, dtype=bool)
        for row, seq in enumerate(sequences):
            ids[row, :len(seq)] = seq
            mask[row, :len(seq)] = True
        assert ids.shape == mask.shape
        assert lengths.shape == (len(sequences),)
        return ids, mask, lengths
    

    词表只能读取训练集;用验证集词元测试时,未登录词应编码为 [UNK]decode(encode(x)) 只能在词元已登录时恢复原词元,不能把 [UNK] 唯一还原。

  6. 可执行的 NumPy 版本为

    import numpy as np
    
    def masked_cross_entropy(scores, targets, mask):
        scores = np.asarray(scores, dtype=np.float64)
        targets = np.asarray(targets, dtype=np.int64)
        mask = np.asarray(mask, dtype=bool)
        if scores.ndim != targets.ndim + 1:
            raise ValueError("scores 应比 targets 多一个类别轴")
        if scores.shape[:-1] != targets.shape or mask.shape != targets.shape:
            raise ValueError("批次与序列轴不一致")
        if not mask.any():
            raise ValueError("至少需要一个有效目标")
        if np.any((targets[mask] < 0) | (targets[mask] >= scores.shape[-1])):
            raise ValueError("目标编号越界")
        maximum = scores.max(axis=-1, keepdims=True)
        log_norm = maximum + np.log(
            np.exp(scores - maximum).sum(axis=-1, keepdims=True)
        )
        log_prob = scores - log_norm
        safe_targets = np.where(mask, targets, 0)
        loss = -np.take_along_axis(
            log_prob, safe_targets[..., None], axis=-1
        )[..., 0]
        return loss[mask].mean()
    

    safe_targets 只用于避免以 -100 等 padding 标签索引类别轴;这些位置随后仍由 mask 排除,不会进入损失。与 torch.nn.functional.cross_entropy 对照时,也可先把有效位置展平后选出,或令 padding 标签为 ignore_index极大绝对值测试应得到有限结果;全无效掩码应明确报错,不能静默除以 0。

  7. 循环版对每个有效位置分别计算 \(m=\max_k z_k\)\(m+\log\sum_k\exp(z_k-m)-z_y\)⁠;批量版沿最后一轴同时完成相同运算。核验应保存每个位置的损失数组,而不只比较平均值:

    np.testing.assert_allclose(loop_terms, batch_terms, rtol=1e-10, atol=1e-12)
    np.testing.assert_allclose(sum(loop_terms), batch_terms.sum())
    

    梯度可与 自动微分结果 对照,或对若干输入元素使用 中心差分 \([J(z+h)-J(z-h)]/(2h)\)⁠。计时前先预热,重复多次并取典型运行时间;两版必须处理同一数组、同一有效词元数。批量版通常更快,但在极小输入上函数调用和数组分配可能占主导,因此不能把单个规模的计时外推到所有规模。

  8. 先一次性确定训练、验证和测试文本,再只在训练文本上建立各自的词表或训练子词器。三个模型使用相同的预测主体、相近参数量、优化器、更新次数和随机种子,例如 {11, 22, 33, 44, 55}若一次更新的有效词元数不同,还应补充总处理词元数相同的结果。记录每个随机种子下如何使用验证集选择模型;模型和训练方法确定后,测试集只评价一次。

    结果表至少包含准确率或困惑度、训练墙钟时间、固定 256 个样本的预测时间、峰值内存、词表大小、平均长度和未登录率。字符表示通常词表小但序列长,词表示序列短但可能有较高未登录率,子词表示常居中;这些是待验证的经验趋势,不是必然的精度排序。不同分词下困惑度的基本单位不同,因此不能脱离词元单位直接比较数值。

  9. 三种 collator 必须接收相同顺序的样本,并使用同一训练/验证/测试划分、种子集合、模型初值、优化器、更新次数和原始批量大小。动态填充与分桶会改变每批最大长度,若希望严格比较优化轨迹,可在评价阶段先比较,或固定样本顺序并记录实际有效词元数。对同一未填充样本,在 eval 模式下三种预测应在容差内一致。

    报告测试指标、训练时间、固定批量预测时间、填充比例 \(1-N_{\mathrm{valid}}/N_{\mathrm{stored}}\) 和峰值内存。全局填充通常浪费最多;分桶可能最高效,但会改变批次组成,并增加数据处理流程的复杂度。如果批次组成也发生变化,就不能认定性能差异完全由填充方法造成。

  10. 在训练集上分别估计二元和三元条件概率,并使用同一种平滑规则。固定词表、时间顺序划分、种子集合、训练语料和评价程序;每种方法遍历相同次数或处理相同总词元数。预测时缓存计数表,使用固定批量前缀测量时间,并把模型内存定义为所有计数表或参数所占字节数。

    报告测试负对数似然、困惑度或下一词准确率,以及训练时间、固定批量预测时间和内存。更长的上下文能表示更具体的依赖关系,但可用的计数会更稀疏,占用内存也更大,对未见过的上下文更依赖平滑或回退方法。词元化方法或词表不同时,困惑度的计算单位也不同,因此不能直接根据困惑度判断哪种方法更好。