LSTM 与 GRU:参考答案

目录

\[ \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} \]

LSTM 与 GRU:参考答案#

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

说明#

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

  1. 细胞状态和隐藏状态分别为

    \[c_t=0.8(2)+0.25(-0.4)=1.5, \qquad h_t=0.5\tanh(1.5)\approx0.45257.\]

    遗忘门保留旧细胞状态的 \(80\%\)⁠,输入门把候选细胞状态的 \(25\%\) 写入,输出门使 \(\tanh(c_t)\) 的一半用于形成隐藏状态。这是给定门值后的状态计算;实际门值仍由输入、旧隐藏状态和可训练参数共同决定。

  2. 直接路径在每一步的局部雅可比近似为 \(\operatorname{diag}(\boldsymbol{f}_t)\)⁠,所以

    \[\frac{\partial\boldsymbol{c}_t}{\partial\boldsymbol{c}_{t-k}} \approx \operatorname{diag}(\boldsymbol{f}_t) \cdot\operatorname{diag}(\boldsymbol{f}_{t-1}) \cdots \operatorname{diag}(\boldsymbol{f}_{t-k+1}).\]

    给定维度的三步系数是 \(0.9(0.8)(0.5)=0.36\)⁠。即使门值都在 \((0,1)\)⁠,持续小于 1 的乘积仍会衰减;门本身也可能因 sigmoid 函数饱和而难以更新,完整雅可比矩阵还包含隐藏状态和门之间的间接路径。因此 LSTM 只是提供了可学习的加性保留通路,并不保证任意长梯度都不消失或爆炸。

  3. \(h_t=z_th_{t-1}+(1-z_t)\widetilde h_t\)⁠,有

    \[h_t=0.75(2)+0.25(-0.4)=1.4.\]

    若暂时把 \(z_t\)\(\widetilde h_t\) 看作不依赖 \(h_{t-1}\) 的给定量,则 \(\partial h_t/\partial h_{t-1}\approx z_t=0.75\)⁠。更新门接近 1 时,旧状态和直接梯度更容易通过;但完整导数还包含门及候选隐藏状态对旧状态的导数。

  4. 忽略输出头时,参数量为

    \[\begin{split}\begin{aligned} N_{\mathrm{RNN}}&=d_h(d_x+d_h+1),\\ N_{\mathrm{GRU}}&=3d_h(d_x+d_h+1),\\ N_{\mathrm{LSTM}}&=4d_h(d_x+d_h+1). \end{aligned}\end{split}\]

    代入 \(d_x=5,d_h=4\) 得到 \(N_{\mathrm{RNN}}=40\)⁠、\(N_{\mathrm{GRU}}=120\)⁠、\(N_{\mathrm{LSTM}}=160\)⁠。相同的隐藏维度只能保证状态宽度相同,三种模型的门数、参数量、需要保存的激活值和乘加次数仍然不同;参数量接近时,隐藏维度又可能不同。因此应分别比较隐藏维度相同和参数量接近这两种情况,并写清每种比较中保持不变的条件。

  5. 可把所有门分别实现,以下为核心骨架。数学公式中的 \(\bx_t\)⁠、\(\boldsymbol{h}_t\)\(\boldsymbol{c}_t\) 均为列向量;程序为了同时处理 \(m\) 个样本,把一批列向量 \(\boldsymbol{u}_1,\ldots,\boldsymbol{u}_B\) 存为 \(\boldsymbol{U}=[\boldsymbol{u}_1\trans;\ldots;\boldsymbol{u}_B\trans]\in\mathbb{R}^{m\times d}\)⁠,因此每一行都是相应列向量的转置:

    import torch
    
    def lstm_step(x, h, c, W_i, W_f, W_o, W_c, b_i, b_f, b_o, b_c):
        if x.ndim != 2 or h.ndim != 2 or c.shape != h.shape:
            raise ValueError("输入或状态维度错误")
        u = torch.cat([x, h], dim=-1)
        i = torch.sigmoid(u @ W_i.T + b_i)
        f = torch.sigmoid(u @ W_f.T + b_f)
        o = torch.sigmoid(u @ W_o.T + b_o)
        candidate = torch.tanh(u @ W_c.T + b_c)
        c_new = f * c + i * candidate
        h_new = o * torch.tanh(c_new)
        return {"i": i, "f": f, "o": o,
                "candidate": candidate, "c": c_new, "h": h_new}
    
    def gru_step(x, h, W_z, U_z, W_r, U_r, W_h, U_h,
                 b_z, b_r, b_h):
        z = torch.sigmoid(x @ W_z.T + h @ U_z.T + b_z)
        r = torch.sigmoid(x @ W_r.T + h @ U_r.T + b_r)
        candidate = torch.tanh(x @ W_h.T + (r * h) @ U_h.T + b_h)
        h_new = z * h + (1.0 - z) * candidate
        return {"z": z, "r": r, "candidate": candidate, "h": h_new}
    

    第 1、3 题可通过直接给定门值测试最后两行;完整单元还要分别测试批量轴、非法特征维度和数据类型。

  6. LSTMCell 常把四个门按框架规定的顺序堆叠,GRUCell 的门排列及 reset 门在线性变换前后的位置也必须查清;先写出参数切片映射,再复制权重和偏置。使用 float64关闭随机操作,并比较 hc输入梯度、旧状态梯度及每个参数梯度:

    torch.testing.assert_close(manual_h, library_h, rtol=1e-8, atol=1e-10)
    torch.testing.assert_close(manual_grad, library_grad,
                               rtol=1e-7, atol=1e-9)
    

    批量大小取 3,并让各行存放的列向量转置互不相同,可以暴露错误广播。若公式约定与框架不同,不能只交换变量名;必须按照框架实际运算重写参数映射。

  7. 通过设置门偏置并把其他门权重置零,可使门近似常数。长度为 \(T\) 的零输入序列中,LSTM 直接路径趋势约为 \(f^T\)⁠,GRU 在候选状态局部不依赖旧状态时约为 \(z^T\)⁠。对 \(0.1,0.5,0.9\) 分别记录 \(\lVert\partial\boldsymbol{s}_T/\partial\boldsymbol{s}_0\rVert\)⁠,应看到保留门越大,衰减越慢。

    使用极大的正、负门输入检查输出有限且无 NaN实际梯度可能偏离简单乘积,因为门、LSTM 的候选细胞状态和 GRU 的候选隐藏状态通常依赖旧隐藏状态,输出损失还会经过其他路径;因此该测试验证的是机制趋势,而不是所有训练情形的精确公式。

  8. 延迟复制任务应先生成固定的训练、验证和测试集合,并让三种模型共享输入编码、输出头形式、种子集合、批量大小、更新次数和每次更新样本数。先比较相同 \(d_h\)⁠,再为每种模型选择参数量接近的宽度;两张表都报告测试准确率、训练时间、固定 512 个序列的预测时间、峰值内存、参数量和循环单元调用次数。

    门控模型通常在长延迟任务上更容易训练,但结果仍取决于优化方法和数据;相同宽度不等于计算成本相同,参数量接近也不等于运行时间接近。如果不同模型需要不同的学习率,应在同一个验证集上尝试相同的一组学习率,并按照相同标准选择,不能只为其中一个模型反复调整。

  9. 用第 4 题的一般参数公式寻找隐藏维度,使包含输入层和输出头后的总参数差最小;实际参数量仍从程序统计并写入表格。固定划分、种子集合、优化器族、更新次数、序列长度与评价代码,各模型的学习率可在同一候选集合中由验证集选择。报告测试指标、训练时间、固定批量预测时间、峰值内存、参数量和估算乘加量。

    参数量相近时,LSTM 和 GRU 的状态宽度不同,可能影响表示能力;相同参数量也不能保证融合算子的硬件效率相同。因此结论只能描述给定实现、设备和任务上的性能—成本关系。

  10. 对每个依赖长度使用同一规则生成互不重叠的训练、验证和测试集,并固定各集合规模、种子集合、最大更新次数、早停规则和指标。GRU 与 LSTM 都运行完整的同一学习率网格,报告所有组合的预测指标、训练时间、固定批量预测时间、峰值内存、梯度范数和实际更新数;最终测试方案只能由验证结果确定。

    随依赖长度增长,任务计算本身也增多,因此训练时间上升不等同于优化更困难。若某学习率只在短序列稳定,应说明梯度尺度和更新次数的变化,不能删除失败运行后只比较成功点。