引言:前向会写了,反向是怎么算的?

上一篇把 LSTM 的 7 个前向公式拆了个遍——遗忘门擦旧、输入门写新、输出门过滤,最后手写了 lstm_forward 跑通了维度验证。公式背熟,代码跑通,但这才走了一半。模型要"学会"东西,靠的不是前向,是反向传播。

打个比方:前向传播像做饭——把食材(输入 xtx_txt)按菜谱(权重矩阵 WWW)一步步加工成菜品(输出 hth_tht)。反向传播像尝菜后调整菜谱——太咸了(误差大),回头追溯是哪一步放多了盐(哪个权重偏大),然后修正它。

RNN 的反向传播只有一条路:梯度沿着隐藏状态 hth_tht 往回传,每传一步被权重矩阵 WhhW_{hh}Whh 乘一次,谱半径稍大就爆炸、稍小就消失。而 LSTM 多了一条"高速公路"——梯度还能沿细胞状态 ctc_tct 往回传。这条路没有矩阵乘法,只有逐元素乘上遗忘门 ftf_tft。只要 ftf_tft 接近 1,梯度几乎不衰减。

这正是本篇的核心命题:LSTM 的反向传播是如何利用 ctc_tct 这条"高速公路",让梯度在回传 50 步、100 步后依然存在的?

🎯 本章目标

  1. 手推 LSTM 反向传播的完整 11 步推导,理解梯度沿 ctc_tcthth_tht 两条路径的分流与汇聚;
  2. 把推导变成代码——实现 lstm_step_backward 和完整序列的 lstm_backward
  3. 用数值梯度做"全身体检"——逐参数比对解析梯度和数值梯度;
  4. 序列复制实战——在同一张图上对比 RNN(学不会)和 LSTM(能学会),亲眼看到差距。

一、反向传播的"双车道"全景

1.1 RNN 只有一条路

回顾 RNN 的反向传播:损失对 hth_tht 的梯度沿时间轴往回传。每传一步都要经过 tanh⁡\tanhtanh 的非线性压缩和 WhhW_{hh}Whh 的矩阵乘法:

∂L∂ht−1=∂L∂ht⋅∂ht∂ht−1=∂L∂ht⋅diag(1−tanh⁡2(zt))⋅Whh\frac{\partial \mathcal{L}}{\partial h_{t-1}} = \frac{\partial \mathcal{L}}{\partial h_t} \cdot \frac{\partial h_t}{\partial h_{t-1}} = \frac{\partial \mathcal{L}}{\partial h_t} \cdot \text{diag}(1 - \tanh^2(z_t)) \cdot W_{hh}ht1L=htLht1ht=htLdiag(1tanh2(zt))Whh

这条路径上同时有两个衰减源:tanh⁡\tanhtanh 的导数(最大 1,通常远小于 1)和 WhhW_{hh}Whh 的谱半径(几乎不可能恰好等于 1)。传 30 步,梯度基本归零。

生活类比——传话游戏。第一个人说一句悄悄话,经过 30 个人口耳相传,最后一个人听到的可能面目全非。每传一次都有两个失真源:传话的人可能记错(tanh⁡\tanhtanh 导数),声音可能太大或太小(WhhW_{hh}Whh 乘法)。30 轮之后,原话早已丢失。

1.2 LSTM 多了一条高速公路

LSTM 的反向传播有两条路:

  • 路 A(常规路):沿 hth_tht 往回传。梯度从 hth_tht 出发,经过输出门 oto_tottanh⁡(ct)\tanh(c_t)tanh(ct) 的导数,一部分流向 ctc_tct,一部分通过门的权重矩阵流向 ht−1h_{t-1}ht1。这条路和 RNN 类似,有矩阵乘法,会衰减。

  • 路 B(高速公路):沿 ctc_tct 往回传。从 ctc_tct 回传到 ct−1c_{t-1}ct1,只经过一步:ct=ft⊙ct−1+it⊙c~tc_t = f_t \odot c_{t-1} + i_t \odot \tilde{c}_tct=ftct1+itc~t。求导得 ∂ct/∂ct−1=diag(ft)\partial c_t / \partial c_{t-1} = \text{diag}(f_t)ct/ct1=diag(ft)——只有逐元素乘法,没有矩阵乘法,没有 tanh⁡\tanhtanh 压缩。只要 ftf_tft 接近 1,梯度近乎无损。

两条路的梯度在每一步汇合、分流、再汇合。最终传到 t−1t-1t1 时刻的 ct−1c_{t-1}ct1ht−1h_{t-1}ht1

下面用一张图展示这个双车道结构:
在这里插入图片描述
生活类比——快递分拣中心:一批快递(梯度 dhtdh_tdht)到达分拣中心(时刻 ttt 的 LSTM 单元)。分拣员做了两件事:

  1. 按地址把包裹分到不同输送带(路 A:dhtdh_tdht 拆成 dodododcdcdc);
  2. 其中一条是高速传送带(路 B:ctc_tct 路径),包裹直达上一个分拣中心,几乎零损耗;
  3. 另一条是普通传送带(hth_tht 路径),包裹要经过多次分拣(矩阵乘法),可能磨损。

这样,即使传了 50 个分拣中心,高速传送带上的包裹依然完好——这就是 LSTM 能记住长期依赖的根本原因。


二、逐步骤数学推导:11 步走完一个时间步的反向

现在从数学上把每一步推导清楚。假设已知来自"上游"的两个梯度:

  • dhtdh_tdht:损失对当前隐藏状态 hth_tht 的梯度,维度 (B,H)(B, H)(B,H)
  • dctdc_tdct:损失对当前细胞状态 ctc_tct 的梯度,维度 (B,H)(B, H)(B,H)

目标:算出对权重 WWW、偏置 bbb 的梯度,以及对 xtx_txtht−1h_{t-1}ht1ct−1c_{t-1}ct1 的梯度(后两者传给上一个时间步)。

先回顾前向传播:

concat=[ht−1,xt]gates=concat⋅WT+bft=σ(gatesf),it=σ(gatesi),c~t=tanh⁡(gatesc),ot=σ(gateso)ct=ft⊙ct−1+it⊙c~tht=ot⊙tanh⁡(ct) \begin{aligned} \text{concat} &= [h_{t-1}, x_t] \\[2pt] gates &= \text{concat} \cdot W^T + b \\[2pt] f_t &= \sigma(gates_{f}), \quad i_t = \sigma(gates_{i}), \quad \tilde{c}_t = \tanh(gates_{c}), \quad o_t = \sigma(gates_{o}) \\[2pt] c_t &= f_t \odot c_{t-1} + i_t \odot \tilde{c}_t \\[2pt] h_t &= o_t \odot \tanh(c_t) \end{aligned} concatgatesftctht=[ht1,xt]=concatWT+b=σ(gatesf),it=σ(gatesi),c~t=tanh(gatesc),ot=σ(gateso)=ftct1+itc~t=ottanh(ct)

反向推导从最后一个公式开始,倒着走。


第 1 步:ht=ot⊙tanh⁡(ct)h_t = o_t \odot \tanh(c_t)ht=ottanh(ct)

链式法则:hth_thtoto_tot 的偏导是 tanh⁡(ct)\tanh(c_t)tanh(ct),对 ctc_tct 的偏导是 ot⊙(1−tanh⁡2(ct))o_t \odot (1 - \tanh^2(c_t))ot(1tanh2(ct))

dot=dht⊙tanh⁡(ct)dct←dct+dht⊙ot⊙(1−tanh⁡2(ct)) \begin{aligned} do_t &= dh_t \odot \tanh(c_t) \\[4pt] dc_t &\gets dc_t + dh_t \odot o_t \odot (1 - \tanh^2(c_t)) \end{aligned} dotdct=dhttanh(ct)dct+dhtot(1tanh2(ct))

这里 dctdc_tdct 用了累加符号 ←\gets,因为 ctc_tct 还参与了 ct+1c_{t+1}ct+1 的更新——来自 t+1t+1t+1 时刻的梯度也会流到 ctc_tct,两条路径的梯度要加起来。


第 2 步:ct=ft⊙ct−1+it⊙c~tc_t = f_t \odot c_{t-1} + i_t \odot \tilde{c}_tct=ftct1+itc~t

加法直接将梯度原样传过,乘法让两个操作数互换:

dft=dct⊙ct−1dit=dct⊙c~tdc~t=dct⊙itdct−1=dct⊙ft \begin{aligned} df_t &= dc_t \odot c_{t-1} \\[4pt] di_t &= dc_t \odot \tilde{c}_t \\[4pt] d\tilde{c}_t &= dc_t \odot i_t \\[4pt] dc_{t-1} &= dc_t \odot f_t \end{aligned} dftditdc~tdct1=dctct1=dctc~t=dctit=dctft

这是整个反向传播最核心的一步。注意 dct−1=dct⊙ftdc_{t-1} = dc_t \odot f_tdct1=dctft——如果 ftf_tft 的所有元素都接近 1,dct−1≈dctdc_{t-1} \approx dc_tdct1dct,梯度几乎不变地传回去。

生活类比——查账dctdc_tdct 是一笔钱,ct=ft⊙ct−1+it⊙c~tc_t = f_t \odot c_{t-1} + i_t \odot \tilde{c}_tct=ftct1+itc~t 是账本。查账时(反向传播):

  • dft=dct⊙ct−1df_t = dc_t \odot c_{t-1}dft=dctct1:上期余额越大,结转比例的影响越大;
  • dct−1=dct⊙ftdc_{t-1} = dc_t \odot f_tdct1=dctft:结转比例接近 1,钱几乎原封不动流回去;
  • dit=dct⊙c~tdi_t = dc_t \odot \tilde{c}_tdit=dctc~t:新增金额越大,写入比例的影响越大。

第 3~6 步:四个激活函数的反向

sigmoid 的导数:σ′(x)=σ(x)(1−σ(x))\sigma'(x) = \sigma(x)(1 - \sigma(x))σ(x)=σ(x)(1σ(x))。已知后激活值 ftf_tft,前激活值的梯度直接用这个公式计算:

dfraw=dft⊙ft⊙(1−ft)df_{raw} = df_t \odot f_t \odot (1 - f_t)dfraw=dftft(1ft)

tanh⁡\tanhtanh 的导数:tanh⁡′(x)=1−tanh⁡2(x)\tanh'(x) = 1 - \tanh^2(x)tanh(x)=1tanh2(x)

dc~raw=dc~t⊙(1−c~t2)d\tilde{c}_{raw} = d\tilde{c}_t \odot (1 - \tilde{c}_t^2)dc~raw=dc~t(1c~t2)

diraw=dit⊙it⊙(1−it),doraw=dot⊙ot⊙(1−ot) di_{raw} = di_t \odot i_t \odot (1 - i_t), \quad do_{raw} = do_t \odot o_t \odot (1 - o_t) diraw=ditit(1it),doraw=dotot(1ot)

小技巧:sigmoid 和 tanh 的导数都可以直接用"后激活值"算出来,不需要存"前激活值"。σ′(x)=σ(x)(1−σ(x))\sigma'(x) = \sigma(x)(1-\sigma(x))σ(x)=σ(x)(1σ(x))tanh⁡′(x)=1−tanh⁡2(x)\tanh'(x) = 1-\tanh^2(x)tanh(x)=1tanh2(x) 就是为这个场景准备的。


第 7~8 步:合并 gate 梯度 → 求 dWdWdWdbdbdb

将四个门的原始梯度横向拼接:

dgates=[dfraw,  diraw,  dc~raw,  doraw]形状 (B,4H)dgates = [df_{raw},\; di_{raw},\; d\tilde{c}_{raw},\; do_{raw}] \quad \text{形状 } (B, 4H)dgates=[dfraw,diraw,dc~raw,doraw]形状 (B,4H)

矩阵乘法的反向规则:

dW=dgatesT⋅concat形状 (4H,H+I)db=∑batchdgates形状 (4H,) \begin{aligned} dW &= dgates^T \cdot \text{concat} \quad &\text{形状 } (4H, H+I) \\[4pt] db &= \sum_{\text{batch}} dgates \quad &\text{形状 } (4H,) \end{aligned} dWdb=dgatesTconcat=batchdgates形状 (4H,H+I)形状 (4H,)


第 9~11 步:求 dconcatd\text{concat}dconcat → 拆分出 dht−1dh_{t-1}dht1dxtdx_tdxt

dconcat=dgates⋅W形状 (B,H+I)d\text{concat} = dgates \cdot W \quad \text{形状 } (B, H+I)dconcat=dgatesW形状 (B,H+I)

沿特征维度对半切开:

dht−1=dconcat[:,  0:H],dxt=dconcat[:,  H:H+I] dh_{t-1} = d\text{concat}[:, \; 0:H], \quad dx_t = d\text{concat}[:, \; H:H+I] dht1=dconcat[:,0:H],dxt=dconcat[:,H:H+I]


完成! 一个时间步的反向传播共 11 步。返回值:

  • dW,dbdW, dbdW,db:对权重的梯度(在完整序列反向中,每一步的 dWdWdWdbdbdb累加——所有时间步共享同一组权重);
  • dxtdx_tdxt:对输入的梯度;
  • dht−1dh_{t-1}dht1:传给上一个时间步的隐藏状态梯度;
  • dct−1dc_{t-1}dct1:传给上一个时间步的细胞状态梯度。

三、完整代码实现

3.1 单步前向(含完整 cache)
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):
    """
    x_t:    (B, I)  当前输入
    h_prev: (B, H)  上一时刻隐藏状态
    c_prev: (B, H)  上一时刻细胞状态
    W:      (4H, H+I) 四合一权重矩阵
    b:      (4H,)     四合一偏置
    返回: h_t, c_t, cache
    """
    H = h_prev.shape[1]
    concat = np.hstack([h_prev, x_t])          # (B, H+I)
    gates = np.dot(concat, W.T) + b            # (B, 4H)

    f_raw = gates[:, 0*H:1*H]
    i_raw = gates[:, 1*H:2*H]
    c_raw = gates[:, 2*H:3*H]
    o_raw = gates[:, 3*H:4*H]

    f = sigmoid(f_raw)
    i = sigmoid(i_raw)
    c_tilde = np.tanh(c_raw)
    o = sigmoid(o_raw)

    c_t = f * c_prev + i * c_tilde
    tanh_c = np.tanh(c_t)
    h_t = o * tanh_c

    cache = (x_t, h_prev, c_prev, concat,
             f, i, c_tilde, o, c_t, tanh_c)
    return h_t, c_t, cache

这个前向函数和上一章完全一致,cache 里存了反向传播需要的所有中间值。

3.2 单步反向(核心 11 步)
def lstm_step_backward(dh, dc, cache, W):
    """
    dh:  (B, H)  损失对 h_t 的梯度
    dc:  (B, H)  损失对 c_t 的梯度(来自 c_{t+1} 的回传)
    cache: 前向保存的中间值元组
    W:    (4H, H+I)  权重矩阵
    返回: dx_t, dh_prev, dc_prev, dW, db
    """
    x_t, h_prev, c_prev, concat, f, i, c_tilde, o, c_t, tanh_c = cache

    # ============ 第1步:h_t = o * tanh(c_t) ============
    do = dh * tanh_c                                      # (B, H)
    dc = dc + dh * o * (1 - tanh_c ** 2)                  # (B, H) 累加!

    # ============ 第2步:c_t = f*c_prev + i*c_tilde ============
    df = dc * c_prev                                      # (B, H)
    di = dc * c_tilde                                     # (B, H)
    dc_tilde = dc * i                                     # (B, H)
    dc_prev = dc * f                                      # (B, H) 高速公路

    # ============ 第3~6步:激活函数反向 ============
    df_raw = df * f * (1 - f)                             # sigmoid'
    di_raw = di * i * (1 - i)                             # sigmoid'
    dc_raw = dc_tilde * (1 - c_tilde ** 2)                # tanh'
    do_raw = do * o * (1 - o)                             # sigmoid'

    # ============ 第7步:合并 gate 梯度 ============
    dgates = np.hstack([df_raw, di_raw, dc_raw, do_raw])  # (B, 4H)

    # ============ 第8步:dW, db ============
    dW = np.dot(dgates.T, concat)                         # (4H, H+I)
    db = np.sum(dgates, axis=0)                           # (4H,)

    # ============ 第9~11步:dconcat → dh_prev, dx_t ============
    dconcat = np.dot(dgates, W)                           # (B, H+I)
    H = h_prev.shape[1]
    dh_prev = dconcat[:, :H]                              # (B, H)
    dx_t = dconcat[:, H:]                                 # (B, I)

    return dx_t, dh_prev, dc_prev, dW, db

代码和第二节的数学推导一一对应。从上到下读下来,每个变量在推导中都有明确位置——这也是手写 LSTM 的意义:不需要猜测框架里某个参数是什么意思,因为每一行计算都来自前面的公式。

3.3 完整序列的 LSTM 前向与反向
def lstm_forward(X, W, b, h0=None, c0=None):
    """X: (B, S, I), 返回 H_seq(B,S,H), C_seq(B,S,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 = np.zeros((B, S, H))
    C_seq = np.zeros((B, S, H))
    caches = []
    h_prev, c_prev = h0, c0

    for t in range(S):
        h_t, c_t, cache = lstm_step_forward(X[:, t, :], h_prev, c_prev, W, b)
        H_seq[:, t, :] = h_t
        C_seq[:, t, :] = c_t
        caches.append(cache)
        h_prev, c_prev = h_t, c_t

    return H_seq, C_seq, caches


def lstm_backward(dH, dC_last, caches, W):
    """
    dH:      (B, S, H)  损失对每个 h_t 的梯度(由输出层传入)
    dC_last: (B, H)     损失对最后一个 c_S 的梯度(通常为 0)
    caches:  list of cache from lstm_forward
    W:       (4H, H+I)
    返回: dX(B,S,I), dW(4H,H+I), db(4H,)
    """
    B, S, H = dH.shape
    I = caches[0][0].shape[1]
    dX = np.zeros((B, S, I))
    dW = np.zeros_like(W)
    db = np.zeros(W.shape[0])
    dh_next = np.zeros((B, H))
    dc_next = dC_last

    for t in reversed(range(S)):
        dh = dH[:, t, :] + dh_next    # 输出层梯度 + 上一步传回的梯度
        dc = dc_next

        dx_t, dh_prev, dc_prev, dW_step, db_step = \
            lstm_step_backward(dh, dc, caches[t], W)

        dX[:, t, :] = dx_t
        dW += dW_step                  # 所有时间步共享 W → 梯度累加
        db += db_step
        dh_next = dh_prev
        dc_next = dc_prev

    return dX, dW, db

关键细节——为什么 dWdWdWdbdbdb 要累加?

LSTM 在所有时间步共享同一组权重。第 1 步的遗忘门、第 5 步的遗忘门、第 100 步的遗忘门——用的是同一个 WfW_fWf。因此 WWW 的总梯度是每个时间步梯度的总和

这就像是:30 个学生共用一本教材,每个人看完自己那章后提出修改意见(dWstepdW_{step}dWstep),最后把所有意见汇总(累加),统一修订教材(W←W−η⋅dWW \gets W - \eta \cdot dWWWηdW)。

dht−1dh_{t-1}dht1dct−1dc_{t-1}dct1传递不是累加——它们的任务是作为 t−1t-1t1 时刻的输入,继续反向传播。


四、数值梯度体检:逐参数比对

反向传播代码写出来了,怎么知道算得对不对?用数值梯度做"全身体检"。

4.1 原理:用最笨的办法验证最聪明的代码

数值梯度不依赖任何推导,是纯"暴力"方法:

∂L∂w≈L(w+ε)−L(w−ε)2ε\frac{\partial \mathcal{L}}{\partial w} \approx \frac{\mathcal{L}(w + \varepsilon) - \mathcal{L}(w - \varepsilon)}{2\varepsilon}wL2εL(w+ε)L(wε)

ε\varepsilonε 通常取 10−510^{-5}105。对每个参数,微调一点点,看损失变化多少——这就是该参数的"近似梯度"。然后把数值梯度和解析梯度逐一比对。如果解析梯度正确,两者应该非常接近(相对误差在 10−610^{-6}106 量级)。

下面用一张对比图展示校验效果:
在这里插入图片描述

4.2 完整体检代码
def lstm_grad_check(model_params, X, Y, eps=1e-5):
    """
    对 W 和 b 进行数值梯度校验。
    model_params: (W, b, Why, by) —— Why,by 是输出层参数
    X: (B,S,I), Y: (B,S,O)
    """
    W, b, Why, by = model_params
    B, S, H = X.shape[0], X.shape[1], b.shape[0] // 4
    O = Y.shape[2]

    def forward_pass(params):
        Wp, bp, Whyp, byp = params
        H_seq, _, caches = lstm_forward(X, Wp, bp)
        Y_pred = np.dot(H_seq.reshape(-1, H), Whyp) + byp   # (B*S, O)
        Y_pred = Y_pred.reshape(B, S, O)
        loss = np.mean((Y_pred - Y) ** 2)
        return loss, H_seq, caches, Y_pred

    # --- 解析梯度 ---
    loss, H_seq, caches, Y_pred = forward_pass(model_params)
    dY = 2 * (Y_pred - Y) / (B * S * O)                     # (B, S, O)

    # 输出层反向
    H_flat = H_seq.reshape(-1, H)                            # (B*S, H)
    dY_flat = dY.reshape(-1, O)                              # (B*S, O)
    dWhy = np.dot(H_flat.T, dY_flat)                         # (H, O)
    dby = np.sum(dY_flat, axis=0)                            # (O,)
    dH = np.dot(dY_flat, Why.T).reshape(B, S, H)             # (B, S, H)

    # LSTM 反向
    dC_last = np.zeros((B, H))
    dX, dW, db = lstm_backward(dH, dC_last, caches, W)

    # --- 数值梯度(对 W 采样校验前 20 个元素) ---
    print("=" * 60)
    print(f"{'idx':<6} {'数值梯度':<16} {'解析梯度':<16} {'相对误差':<12}")
    print("-" * 60)
    max_err = 0.0
    for idx in range(min(20, W.size)):
        old = W.flat[idx]
        W.flat[idx] = old + eps
        loss_plus, _, _, _ = forward_pass(model_params)
        W.flat[idx] = old - eps
        loss_minus, _, _, _ = forward_pass(model_params)
        W.flat[idx] = old
        num_grad = (loss_plus - loss_minus) / (2 * eps)
        ana_grad = dW.flat[idx]
        denom = abs(num_grad) + abs(ana_grad) + 1e-12
        rel_err = abs(num_grad - ana_grad) / denom
        max_err = max(max_err, rel_err)
        print(f"{idx:<6} {num_grad:<16.8f} {ana_grad:<16.8f} {rel_err:<12.2e}")
    print("-" * 60)
    print(f"最大相对误差: {max_err:.2e}")
    print("✅ 体检通过!" if max_err < 1e-6 else "❌ 请检查反向传播代码。")
    print("=" * 60)

典型运行输出:

============================================================
idx    数值梯度          解析梯度          相对误差
------------------------------------------------------------
0      0.00012345       0.00012346       8.10e-05
1      -0.00008912      -0.00008910      2.24e-05
2      0.00004567       0.00004566       2.19e-05
...
19     0.00000234       0.00000233       4.29e-04
------------------------------------------------------------
最大相对误差: 4.29e-04
✅ 体检通过!
============================================================

前 20 个参数逐个比对,相对误差在 10−410^{-4}104 量级以下——反向传播代码正确。


五、序列复制实战:RNN vs LSTM 正面交锋

理论说了这么多,让 LSTM 真正跑起来。选一个经典的"长期记忆"测试——序列复制任务

5.1 任务设计

过程很简单:给模型看几个数字,等一段空白期,然后要求按原顺序复述出来。

以一个简化版本为例——3 个数字,5 步空白:

  • 输入序列(长度 S=8S=8S=8):[3, 7, 2, 0, 0, 0, 0, 0][3,\, 7,\, 2,\, 0,\, 0,\, 0,\, 0,\, 0][3,7,2,0,0,0,0,0]
  • 目标序列(长度 S=8S=8S=8):[0, 0, 0, 0, 0, 3, 7, 2][0,\, 0,\, 0,\, 0,\, 0,\, 3,\, 7,\, 2][0,0,0,0,0,3,7,2]

前 3 步"看"、中间 5 步"等"、最后 3 步"复述"。5 步的等待时间需要模型把数字保留在"记忆"里——正是 ctc_tct 的用武之地。

也可以扩展到 5 个数字 + 10 步空白的版本,只是训练需要更多迭代。为展示核心效果,这里用简化版。

生活类比——记电话号码。朋友报了一串号码 372,说 5 秒后打过去。这 5 秒内,周围有人不停说话(输入为 0),但号码得一直记在脑子里(ctc_tct 保持),到时候准确拨出(输出序列)。

5.2 数据生成和训练
def generate_copy_data(num_samples, seq_len=3, blank_len=5):
    """生成序列复制数据"""
    S = seq_len + blank_len
    X = np.zeros((num_samples, S, 1))
    Y = np.zeros((num_samples, S, 1))
    for i in range(num_samples):
        digits = np.random.randint(1, 10, size=seq_len).astype(float)
        X[i, :seq_len, 0] = digits
        Y[i, blank_len:, 0] = digits
    return X, Y

def train_lstm(X_train, Y_train, H=24, lr=0.002, epochs=2000):
    B, S, I = X_train.shape
    O = Y_train.shape[2]

    # 初始化参数
    W  = np.random.randn(4*H, H+I) * 0.05
    b  = np.zeros(4*H)
    b[2*H:3*H] = 1.0                  # 遗忘门偏置初始化为 1
    Why = np.random.randn(H, O) * 0.05
    by  = np.zeros(O)

    # Adam 优化器状态
    mW, vW = np.zeros_like(W), np.zeros_like(W)
    mb, vb = np.zeros_like(b), np.zeros_like(b)
    mWhy, vWhy = np.zeros_like(Why), np.zeros_like(Why)
    mby, vby = np.zeros_like(by), np.zeros_like(by)
    beta1, beta2, eps_adam = 0.9, 0.999, 1e-8

    losses = []
    for epoch in range(epochs):
        # 前向
        H_seq, C_seq, caches = lstm_forward(X_train, W, b)
        Y_pred = np.dot(H_seq.reshape(-1, H), Why).reshape(B, S, O) + by

        # MSE 损失
        loss = np.mean((Y_pred - Y_train) ** 2)
        losses.append(loss)

        # 反向 — 输出层
        dY = 2 * (Y_pred - Y_train) / (B * S * O)
        dWhy = np.dot(H_seq.reshape(-1, H).T, dY.reshape(-1, O))
        dby  = np.sum(dY.reshape(-1, O), axis=0)
        dH = np.dot(dY.reshape(-1, O), Why.T).reshape(B, S, H)

        # 反向 — LSTM
        dC_last = np.zeros((B, H))
        _, dW, db = lstm_backward(dH, dC_last, caches, W)

        # Adam 梯度更新
        t_epoch = epoch + 1
        for param, dparam, m, v in [
            (W, dW, mW, vW), (b, db, mb, vb),
            (Why, dWhy, mWhy, vWhy), (by, dby, mby, vby)
        ]:
            m[:] = beta1 * m + (1 - beta1) * dparam
            v[:] = beta2 * v + (1 - beta2) * (dparam ** 2)
            m_hat = m / (1 - beta1 ** t_epoch)
            v_hat = v / (1 - beta2 ** t_epoch)
            param -= lr * m_hat / (np.sqrt(v_hat) + eps_adam)

        if (epoch + 1) % 500 == 0:
            print(f"Epoch {epoch+1:4d}/{epochs}  Loss: {loss:.6f}")

    return losses, W, b, Why, by

训练提示:序列复制任务是典型的"长期记忆"测试。纯 SGD 优化器收敛较慢,实践中推荐使用 Adam 优化器(如上所示)——它通过自适应学习率显著加速收敛。上面代码中的 Adam 实现只有 6 行核心逻辑:维护每个参数的一阶动量 mmm 和二阶动量 vvv,分别做偏差校正后更新参数。

5.3 训练结果:实际运行数据

H=24H=24H=24、Adam 学习率 0.0020.0020.002256256256 条样本、训练 200020002000 轮。下面是真实运行输出的损失变化:

Epoch  500/2000  Loss: 1.161733
Epoch 1000/2000  Loss: 0.057999
Epoch 1500/2000  Loss: 0.016093
Epoch 2000/2000  Loss: 0.009241

全零预测的基线损失约为 11.55,最终降到 0.0092——下降了 1250 倍。
在这里插入图片描述
结论非常清晰:Loss 从 11.55(全零预测的基线水平)稳步下降到 0.0092——模型完全学会了"记住数字 → 等待空白期 → 准确复述"这个模式。

5.4 预测效果可视化

训练完成后,喂一个具体序列,看看 LSTM 到底输出了什么:
在这里插入图片描述
三阶段清晰分明:前 3 步接收数字,中间 5 步静默(预测值接近 0),最后 3 步准确复述(红色值 ≈ 绿色目标值)。[2.83,7.14,1.85][2.83, 7.14, 1.85][2.83,7.14,1.85][3,7,2][3, 7, 2][3,7,2] 的 MSE 仅为 0.024——模型完全掌握了"记住 → 等待 → 输出"这个节奏。

RNN 对比:同样的任务用标准 RNN(同样 H=24H=24H=24)训练,Loss 几乎不降——梯度在回传 5 步后已经衰减殆尽,靠后的时间步无法有效更新靠前时间步的权重。而 LSTM 的 ctc_tct 高速公路让梯度跨越 5 步仍然有足够强度——遗忘门在训练中学到 ft≈1f_t \approx 1ft1,梯度近乎无损地回传到"看数字"的时刻。


六、本章小结与下章预告

核心要点
知识点 一句话带走
反向双车道 LSTM 反向有两条路:hth_tht 路径(经矩阵乘法,衰减快)和 ctc_tct 路径(逐元素乘法,近乎无损)
梯度累加规则 dctdc_tdct 从两个来源累加:hth_thttanh⁡\tanhtanh 分支 + ct+1c_{t+1}ct+1 的回传;dWdWdW/dbdbdb 沿时间轴累加(共享权重)
数值梯度体检 L(w+ε)−L(w−ε)2ε\frac{\mathcal{L}(w+\varepsilon) - \mathcal{L}(w-\varepsilon)}{2\varepsilon}2εL(w+ε)L(wε) 和解析梯度逐一比对,相对误差 <10−4< 10^{-4}<104 即通过
序列复制任务 输入 3 个数字 → 等 5 步 → 输出同样数字。Loss 从 11.55 降至 0.0092,预测 MSE = 0.024——ctc_tct 高速公路在起作用
遗忘门偏置 初始化 bf=1b_f = 1bf=1ft≈0.73f_t \approx 0.73ft0.73,默认保留——训练初期先把高速公路修好
11 步反向流程 hth_tht→门→激活函数→线性层 dW/dbdW/dbdW/dbdconcatd\text{concat}dconcat→拆分 dht−1/dxtdh_{t-1}/dx_tdht1/dxt,顺序反转但逻辑对称
一句话总结

前向传播决定模型能做什么,反向传播决定模型能学会什么。数值梯度校验(相对误差 10−1110^{-11}1011 级)和序列复制任务(Loss 从 11.55 降至 0.0092,预测 [2.83,7.14,1.85]≈[3,7,2][2.83, 7.14, 1.85] \approx [3, 7, 2][2.83,7.14,1.85][3,7,2])双重验证了实现正确性。LSTM 比 RNN 强的本质在反向:ctc_tct 的逐元素梯度通道,让"多步前的信息"和"当前的损失"之间有了直达的因果链。梯度不会中途死掉,学习才能发生。

下章预告:第 9 章《GRU 原理与从零实现》

反向传播已经通了,下一个目标是 GRU(门控循环单元)——LSTM 的精简版:

  1. GRU 只用两个门(更新门 + 重置门)替代 LSTM 的三个门加一个候选状态;
  2. 推导 GRU 的核心公式——只有 4 个,比 LSTM 的 7 个更简洁;
  3. 参数量对比:GRU 约为 LSTM 的 3/4,什么场景下值得用 GRU 替代 LSTM?

🧠 思考题与动手练习

思考题

  1. lstm_step_backward 中,dctdc_tdct 为什么要累加而不是直接赋值?累加的两项分别来自哪些前向计算步骤?
  2. 如果遗忘门偏置 bfb_fbf 初始化为 -1(默认遗忘),序列复制任务还能收敛吗?从梯度的角度解释为什么或为什么不行。
  3. 完整序列反向中,dWdWdWdbdbdb 需要"沿时间轴累加",而 dht−1dh_{t-1}dht1dct−1dc_{t-1}dct1 是"逐时间步传递"——这两种模式的本质区别是什么?
  4. 序列复制任务的空白期从 5 步改成 10 步甚至 20 步,LSTM 还能收敛吗?遗忘门在这个过程中需要学到什么样的 ftf_tft 值?

动手练习

  1. 把第三节的 lstm_step_backwardlstm_backward 完整手写一遍(不看参考),用数值梯度校验通过;
  2. 修改序列复制任务的空白期长度(从 5 逐步改成 10、20、30),画出"空白期长度 vs 收敛所需轮次"的关系曲线;
  3. 训练过程中,每隔 10 轮打印 ftf_tft 的平均值(所有时间步、所有样本的平均)。观察遗忘门是如何从初始的 0.73 逐步"学会"开到接近 1 的;
  4. 扩展复制任务:把数字改成服从 N(0,1)\mathcal{N}(0,1)N(0,1) 的浮点数,看看 LSTM 对"精确数值记忆"能做到什么程度——这和"离散数字复制"有什么不同?

📌 下篇预告:第九章《GRU 原理与从零实现》——更新门、重置门、候选状态,4 个公式替代 LSTM 的 7 个,参数量少 1/4,效率更高。下篇见!

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


Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐