【RustyML入门】3.7. 循环层

3.7. 循环层

RustyML 提供 3 种循环层:SimpleRNNLSTMGRU。它们都位于 rustyml::neural_network::layers::recurrent,并通过 prelude 重新导出。三者共用同一套约定:每个层都接收一个三维序列张量,沿时间轴从左到右跑一遍递归,返回最后的隐藏状态。

如果你熟悉 Keras,可以把这些层理解成把 return_sequences=False 固定死的 SimpleRNNLSTMGRU。这个固定设置会影响你如何堆叠这些层。搭建深层循环网络之前,请先读 [3.7.7 节](#3.7.7 节)。

3.7.1. 输入/输出约定

每个循环层都要求一个三维输入张量,形状为 (batch_size, timesteps, features),输出则是二维的 (batch_size, units)features 轴必须等于你传给构造函数的 input_dimunits 就是你设定的隐藏层宽度。递归会把时间步轴吃掉,所以它不会出现在输出里,只有最后一个隐藏状态 h_T 会留下来。

rust 复制代码
use ndarray::Array;
use rustyml::neural_network::layers::activation::Tanh;
use rustyml::neural_network::layers::recurrent::SimpleRNN;
use rustyml::neural_network::traits::Layer;

fn main() {
    // input_dim = 4 个每步特征,units = 3 个隐藏神经元
    let rnn = SimpleRNN::new(4, 3, Tanh::new()).unwrap();

    // (batch = 2, timesteps = 5, features = 4)
    let x = Array::zeros((2, 5, 4)).into_dyn();

    // predict 只跑递归,不记录反向传播所需的缓存
    let out = rnn.predict(&x).unwrap();

    // 时间步轴消失了:只有最后一个隐藏状态留了下来。
    println!("output shape: {:?}", out.shape()); // [2, 3] == (batch, units)
}
text 复制代码
output shape: [2, 3]

这 3 种层共用同一个构造函数签名:new(input_dim, units, activation) -> Result<Self, Error>。激活参数的类型是 impl Into<Activation>,你可以传一个 Activation 枚举变体,比如 Activation::Tanh,也可以传某个轻量的层包装,比如 Tanh::new()ReLU::new()Sigmoid::new()Linear::new()Softmax::new()。每一种包装都会转换成同一个枚举。完整的激活函数目录见 3.2. 全连接层与激活函数

input_dimunits 为 0 时,new 返回 Error::InvalidParameter。RustyML 没有 return_sequencesreturn_statebidirectional,也没有 cell 内 dropout 选项。每个层永远只返回最后的状态,永远沿时间正向推进,永远使用同一套稠密实现。

二维或四维输入是硬错误,不会被悄悄 reshape:只要输入不是三维,forwardpredict 就返回 Error::InvalidInput。层会缓存前向激活值供反向传播使用,所以在 forward 之前调用 backward 会返回 Error::NeuralNetwork(NnError::ForwardPassNotRun("SimpleRNN"))(对 LSTM、GRU 则是 "LSTM""GRU")。这些错误变体如何配合,见 1.6. 错误处理

RustyML 和 Keras 有一处不同:activation 参数只控制候选值或输出那一处非线性,不涉及门。在 SimpleRNN 里,这个激活作用于每一个隐藏状态;在 LSTMGRU 里,它作用于候选值,在 LSTM 里还会作用于输出门之前的细胞状态。

门本身永远使用 sigmoid。RustyML 不像 Keras 那样提供独立的 recurrent_activation 选项,门的非线性是固定的。Tanh 是默认的激活函数,也是几乎所有已发表架构采用的选择。

3.7.2. SimpleRNN 及门控单元存在的理由

SimpleRNN 就是教科书里的 Elman 递归:从全零隐藏状态出发,每个时间步用 2 个权重矩阵和 1 个偏置,把当前输入和上一步的隐藏状态混合起来:

text 复制代码
h_0 = 0
h_t = activation( x_t @ W + h_{t-1} @ U + b )     for t = 1..T
output = h_T

这里 W 是输入核 (input_dim, units)U 是递归核 (units, units)b 是偏置 (1, units)@ 是沿 batch 维度进行的矩阵乘法。

RustyML 用 Xavier/Glorot 均匀分布初始化 W,用正交矩阵(Gram-Schmidt)初始化 U。正交递归核是刻意为之,它让状态转移在初始化时保持范数不变,这是延缓下文要讲的梯度消失问题的一种省事办法。

这个问题就是梯度消失,以及它的对立面梯度爆炸。从 h_T 反向传播到 h_1,每一步都要把上游梯度乘上一个新的雅可比矩阵,大致是 grad_{t-1} = (activation'(h_t) * grad_t) @ U^T。把 T 步这样的运算串起来,相当于给一个矩阵求 T 次幂:如果它的有效量级小于 1,梯度就会指数衰减,网络学不到往前超过几步的依赖关系;如果量级大于 1,梯度反而会爆炸。

正交的 UU^T 这个因子保持范数不变,tanh' <= 1 又把乘积控制在有界范围内。即便如此,这个乘积在长序列上依然会趋向于零。这正是 SimpleRNN 只能应付短序列(至多几十步)、碰到长程依赖就失效的原因,也是 LSTM 和 GRU 存在的原因。

3.7.3. LSTM:一条加性的记忆高速路

LSTM 增加了第二个状态:细胞状态 c_t,它的更新是加性的。LSTM 用 3 个 sigmoid 门来决定写入什么、保留什么、读出什么。RustyML 把 4 个权重块并排融合存放,列的顺序和 Keras 一致,[input | forget | cell | output],记作 [i | f | g | o]

text 复制代码
i_t = sigmoid( x_t @ W_i + h_{t-1} @ U_i + b_i )     输入门(写入多少候选值)
f_t = sigmoid( x_t @ W_f + h_{t-1} @ U_f + b_f )     遗忘门(保留多少旧细胞状态)
g_t = act(     x_t @ W_g + h_{t-1} @ U_g + b_g )     候选值("cell gate")
o_t = sigmoid( x_t @ W_o + h_{t-1} @ U_o + b_o )     输出门(暴露多少细胞状态)
c_t = f_t * c_{t-1} + i_t * g_t                      细胞状态(加性更新)
h_t = o_t * act(c_t)                                 隐藏状态

这里 * 是逐元素乘法,act 是可配置的激活函数(默认 Tanh),同时作用于候选值和细胞状态。关键的一行是 c_t = f_t * c_{t-1} + i_t * g_t:它对 c_{t-1} 的梯度就是 f_t,一次逐元素乘法,没有反复的矩阵乘法。

当遗忘门打开,也就是 f 接近 1 时,细胞状态几乎无损地把梯度往回传,是一条近乎恒等的高速路,加性项则不断汇入这条路。记忆之所以留得住、梯度之所以流得动,是因为主干路径是加法,而不是反复的矩阵相乘。

RustyML 把遗忘门的偏置初始化为 1.0 ,其余偏置全部为零,这给这条记忆高速路一个先发优势。在训练把门调整成形之前,f 就已经偏向打开,于是细胞状态连同它的梯度,从第一个 epoch 起就能存活下来。集成测试 lstm_forget_bias_is_one_not_zero 就是用来检验这一行为的。LSTM 的参数量相当于 4 个门:param_count = 4 * (input_dim * units + units * units + units)

3.7.4. GRU:同一思路,门更少

GRU 保留了 LSTM 的加性混合思路,把输入门和遗忘门合并成一个更新 门,并去掉了独立的细胞状态,所以只需要 3 个权重块而不是 4 个。RustyML 把它们融合存放,顺序是 [update | reset | candidate],记作 [z | r | h],这个顺序和 Keras 一致:

text 复制代码
z_t = sigmoid( x_t @ W_z + h_{t-1} @ U_z + b_z )          更新门
r_t = sigmoid( x_t @ W_r + h_{t-1} @ U_r + b_r )          重置门
n_t = act( x_t @ W_h + (r_t * h_{t-1}) @ U_h + b_h )      候选值
h_t = z_t * h_{t-1} + (1 - z_t) * n_t                     隐藏状态

更新门 z_t 做的是一次凸混合。当 z 接近 1 时,层会把上一步的隐藏状态原封不动地拷过来,这是一条梯度高速路,和关闭的 LSTM 遗忘门效果一样。当 z 接近 0 时,层会用新的候选值替换掉上一步的状态。

这是 Keras 的约定。有些资料用的是它的补,从那类资料里搬过来的 z 需要翻转。测试 gru_update_gate_one_keeps_previous_hiddengru_update_gate_zero_takes_the_candidate 分别检验这两个极端,测试 gru_fused_kernel_first_block_is_the_update_gate 则检验列的顺序。

要留意重置门作用的确切位置:RustyML 在候选值的递归矩阵乘法之前 就算好了 r_t * h_{t-1},即 (r_t * h_{t-1}) @ U_h。这对应 Cho 等人最初的公式,也就是 Keras 的 reset_after=False,而不是 CuDNN 的 reset_after=True 变体(后者在矩阵乘法之后才施加重置,并且每个门需要 2 套偏置)。这里每个门只有一套偏置。

GRU 的 param_count = 3 * (input_dim * units + units * units + units),是同等宽度 LSTM 的四分之三。实践中 GRU 训练略快,在很多任务上和 LSTM 打平;当任务需要一段长而精准可控的记忆时,LSTM 有时表现更好。两者的输入核都用 Xavier/Glorot 初始化,用的是每个门各自的扇入 input_dim + units,而不是融合后的宽度;每个门的递归块都各自初始化成一个独立的正交矩阵。

3.7.5. 权重、形状,以及手动设置

可训练张量及其形状:

kernel recurrent_kernel bias 融合的列块
SimpleRNN (input_dim, units) (units, units) (1, units)
LSTM (input_dim, 4 * units) (units, 4 * units) (1, 4 * units) `[i
GRU (input_dim, 3 * units) (units, 3 * units) (1, 3 * units) `[z

把每个门融合进一个矩阵不只是好看:它让输入投影和递归投影在每个时间步都能跑成一个大 GEMM,而不是每个门各跑一个,这在缓存和 SIMD 上带来实打实的收益。你可以用 get_weights() 检视这些实时数组,它返回一个 LayerWeight::{SimpleRNN,LSTM,GRU} 值,里面带着借用的 kernelrecurrent_kernelbias 字段:

rust 复制代码
use rustyml::neural_network::layers::activation::Tanh;
use rustyml::neural_network::layers::layer_weight::LayerWeight;
use rustyml::neural_network::layers::recurrent::LSTM;
use rustyml::neural_network::traits::Layer;

fn main() {
    // input_dim = 4,units = 8。with_random_state 让初始化可复现。
    let lstm = LSTM::new(4, 8, Tanh::new()).unwrap().with_random_state(42);

    match lstm.get_weights() {
        LayerWeight::LSTM(w) => {
            // 4 个门并排融合在一起:宽度 == 4 * units。
            println!("kernel           {:?}", w.kernel.shape()); // [4, 32]
            println!("recurrent_kernel {:?}", w.recurrent_kernel.shape()); // [8, 32]
            println!("bias             {:?}", w.bias.shape()); // [1, 32]
        }
        _ => unreachable!(),
    }
}
text 复制代码
kernel           [4, 32]
recurrent_kernel [8, 32]
bias             [1, 32]

想要可复现的初始化,调用 with_random_state(seed):它会以确定的方式重跑核与递归核的采样,遗忘门偏置为 1.0 的规则依然生效。不传种子的话,权重就从全局种子或系统熵取种。见 7.1. 可复现性与随机种子

要手动装入权重,比如从别的框架移植过来,或者对一段精确的递归做单元测试,每个层都提供 set_weights(kernel, recurrent_kernel, bias),直接接收融合后的矩阵。LSTM 和 GRU 还额外提供 set_gate_weights(...),它按门接收一组 (kernel, recurrent_kernel, bias) 三元组,并替你拼接成融合布局。

按门传参的顺序,LSTM 是 (input, forget, cell, output),共 12 个数组;GRU 是 (reset, update, candidate),共 9 个数组。任何对不上的形状都会返回 Error::NeuralNetwork(NnError::WeightShape { .. })。保存或加载整个模型时,这些数组会原样序列化。见 3.9. 权重保存与加载

3.7.6. BPTT 与开销模型

训练用的是随时间反向传播(BPTT),而且是完整展开 ,没有截断窗口。前向传播时,层会把反向传播需要的一切都缓存下来:SimpleRNN 存每一个隐藏状态(前面补上 h_0 = 0);LSTM 还会存细胞状态、activation(c_t),以及每个时间步的 4 个门激活值;GRU 存重置门、更新门、候选值,以及 r_t * h_{t-1} 这个乘积。

predict 给这些缓存传 None,跳过记录和相应的克隆,这就是推理比训练时的 forward 更省的原因。内存开销随 timesteps 线性增长:长序列吃的是内存,不只是时间。

backward 反向遍历时间步。它要求一个二维的上游梯度 (batch, units),也就是损失对最终隐藏状态的梯度,这是层唯一需要的梯度,因为层对外吐出来的就只有这个最终状态。每一步,backward 都会算出该时间步的预激活梯度,把 grad_h(LSTM 还有 grad_c)串回上一步,这部分天生沿时间串行

接着 backward 把权重梯度的归约批量化:把逐时间步的 dz(batch, timesteps) 上折叠起来,kernel 和 bias 梯度各自从一个大 GEMM 里算出。递归核的梯度通常也是这样,但 GRU 例外:它有 2 个不同的递归输入,要用 2 个 GEMM 才能算出递归核的梯度。RustyML 以替换语义存储梯度,层内部不做 裁剪;如果需要梯度裁剪,用一个提供该功能的优化器。见 3.4. 优化器

要记住的性能形态是:沿时间轴串行,沿 batch 轴和融合的门轴并行 。输入投影 x @ W 不依赖递归,RustyML 会一次性把它算完,作为一个跨所有时间步的批量 GEMM。只有 h_{t-1} @ U 这一项必须一步一步来,而其中每一步本身又是一个 batch 并行的 GEMM。

GRU 在这上面还能再省一点:它把重置门和更新门的递归投影融合成一个 GEMM,因为两者都要读 h_{t-1}。只有候选值的递归投影仍然单独计算,因为它的输入 r_t * h_{t-1} 依赖刚算好的重置门。

每一次矩阵乘法都交给 gemmkit 后端(见 6.2. 矩阵乘法)。gemmkit 会根据工作量自行决定并行规模:宽层和大 batch 会自动获得线程并行,小层则保持串行以避免额外开销。一个时间步里融合的门投影,总会在计算乘积的同一趟里施加偏置。只有 SimpleRNN 搭配 ReLU 时,激活函数才会融合进这一趟里,LSTMGRU,以及其他任何激活函数都不会融合。

实用的结论是:更多序列,也就是更大的 batch,能很好地并行;更多时间步则不行,因为那个轴是一条串行的依赖链。线程池的各项设置见 7.3. 性能调优与并行

3.7.7. 堆叠与构建模型

循环层只返回最后一个隐藏状态,一个二维的 (batch, units) 张量。你不能 把这个输出直接喂给另一个循环层:循环层要的是三维的 (batch, timesteps, features) 输入,面对二维张量只会返回 Error::InvalidInput。RustyML 没有 return_sequences 选项,层没法吐出一条逐时间步的序列,所以 Keras 意义上那种循环层的深层堆叠,在这里搭不出来。请提前为这个限制做好规划。

标准的写法是:拿一个循环层当序列编码器,后面接一个稠密头,把最终状态映射到目标上。这样组合起来很干净:循环层把 (batch, timesteps, features) 变成 (batch, units)Dense 要的恰好就是这个二维形状。

下面的示例学的是一个真实的序列任务:预测一条长度为 3 的标量序列之和。它用 LSTM 编码器、一个 Dense 读出层、Adam、均方误差搭起来,能在几百个全批量 epoch 内收敛(回想 3.1. Sequential模型 中提到,fit 每个 epoch 只走一步全批量梯度):

rust 复制代码
use ndarray::Array;
use rustyml::neural_network::sequential::Sequential;
use rustyml::prelude::*;

fn main() {
    // 4 条序列,3 个时间步,1 个特征。目标 = 3 个标量之和。
    let x = Array::from_shape_vec(
        (4, 3, 1),
        vec![
            0.1, 0.2, 0.1, // 和为 0.4
            0.3, 0.1, 0.2, // 和为 0.6
            0.0, 0.2, 0.2, // 和为 0.4
            0.2, 0.2, 0.1, // 和为 0.5
        ],
    )
    .unwrap()
    .into_dyn();
    let y = Array::from_shape_vec((4, 1), vec![0.4, 0.6, 0.4, 0.5])
        .unwrap()
        .into_dyn();

    let mut model = Sequential::new();
    model
        // 循环特征提取器:(batch, 3, 1) -> (batch, 16)
        .add(LSTM::new(1, 16, Tanh::new()).unwrap().with_random_state(42))
        // 把最后的隐藏状态读出成单个标量。
        .add(Dense::new(16, 1, Linear::new()).unwrap())
        .compile(
            Adam::new(0.01, 0.9, 0.999, 1e-8, 0.0).unwrap(),
            MeanSquaredError::new(),
        );

    model.fit(&x, &y, 300).unwrap();

    let pred = model.predict(&x).unwrap();
    println!("target      : {:?}", y.as_slice().unwrap());
    println!("prediction  : {:?}", pred.as_slice().unwrap());
}

300 个 epoch 之后,4 个预测值都落在 [0.4, 0.6, 0.4, 0.5] 附近的几个百分点以内。LSTM 学会了累加这条序列,稠密头再把累加结果读出来。把 LSTM 换成 GRUSimpleRNN,同一套代码照样能编译、能训练:在这条短序列上,3 种层都会收敛,这恰恰说明门控单元在梯度消失上的优势,只有在长序列上才会显现出来。

model.summary() 会把每个循环层的输出形状打印成 (None, units),其中 None 代表动态的 batch 维度。数据集只要比玩具规模大一些,就应该用 fit_with_batches 代替 fit,这样每个 epoch 会走好几步小批量,并重新洗牌数据。

相关推荐
言乐65 小时前
Python游戏水平测试辅助系统
开发语言·python·游戏·django·pygame
东风破_12 小时前
ESLint 是什么?为什么你的项目需要它?
前端·后端·代码规范
嘻哈∠※12 小时前
0061基于 SpringBoot 的投稿与稿件处理系统设计与实现
java·spring boot·后端
WWJA王文举12 小时前
I²C通信完整流程详解:START、地址、ACK、数据、Repeated START和STOP一次讲透
c语言·开发语言
卷无止境12 小时前
在 awesome-fastapi 里,哪些库值得一看?
后端·python
ltl13 小时前
LLM 训练全景:Pre-train、SFT、RLHF、DPO 与蒸馏
机器学习
zhanghaha131413 小时前
Python进阶教程:6_JSON 数据解析 —— 新手完全指南
开发语言·python·json
知识分享小能手14 小时前
线性代数学习教程,从入门到精通,向量组的线性相关性 — 完整知识点梳理(7)
学习·线性代数·机器学习
卷无止境14 小时前
FastAPI 的Admin面板生态
后端·python