LSTM 与 GRU:参考答案#
说明#
以下答案与正文10道题逐题对应。前四题给出主要推导与计算;第 5--7 题给出可以运行的程序和检查方法;第 8--10 题说明比较时需要保持哪些条件相同,以及实验结果能够说明什么、不能说明什么。
细胞状态和隐藏状态分别为
\[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)\) 的一半用于形成隐藏状态。这是给定门值后的状态计算;实际门值仍由输入、旧隐藏状态和可训练参数共同决定。
直接路径在每一步的局部雅可比近似为 \(\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 只是提供了可学习的加性保留通路,并不保证任意长梯度都不消失或爆炸。
按 \(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 时,旧状态和直接梯度更容易通过;但完整导数还包含门及候选隐藏状态对旧状态的导数。
忽略输出头时,参数量为
\[\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\)。相同的隐藏维度只能保证状态宽度相同,三种模型的门数、参数量、需要保存的激活值和乘加次数仍然不同;参数量接近时,隐藏维度又可能不同。因此应分别比较隐藏维度相同和参数量接近这两种情况,并写清每种比较中保持不变的条件。
可把所有门分别实现,以下为核心骨架。数学公式中的 \(\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 题可通过直接给定门值测试最后两行;完整单元还要分别测试批量轴、非法特征维度和数据类型。
LSTMCell常把四个门按框架规定的顺序堆叠,GRUCell的门排列及 reset 门在线性变换前后的位置也必须查清;先写出参数切片映射,再复制权重和偏置。使用float64、关闭随机操作,并比较h、c、输入梯度、旧状态梯度及每个参数梯度: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,并让各行存放的列向量转置互不相同,可以暴露错误广播。若公式约定与框架不同,不能只交换变量名;必须按照框架实际运算重写参数映射。
通过设置门偏置并把其他门权重置零,可使门近似常数。长度为 \(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 的候选隐藏状态通常依赖旧隐藏状态,输出损失还会经过其他路径;因此该测试验证的是机制趋势,而不是所有训练情形的精确公式。延迟复制任务应先生成固定的训练、验证和测试集合,并让三种模型共享输入编码、输出头形式、种子集合、批量大小、更新次数和每次更新样本数。先比较相同 \(d_h\),再为每种模型选择参数量接近的宽度;两张表都报告测试准确率、训练时间、固定 512 个序列的预测时间、峰值内存、参数量和循环单元调用次数。
门控模型通常在长延迟任务上更容易训练,但结果仍取决于优化方法和数据;相同宽度不等于计算成本相同,参数量接近也不等于运行时间接近。如果不同模型需要不同的学习率,应在同一个验证集上尝试相同的一组学习率,并按照相同标准选择,不能只为其中一个模型反复调整。
用第 4 题的一般参数公式寻找隐藏维度,使包含输入层和输出头后的总参数差最小;实际参数量仍从程序统计并写入表格。固定划分、种子集合、优化器族、更新次数、序列长度与评价代码,各模型的学习率可在同一候选集合中由验证集选择。报告测试指标、训练时间、固定批量预测时间、峰值内存、参数量和估算乘加量。
参数量相近时,LSTM 和 GRU 的状态宽度不同,可能影响表示能力;相同参数量也不能保证融合算子的硬件效率相同。因此结论只能描述给定实现、设备和任务上的性能—成本关系。
对每个依赖长度使用同一规则生成互不重叠的训练、验证和测试集,并固定各集合规模、种子集合、最大更新次数、早停规则和指标。GRU 与 LSTM 都运行完整的同一学习率网格,报告所有组合的预测指标、训练时间、固定批量预测时间、峰值内存、梯度范数和实际更新数;最终测试方案只能由验证结果确定。
随依赖长度增长,任务计算本身也增多,因此训练时间上升不等同于优化更困难。若某学习率只在短序列稳定,应说明梯度尺度和更新次数的变化,不能删除失败运行后只比较成功点。