python神经网络编程入门(二十)——RNN LSTM 数学原理与结构拆解

引言:直觉有了,现在把每个零件拆开看

上一篇用"传送带 + 收费站"的比喻建立了 LSTM 的整体直觉:细胞状态 ctc_tct 是高速公路,三个门(遗忘门、输入门、输出门)是收费站。R-N-N 在 30 步后梯度归零,而 L-S-T-M 在 50 步后还剩 7.7%------差距一目了然。

但比喻只能帮理解"为什么",要真正写出代码,得把每个零件的规格搞清楚。

打个比方:上一篇相当于从外观 看一辆车------知道它有发动机、轮胎、方向盘,大概知道怎么开。这一篇是打开引擎盖,看每个零件的规格型号:遗忘门的输入是哪些信号?权重矩阵长什么样?sigmoid 和 tanh 各管什么?一个 LSTM 单元到底有多少参数?

这些问题的答案,汇总起来就是 LSTM 的 7 个核心公式。这 7 个公式是面试必考题,也是手写代码的"施工图纸"------图纸上画了什么,代码就写什么。

🎯 本章目标

  1. 默写出 LSTM 完整的 7 个前向公式;
  2. 理解 sigmoid(控制"开/关比例")和 tanh(承载"正负信息")的分工逻辑;
  3. 算出 LSTM 的总参数量,和 RNN 做量化对比;
  4. 手写 lstm_forward 函数,打印每一步张量的形状变化。

一、LSTM 完整公式全景

1.1 7 个公式一览

先把 LSTM 一个时间步的完整计算流程摆出来。输入是 xtx_txt(当前时刻的数据)和 ht−1h_{t-1}ht−1(上一时刻的隐藏状态),输出是 hth_tht(新隐藏状态)和 ctc_tct(新细胞状态)。

阶段一:三个门 + 候选信息(并行计算)

这四个值的输入完全相同------都是 ht−1,xth_{t-1}, x_tht−1,xt 的拼接,区别只在各自的权重矩阵和激活函数:

ft=σ(Wf⋅ht−1,xt+bf)遗忘门it=σ(Wi⋅ht−1,xt+bi)输入门c~t=tanh⁡(Wc⋅ht−1,xt+bc)候选细胞状态ot=σ(Wo⋅ht−1,xt+bo)输出门 \begin{aligned} f_t &= \sigma(W_f \cdot h_{t-1}, x_t + b_f) \quad &\text{遗忘门} \\4pt i_t &= \sigma(W_i \cdot h_{t-1}, x_t + b_i) \quad &\text{输入门} \\4pt \tilde{c}_t &= \tanh(W_c \cdot h_{t-1}, x_t + b_c) \quad &\text{候选细胞状态} \\4pt o_t &= \sigma(W_o \cdot h_{t-1}, x_t + b_o) \quad &\text{输出门} \end{aligned} ftitc~tot=σ(Wf⋅ht−1,xt+bf)=σ(Wi⋅ht−1,xt+bi)=tanh(Wc⋅ht−1,xt+bc)=σ(Wo⋅ht−1,xt+bo)遗忘门输入门候选细胞状态输出门

阶段二:状态更新(串行,依赖阶段一的结果)

ct=ft⊙ct−1+it⊙c~t c_t = f_t \odot c_{t-1} + i_t \odot \tilde{c}_t ct=ft⊙ct−1+it⊙c~t

阶段三:输出

ht=ot⊙tanh⁡(ct) h_t = o_t \odot \tanh(c_t) ht=ot⊙tanh(ct)

一共 7 个公式。前 4 个可以并行计算 (输入相同,只是权重不同),后 3 个必须按顺序来 ------ctc_tct 依赖 ft,it,c~tf_t, i_t, \tilde{c}_tft,it,c~t,hth_tht 依赖 oto_tot 和 ctc_tct。

1.2 "四合一"权重矩阵设计

前 4 个公式的输入全是 ht−1,xth_{t-1}, x_tht−1,xt(维度 HHH 的旧状态 + 维度 III 的新输入 → 拼接成 H+IH+IH+I)。如果用四组独立矩阵,代码会非常冗长。实际实现时,把四组权重横向摞成一个大矩阵

W=WfWiWcWo,b=bfbibcboW = \begin{bmatrix} W_f \\ W_i \\ W_c \\ W_o \end{bmatrix}, \quad b = \begin{bmatrix} b_f \\ b_i \\ b_c \\ b_o \end{bmatrix}W= WfWiWcWo ,b= bfbibcbo

每个 W∗W_*W∗ 的形状都是 H×(H+I)H \times (H+I)H×(H+I),四个摞一起就是 4H×(H+I)4H \times (H+I)4H×(H+I)。一次矩阵乘法算完,再按行切成四段:

python 复制代码
# 拼接后的输入: (B, H+I)
concat = np.hstack([h_prev, x_t])
# 一次矩阵乘法算出四个门的原始值: (B, 4H)
gates = np.dot(concat, W.T) + b
# 切成四段,每段 (B, H)
f = sigmoid(gates[:, 0:H])
i = sigmoid(gates[:, H:2*H])
c_tilde = np.tanh(gates[:, 2*H:3*H])
o = sigmoid(gates[:, 3*H:4*H])

这个"四合一"技巧把四次矩阵乘法压缩成一次,后续手写 lstm_forward 时直接用到。

上图是 LSTM 一个时间步的完整数据流。顶层是 ht−1h_{t-1}ht−1 和 xtx_txt 拼接后送入四个并行分支------每个分支对应一组线性变换 + 激活函数。中间是细胞状态更新:遗忘门擦旧(ft⊙ct−1f_t \odot c_{t-1}ft⊙ct−1)、输入门写新(it⊙c~ti_t \odot \tilde{c}_tit⊙c~t),两者相加得到新 ctc_tct。右侧是隐藏状态输出:ctc_tct 先过 tanh⁡\tanhtanh 压缩,再由输出门 oto_tot 过滤。注意图中两个 tanh⁡\tanhtanh 操作------候选状态里的 tanh⁡\tanhtanh 和 hth_tht 输出路径上的 tanh⁡(ct)\tanh(c_t)tanh(ct) 是两个不同的操作,只是用了同一个函数。


二、逐门拆解

2.1 遗忘门:记忆的"橡皮擦"

公式

ft=σ(Wf⋅ht−1,xt+bf)f_t = \sigma(W_f \cdot h_{t-1}, x_t + b_f)ft=σ(Wf⋅ht−1,xt+bf)

输入 :上一时刻的隐藏状态 ht−1h_{t-1}ht−1(维度 HHH)+ 当前输入 xtx_txt(维度 III),拼接后维度 H+IH+IH+I。

计算 :拼接向量 × 权重矩阵 WfW_fWf(形状 H×(H+I)H \times (H+I)H×(H+I))+ 偏置 bfb_fbf(维度 HHH),再过 sigmoid,输出 ftf_tft(维度 HHH,每个元素在 0∼10 \sim 10∼1 之间)。

作用 :ftf_tft 和旧细胞状态 ct−1c_{t-1}ct−1 逐元素相乘。ftf_tft 的某个元素接近 0 → 对应位置的旧记忆被"擦除";接近 1 → 原样保留。

生活类比------复习笔记 :期末考试前翻笔记本。看到"第三章的公式推导"------这部分已经会了,划掉(ft≈0f_t \approx 0ft≈0);看到"第五章的关键结论"------这个要记牢,留着(ft≈1f_t \approx 1ft≈1)。遗忘门就是这支荧光笔:哪些该忘、哪些该记,看一眼旧笔记(ht−1h_{t-1}ht−1)和当前复习进度(xtx_txt)综合判断。

⚠️ 关键细节------偏置 bfb_fbf 的初始化 :实践中通常把 bfb_fbf 初始化为 1(或接近 1 的正数),而不是 0。因为训练刚开始时,模型还不知道哪些信息重要,默认"全保留"比默认"全遗忘"更安全------信息留着后面可以再忘,但一旦忘掉就永远回不来。这个技巧在面试中也经常被问到。

2.2 输入门 + 候选状态:新信息的"安检通道"

公式

it=σ(Wi⋅ht−1,xt+bi)i_t = \sigma(W_i \cdot h_{t-1}, x_t + b_i)it=σ(Wi⋅ht−1,xt+bi)

c~t=tanh⁡(Wc⋅ht−1,xt+bc)\tilde{c}_t = \tanh(W_c \cdot h_{t-1}, x_t + b_c)c~t=tanh(Wc⋅ht−1,xt+bc)

输入 :和遗忘门完全一样------ht−1h_{t-1}ht−1 和 xtx_txt 的拼接。

计算 :同样是拼接向量 × 权重矩阵 + 偏置,区别在于激活函数------iti_tit 用 sigmoid(输出 0∼10 \sim 10∼1),c~t\tilde{c}_tc~t 用 tanh⁡\tanhtanh(输出 −1∼1-1 \sim 1−1∼1)。

作用 :iti_tit 决定"写多少",c~t\tilde{c}_tc~t 提供"写什么"。两者逐元素相乘后加到细胞状态上。

生活类比------机场安检 :旅客(新输入 xtx_txt)来到安检口。安检员看一眼旅客的登机牌(ht−1h_{t-1}ht−1),做出两个判断:

  1. 这个旅客可以过吗? (输入门 iti_tit)------签证有效、没带违禁品 → 放行(it≈1i_t \approx 1it≈1);签证过期 → 拦下(it≈0i_t \approx 0it≈0)。
  2. 过了之后,哪些行李要带上飞机? (候选状态 c~t\tilde{c}_tc~t)------笔记本电脑放托盘、外套脱掉、液体扔掉。这是对原始输入的一次"重新打包"。

iti_tit 和 c~t\tilde{c}_tc~t 配合:安检员决定"放行"(iti_tit),同时对行李做"重新打包"(c~t\tilde{c}_tc~t),打包好的行李才送上飞机(加入 ctc_tct)。

2.3 细胞状态更新:传送带上的"加减货"

公式

ct=ft⊙ct−1+it⊙c~tc_t = f_t \odot c_{t-1} + i_t \odot \tilde{c}_tct=ft⊙ct−1+it⊙c~t

这是整个 LSTM 的灵魂。两个操作,一个符号:

  • ft⊙ct−1f_t \odot c_{t-1}ft⊙ct−1:减货------遗忘门把传送带上不重要的旧货扔掉;
  • it⊙c~ti_t \odot \tilde{c}_tit⊙c~t:加货------输入门把新打包好的货放上来。

核心 :这是逐元素加法和乘法 ,不是矩阵乘法。上一章花了大幅论证:加法让梯度能沿 ftf_tft 这条路径近乎无损回传------∂ct/∂ct−1=diag(ft)\partial c_t / \partial c_{t-1} = \mathrm{diag}(f_t)∂ct/∂ct−1=diag(ft),只要 ftf_tft 接近 1,梯度就不衰减。

一个边界情况 :如果 ft=1f_t = \mathbf{1}ft=1(全 1 向量)且 it=0i_t = \mathbf{0}it=0(全 0 向量),那么 ct=ct−1c_t = c_{t-1}ct=ct−1------信息完全不变 地传递下去。RNN 做不到这一点,因为 ht=tanh⁡(Wht−1+... )h_t = \tanh(W h_{t-1} + \dots)ht=tanh(Wht−1+...) 里矩阵乘法必然改变信息。LSTM 却能做到------这就是"高速公路"的字面含义。

2.4 输出门:对外发布"精选摘要"

公式

ot=σ(Wo⋅ht−1,xt+bo)o_t = \sigma(W_o \cdot h_{t-1}, x_t + b_o)ot=σ(Wo⋅ht−1,xt+bo)

ht=ot⊙tanh⁡(ct)h_t = o_t \odot \tanh(c_t)ht=ot⊙tanh(ct)

输入 :和前面三个门一样------ht−1h_{t-1}ht−1 和 xtx_txt 拼接。

计算 :oto_tot 通过 sigmoid(0∼10 \sim 10∼1),然后和 tanh⁡(ct)\tanh(c_t)tanh(ct)(−1∼1-1 \sim 1−1∼1)逐元素相乘得到 hth_tht。

作用 :细胞状态 ctc_tct 是 LSTM 的"内部数据库"------什么都有。但外界(下一层网络、下一个时间步的四个门)不需要看到全部。输出门做两件事:

  1. 先用 tanh⁡\tanhtanh 把 ctc_tct 压缩到 (−1,1)(-1, 1)(−1,1)------数值太大会让后续计算不稳定;
  2. 再用 oto_tot 过滤------哪些维度的信息对外发布、哪些藏起来。

生活类比------公司月报 :公司内部数据库(ctc_tct)记录了所有项目的详细数据。月底向管理层汇报时(hth_tht),数据员(输出门 oto_tot)筛选出"本月关键指标"------不需要把所有原始数据都打印出来,只挑管理层关心的部分(ot≈1o_t \approx 1ot≈1 的维度),其余归档(ot≈0o_t \approx 0ot≈0 的维度)。


三、激活函数分工:为什么门用 sigmoid,状态用 tanh?

这可能是 LSTM 设计中最常被问到的问题。答案只有三句话,但每一句都有分量。

sigmoid 输出 (0,1)(0, 1)(0,1)------天然适合做"比例阀"。

门的本质是"控制流量":让多少旧信息通过(遗忘门)、写入多少新信息(输入门)、对外暴露多少(输出门)。这些操作的语义天然是"0 到 1 的比例"。sigmoid 恰好把任意实数映射到 (0,1)(0,1)(0,1),和"比例阀"完美对接。

tanh⁡\tanhtanh 输出 (−1,1)(-1, 1)(−1,1)------天然适合做"信息载体"。

候选状态 c~t\tilde{c}_tc~t 和 tanh⁡(ct)\tanh(c_t)tanh(ct) 都是信息的载体。信息有正有负------"喜欢"是正信号,"讨厌"是负信号,两者都需要表达。tanh⁡\tanhtanh 以 0 为中心、左右对称,正负信息都能承载。

如果用反了会怎样?

  • 如果门用 tanh⁡\tanhtanh:输出范围 (−1,1)(-1,1)(−1,1),出现负数时"比例阀"的含义崩了------"-0.5 的比例"是什么意思?让 -50% 的信息通过?
  • 如果状态用 sigmoid:输出范围 (0,1)(0,1)(0,1),全是正数。隐藏状态的均值会逐渐偏向正数,经过多层叠加后分布越来越偏------这就是所谓的"covariate shift",大白话就是"数字越跑越歪"。

    上图一目了然:左图 sigmoid 输出在 0~1 之间,天然是"阀门刻度";右图 tanh⁡\tanhtanh 输出在 -1~1 之间,以 0 为中心对称,天然是"信息载体,可正可负"。

一句话记住:sigmoid 管"开多少",tanh 管"装什么"。前者是阀门刻度(0~1),后者是货物(可正可负)。


四、参数量统计:LSTM 到底比 RNN 多了多少参数

搞清楚参数量有两个好处:一是面试问到能秒答,二是实际选型时心里有数。

4.1 计算过程

LSTM 有 4 组权重矩阵,每组包含两个部分:

  • W∗W_*W∗:形状 H×(H+I)H \times (H+I)H×(H+I),连接 ht−1,xth_{t-1}, x_tht−1,xt 到门;
  • b∗b_*b∗:形状 HHH,偏置。

单组参数 = H(H+I)+H=H(H+I+1)H(H+I) + H = H(H+I+1)H(H+I)+H=H(H+I+1)。

四组合计 = 4H(H+I+1)4H(H+I+1)4H(H+I+1)。

4.2 与 RNN 对比

RNN 有 3 组权重矩阵(WxhW_{xh}Wxh、WhhW_{hh}Whh、WhyW_{hy}Why)加两组偏置(bhb_hbh、byb_yby),合计:

RNN 参数=I⋅H+H2+H⋅O+H+O\text{RNN 参数} = I \cdot H + H^2 + H \cdot O + H + ORNN 参数=I⋅H+H2+H⋅O+H+O

假如隐层维度 H=128H=128H=128,输入维度 I=100I=100I=100,输出维度 O=2O=2O=2:

模型 参数量 倍数
RNN 100×128+1282+128×2+128+2≈29.4K100\times128 + 128^2 + 128\times2 + 128 + 2 \approx 29.4\text{K}100×128+1282+128×2+128+2≈29.4K
LSTM 4×128×(128+100+1)≈117.2K4\times128\times(128+100+1) \approx 117.2\text{K}4×128×(128+100+1)≈117.2K 约 4×

LSTM 参数大约是 RNN 的 4 倍。这个倍数很好记------LSTM 有 4 个门(遗忘、输入、候选、输出),每个门的参数量和 RNN 的一个隐藏层差不多。如果用"四合一"矩阵 WWW(形状 4H×(H+I)4H \times (H+I)4H×(H+I)),参数量就是 4H(H+I)+4H=4H(H+I+1)4H(H+I) + 4H = 4H(H+I+1)4H(H+I)+4H=4H(H+I+1)------和分开算完全一样,只是存储方式不同。


五、实操:手写 lstm_forward,亲眼看到数据流动

光说不练假把式。下面用 NumPy 实现 LSTM 的单步前向计算,并打印每一步张量的形状------这是真正理解 LSTM 的"最后一公里"。

5.1 单步前向计算
python 复制代码
import numpy as np

def sigmoid(x):
    return 1.0 / (1.0 + np.exp(-x))

def lstm_step_forward(x_t, h_prev, c_prev, W, b):
    """
    LSTM 单个时间步的前向计算。
    参数:
      x_t:    (B, I)  当前输入
      h_prev: (B, H)  上一时刻隐藏状态
      c_prev: (B, H)  上一时刻细胞状态
      W:      (4H, H+I) 四合一权重矩阵
      b:      (4H,)     四合一偏置
    返回:
      h_t: (B, H)  新隐藏状态
      c_t: (B, H)  新细胞状态
      cache: dict  中间值(供反向传播使用)
    """
    H = h_prev.shape[1]
    # 1. 拼接输入
    concat = np.hstack([h_prev, x_t])          # (B, H+I)
    # 2. 一次矩阵乘法算出四个门的原始值
    gates_raw = np.dot(concat, W.T) + b        # (B, 4H)
    # 3. 切成四段
    f_raw = gates_raw[:, 0*H:1*H]              # (B, H)
    i_raw = gates_raw[:, 1*H:2*H]              # (B, H)
    c_raw = gates_raw[:, 2*H:3*H]              # (B, H)
    o_raw = gates_raw[:, 3*H:4*H]              # (B, H)
    # 4. 激活
    f = sigmoid(f_raw)                         # 遗忘门: (B, H), 0~1
    i = sigmoid(i_raw)                         # 输入门: (B, H), 0~1
    c_tilde = np.tanh(c_raw)                   # 候选状态: (B, H), -1~1
    o = sigmoid(o_raw)                         # 输出门: (B, H), 0~1
    # 5. 状态更新
    c_t = f * c_prev + i * c_tilde             # 新细胞状态: (B, H)
    h_t = o * np.tanh(c_t)                     # 新隐藏状态: (B, H)
    # 6. 保存中间值
    cache = (x_t, h_prev, c_prev, concat,
             f, i, c_tilde, o, c_t, h_t)
    return h_t, c_t, cache

代码虽短,但每一行都能在前面的公式里找到对应。关键步骤注释了编号,对照公式看一遍就通:

  1. 拼接 → 对应所有公式里的 ht−1,xth_{t-1}, x_tht−1,xt
  2. 矩阵乘法 → 对应 W∗⋅ht−1,xt+b∗W_* \cdot h_{t-1}, x_t + b_*W∗⋅ht−1,xt+b∗ 四合一版本
  3. 切片 → 四组门各取一段
  4. 激活 → 遗忘门/输入门/输出门走 sigmoid,候选状态走 tanh
  5. 更新 → ct=f⊙ct−1+i⊙c~tc_t = f \odot c_{t-1} + i \odot \tilde{c}_tct=f⊙ct−1+i⊙c~t,ht=o⊙tanh⁡(ct)h_t = o \odot \tanh(c_t)ht=o⊙tanh(ct)
5.2 维度验证

用一组小数值跑一遍,打印每个中间结果的形状:

python 复制代码
# 小模型:B=2, I=3, H=4
B, I, H = 2, 3, 4
x_t = np.random.randn(B, I)
h_prev = np.random.randn(B, H)
c_prev = np.random.randn(B, H)
W = np.random.randn(4*H, H+I) * 0.1
b = np.zeros(4*H)

h_t, c_t, cache = lstm_step_forward(x_t, h_prev, c_prev, W, b)

print("x_t     shape:", x_t.shape)       # (2, 3)
print("h_prev  shape:", h_prev.shape)    # (2, 4)
print("concat  shape:", cache[3].shape)  # (2, 7) = H+I
print("gates   shape:", (B, 4*H))        # (2, 16) = 4H
print("f/i/c̃/o  各 shape:", (B, H))      # (2, 4)
print("c_t     shape:", c_t.shape)       # (2, 4)
print("h_t     shape:", h_t.shape)       # (2, 4)

输出:

text 复制代码
x_t     shape: (2, 3)
h_prev  shape: (2, 4)
concat  shape: (2, 7)
gates   shape: (2, 16)
f/i/c̃/o  各 shape: (2, 4)
c_t     shape: (2, 4)
h_t     shape: (2, 4)

几个关键检查点:

  • concat 的维度是 H+I=4+3=7H+I = 4+3 = 7H+I=4+3=7 ✓
  • gates_raw 的维度是 4H=164H = 164H=16 ✓
  • 四个门的输出都是 (B,H)=(2,4)(B, H) = (2, 4)(B,H)=(2,4) ✓
  • ctc_tct 和 hth_tht 形状相同------隐藏状态的维度始终是 HHH,细胞状态作为"内部记忆"也是同样大小。
5.3 扩展到完整序列

单步搞定了,多步就是套一层循环。和 RNN 的 rnn_forward 思路完全一样:初始化 h0h_0h0 和 c0c_0c0(通常都是全零),然后逐时间步调用 lstm_step_forward,把每一步的 hth_tht 和 ctc_tct 传给下一步:

python 复制代码
def lstm_forward(x, W, b, h0=None, c0=None):
    """
    完整序列的 LSTM 前向。
    x: (B, S, I)  输入序列,S 为序列长度
    返回: h_seq (S, B, H), c_seq (S, B, H), caches
    """
    B, S, I = x.shape
    H = b.shape[0] // 4
    if h0 is None:
        h0 = np.zeros((B, H))
    if c0 is None:
        c0 = np.zeros((B, H))

    h_seq, c_seq, caches = [], [], []
    h_prev, c_prev = h0, c0
    for t in range(S):
        h_t, c_t, step_cache = lstm_step_forward(
            x[:, t, :], h_prev, c_prev, W, b)
        h_seq.append(h_t)
        c_seq.append(c_t)
        caches.append(step_cache)
        h_prev, c_prev = h_t, c_t

    h_seq = np.stack(h_seq, axis=0)   # (S, B, H)
    c_seq = np.stack(c_seq, axis=0)   # (S, B, H)
    return h_seq, c_seq, caches

和 RNN 前向的核心区别就三点:

  1. 多了一个细胞状态 ctc_tct 在时间步之间传递;
  2. 每一步计算用四个门而不是一个 tanh⁡\tanhtanh;
  3. 输出 hth_tht 取自 ot⊙tanh⁡(ct)o_t \odot \tanh(c_t)ot⊙tanh(ct),而不是直接 tanh⁡(zt)\tanh(z_t)tanh(zt)。

六、本章小结与下章预告

本章小结
知识点 一句话带走
LSTM 7 个公式 前 4 个并行为门控(ft,it,c~t,otf_t, i_t, \tilde{c}_t, o_tft,it,c~t,ot),后 3 个串行为状态更新(ct,htc_t, h_tct,ht)
四合一权重矩阵 WWW 形状 4H×(H+I)4H \times (H+I)4H×(H+I),一次矩阵乘法算四个门,再按行切片
遗忘门偏置初始化 通常初始化为 1(默认保留),而不是 0------防止训练初期丢失重要信息
sigmoid vs tanh sigmoid 管"开多少"(0~1 比例阀),tanh 管"装什么"(-1~1 信息载体)
参数量 4H(H+I+1)4H(H+I+1)4H(H+I+1),约为同尺寸 RNN 的 4 倍
单步实现 拼接 → 一次矩阵乘法 → 切片 → 四个激活 → 状态更新 → 输出
核心直觉(三段式流水线)

LSTM 的一个时间步可以浓缩为"三段式流水线":

  1. 并行阶段 :ht−1h_{t-1}ht−1 和 xtx_txt 拼接,同时算四个门------各自独立、互不干扰;
  2. 汇合阶段 :遗忘门擦旧(ft⊙ct−1f_t \odot c_{t-1}ft⊙ct−1),输入门写新(it⊙c~ti_t \odot \tilde{c}_tit⊙c~t),两者相加更新 ctc_tct;
  3. 输出阶段 :tanh⁡(ct)\tanh(c_t)tanh(ct) 压缩,输出门 oto_tot 过滤,产生 hth_tht。

记住这个流水线,代码就不会写错顺序。

下章预告:第 8 章《从零实现 LSTM 前后向传播》

公式和单步代码都有了,下一章做三件事:

  1. 手写完整的 lstm_backward------反向传播涉及 4 组权重矩阵,梯度如何沿 ctc_tct 和 hth_tht 两条路径分流汇聚;
  2. 用数值梯度给 LSTM 做"全身体检";
  3. 序列复制任务 ------输入 [3, 7, 2, 1, 9, 0, 0, ...],要求 10 步后输出同样的数字序列。这是检验长期记忆的经典任务,在同一张图上对比 RNN(loss 不降)和 LSTM(loss 稳步下降)。

公式背熟了,代码写出来了。下一章,让 LSTM 真正"跑起来"。


🧠 思考题与动手练习

思考题

  1. 为什么四个门的输入都是 ht−1,xth_{t-1}, x_tht−1,xt 的拼接,而不是分别用不同的输入?
  2. 如果遗忘门偏置 bfb_fbf 初始化为 -1(默认遗忘),训练初期会发生什么?为什么实践中不这样做?
  3. LSTM 的参数量公式是 4H(H+I+1)4H(H+I+1)4H(H+I+1),推导一下:每个 HHH、III、111 分别对应什么?
  4. 四合一权重矩阵 WWW 的形状是 4H×(H+I)4H \times (H+I)4H×(H+I),四种门的子矩阵排列顺序有要求吗?能不能把遗忘门放在第三段、输入门放在第一段?

动手练习

  1. 把第五节代码中的 HHH 从 4 改成 128,III 改成 100,打印 WWW 的形状并验算参数量是否为 4×128×(128+100+1)=1172484\times128\times(128+100+1)=1172484×128×(128+100+1)=117248;
  2. lstm_step_forward 中故意把 ctc_tct 的更新顺序写反------先算 hth_tht 再算 ctc_tct------看看运行会不会报错、报什么错;
  3. 改造第五节代码,把遗忘门偏置 bfb_fbf 初始化为 1.0(其余偏置为 0),跑一遍前向,观察 ftf_tft 的初始值是否接近 0.73(sigmoid(1.0));
  4. 用纸笔把 7 个公式完整默写一遍,标出每个公式的输入输出维度。

📌 下篇预告:第八章《从零实现 LSTM 前后向传播》------手写反向传播、数值梯度体检、序列复制任务实战对比。下篇见!

本文为原创,遵循 CC 4.0 BY-SA 版权协议,转载需附原文链接。


相关推荐
满怀冰雪2 小时前
15-Paddle 高层 API 入门:paddle.Model 的训练与评估流程
人工智能·python·深度学习·机器学习·paddle
妍妍爱学习16 小时前
自动编码器与变分自动编码器:同一脉络下的进化
深度学习·无监督学习·概率模型·自动编码器·变分自动编码器
赤羽尾风16 小时前
NumPy快速入门
python·numpy
小白说大模型17 小时前
人工智能:深度学习中的卷积神经网络(CNN)实战应用
人工智能·深度学习·cnn
Keepreal49617 小时前
《从零构建大模型》——从头实现GPT模型进行文本生成
gpt·深度学习·llm
卡梅德生物科技小能手17 小时前
卡梅德生物科普 TPBG(5T4)
经验分享·深度学习·生活
m沐沐19 小时前
【机器学习】DBSCAN聚类算法——原理、参数调优与实战
人工智能·python·深度学习·算法·机器学习·聚类·dbscan
手写码匠20 小时前
华为云征文|DeepSeek-R1 智能问数 Agent 实战:Flexus X 实例 + Dify 构建企业级 Text-to-SQL 查询助手
人工智能·深度学习·算法·aigc
donoot20 小时前
《大话文渊慧典》:六
人工智能·深度学习·aigc·ppstructure·文渊慧典·大话系列