资讯动态

Flax 高级 RNN 层设计全解:从 FLIP 2396 到 nn.RNN 与 Bidirectional 的源码级剖析

发布时间:2026/9/16 18:58:36 来源:尧图企业网站定制
Flax 高级 RNN 层设计全解从 FLIP 2396 到 nn.RNN 与 Bidirectional 的源码级剖析【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax本篇以 Flax 的 FLIP 2396RNN Flip为主线完整梳理「高层循环层」的设计动机、三层抽象结构、RNN/Bidirectional/RNNBase的 API 语义与 masking 机制并对照当前仓库flax/linen/recurrent.py的实际实现、seq2seq 示例与单元测试说明这套抽象从提案落地为可运行代码的完整过程。读完你既能看懂这套 API 的设计权衡也能直接上手编写带 padding、双向扫描、循环 dropout 的 Flax 循环网络。1. 背景与动机手动 nn.scan 写 LSTM 有多繁琐FLIP 23962022-08-18作者 Jasmijn Bastings 与 Cristian Garcia后续由 Cristian Garcia、Marcus Chiam 等人跟进提出的核心目标是为已有循环单元RNNCellBase 子类之上提供更高层的 RNN、GRU、LSTM 层帮助用户更方便地处理输入序列。提案给出的动机非常具体即便是一个简单的 LSTM 层用户也必须手动创建和管理 carry记忆状态并正确配置nn.scan例如nn.compact def __call__(self, x): LSTM nn.scan( nn.LSTMCell, variable_broadcastparams, split_rngs{params: False} ) carry LSTM.initialize_carry( jax.random.key(0), batch_dimsx.shape[:1], sizeself.hidden_size ) carry, x LSTM()(carry, x) return x而一旦涉及 padding 等更复杂的场景比如 seq2seq 示例 examples/seq2seq/models.py手动代码的工作量会成倍增长。FLIP 因此提出应为用户提供干净、正确且高效的循环单元抽象。从仓库现状看这个 FLIP 已经完整落地flax/linen/recurrent.py 中实现了RNNCellBase、LSTMCell、GRUCell、MGUCell、SimpleCell、ConvLSTMCell、OptimizedLSTMCell等单元以及 FLIP 提出的RNN、Bidirectional两个高层层并在 flax/linen/init.py 中以nn.RNN、nn.Bidirectional等名字导出。2. 设计需求四条硬性要求FLIP 在 Requirements 一节列出了四个必须满足的需求它们也构成了nn.RNN全部参数设计的出发点Masking掩码必须支持批次内每条序列尾部带 padding 的情形。出于性能考虑不支持非连续 padding即 padding 不在序列末尾的情况除非采用 packing见第 7 节「未来想法」。Bidirectionality双向能够沿正向与反向两个方向处理序列且必须尊重 padding——反向方向应当从真实输入而非 padding 值开始。Performance性能提案要求对候选类做基准测试在步长时间与/或内存占用上取得最佳表现。Recurrent Dropout循环 dropout支持单元内部对状态施加的 dropout。3. 三层抽象结构Cells / Layers / BidirectionalFLIP 提议采用三层抽象这是整个设计的骨架Cells不改动所有RNNCellBase子类LSTMCell、GRUCell等实现单步stepwise逻辑。Flax 当时已具备这些单元。Layers新增一个RNN类接收一个 cell 实例并沿序列扫描尊重可能的 padding 值可选支持打包packed序列。Bidirectional新增单个类接收前向与反向两个RNN实例正确地以两个方向处理输入序列并合并结果。FLIP 中给出的目标 API 示例如下提案阶段的原始形式cell nn.LSTMCell() # 编码一批输入序列。 carry, outputs nn.RNN(cell, cell_size)(inputs, seq_lengths)双向层前向、反向均为 LSTM的用法forward_rnn nn.RNN(nn.LSTMCell(), cell_size32) backward_rnn nn.RNN(nn.LSTMCell(), cell_size32) # 双向组合器。 bi_rnn nn.Bidirectional(forward_rnn, backward_rnn) # 双向编码一批输入序列。 carry, outputs bi_rnn(inputs, seq_lengths)3.1 与当前实现的差异cell_size 去哪了需要注意当前仓库的RNN签名中已没有cell_size参数。这是后续 FLIP 3099docs/flip/3099-rnnbase-refactor.md状态Implemented重构的结果隐藏层大小直接作为features传入 cell 构造器而initialize_carry被改为实例方法由模块自身推断 batch 维与特征维形状。当前仓库中的实际用法与 flax/linen/recurrent.py 中RNN的 docstring 示例一致为import jax import jax.numpy as jnp import flax.linen as nn x jnp.ones((10, 50, 32)) # (batch, time, features) lstm nn.RNN(nn.LSTMCell(64)) # 隐藏单元数在 cell 上指定 variables lstm.init(jax.random.key(0), x) y lstm.apply(variables, x) print(y.shape) # (10, 50, 64)对于带空间维的ConvLSTMCell则不需要任何额外形状参数因为形状可以从输入推断x jnp.ones((10, 50, 32, 32, 3)) # (batch, time, height, width, features) conv_lstm nn.RNN(nn.ConvLSTMCell(64, kernel_size(3, 3))) y, variables conv_lstm.init_with_output(jax.random.key(0), x) print(y.shape) # (10, 50, 32, 32, 64)从源码结构看RNN.__call__通过 cell 的num_feature_axes属性自动推导时间轴位置time_axis inputs.ndim - (self.cell.num_feature_axes 1)见 flax/linen/recurrent.py#L1071-L1085这正是 FLIP 3099 中num_feature_dims机制的落地——普通LSTMCell/GRUCell返回 1而ConvLSTMCell返回len(kernel_size) 1。4. RNNBase 协议call参数逐一解析FLIP 定义了RNNBase作为RNN的基类/协议它规定了所有 RNN 层必须实现的 API以便与Bidirectional组合。FLIP 中的原始定义class RNNBase(Protocol): def __call__( self, inputs: jax.Array, *, initial_carry: Optional[Carry] None, init_key: Optional[random.KeyArray] None, seq_lengths: Optional[Array] None, return_carry: Optional[bool] None, time_major: Optional[bool] None, reverse: Optional[bool] None, keep_order: Optional[bool] None, ) - Union[Output, Tuple[Carry, Output]]: ...当前仓库中的RNNBase同样是一个typing_extensions.Protocol见 flax/linen/recurrent.py#L1246-L1259签名与提案完全一致。各参数语义如下FLIP 原文定义当前实现逐字保留参数语义默认值inputs输入序列—initial_carry初始 carry未提供时通过 cell 的initialize_carry方法初始化Noneinit_key用于初始化 carry 的 PRNG key未提供时使用jax.random.key(0)。大多数 cell 会忽略该参数Noneseq_lengths可选的整型数组形状为(*batch)指示每条序列的长度时间维度上索引大于对应长度的元素被视为 padding 并被忽略Nonereturn_carryFalse时仅返回输出序列True时返回最终 carry, 输出序列元组Falsetime_majorFalse默认时期望输入形状为(*batch, time, *features)True时期望(time, *batch, *features)FalsereverseFalse时从左到右处理并按原始顺序返回True时从右到左处理、按反转顺序返回。若传入seq_lengthspadding 始终留在序列末尾Falsekeep_orderTrue且reverseTrue时处理完成后将输出翻回原始顺序便于在双向 RNN 中对齐序列默认False保持reverse指定的顺序False返回值return_carryFalse时仅输出序列否则为最终 carry, 输出序列元组—RNN的构造函数属性当前实现见 flax/linen/recurrent.py#L1001-L1014cell: RNNCellBase time_major: bool False return_carry: bool False reverse: bool False keep_order: bool False unroll: int 1 variable_axes: Mapping[CollectionFilter, InOutScanAxis] FrozenDict() variable_broadcast: CollectionFilter params variable_carry: CollectionFilter False split_rngs: Mapping[PRNGSequenceFilter, bool] FrozenDict({params: False})FLIP 原文说明variable_axes、variable_broadcast、variable_carry、split_rngs这些属性直接透传给nn.scan其默认值被设置为让LSTMCell、GRUCell等常见单元开箱即用即variable_broadcastparams让参数在时间步间共享split_rngs{params: False}防止参数集合被当作逐时间步 RNG 拆分。当前实现中这组默认值与提案完全一致并额外暴露了unroll控制展开程度scan内一次迭代展开的步数默认 1。覆盖 scan 默认值的用法RNNdocstring 示例lstm nn.RNN( nn.LSTMCell(64), unroll1, variable_axes{}, variable_broadcastparams, variable_carryFalse, split_rngs{params: False})time_majorTrue的形态切换同样支持输出形状随之变为(time, batch, cell_size)x jnp.ones((50, 10, 32)) # (time, batch, features) lstm nn.RNN(nn.LSTMCell(64), time_majorTrue) variables lstm.init(jax.random.key(0), x) y lstm.apply(variables, x) print(y.shape) # (50, 10, 64)5. Masking为什么选择 seq_lengths 这种掩码格式FLIP 的 Masking 小节将seq_lengths定义为形状(*batch,)的整型数组指示每条序列的长度。提案还专门讨论了业界主流的三种掩码表示这是理解该 API 取舍的关键Binary masking二值掩码逐样本、逐时间步指明该数据点是否参与计算允许非连续如[1, 1, 0, 1]。Keras 采用这种格式。Sequence length masking序列长度掩码逐样本指明序列中非 padding 样本的数量padding 必须堆叠在末尾。FlaxFormer 采用这种格式。Segmentation Mask分段掩码指明每个时间步属于哪一条样本允许一行中包含多条序列从而减少总 padding 量如[1, 1, 1, 2, 2, 0, 0]。PyTorch 使用这种表示其pack_padded_sequence工具即基于此。提案结论序列打包sequence packingFlax 的 LM1B 示例 examples/lm1b/input_pipeline.py 中即有应用虽然更强大但实现更复杂是否值得存疑最简单的序列长度掩码是最终选择。这一取舍直接体现在nn.RNN的实现中当前源码的行为可以归纳为传入seq_lengths后padding 步的输出不会被置零RNNdocstring 明确说明 The output elements corresponding to padding elements are NOT zeroed out若同时return_carryTrue返回的 carry 是每条序列最后一个有效时间步的状态而非整段 padded 序列末尾的状态。第 2 点在源码中通过_select_last_carry实现见 flax/linen/recurrent.py#L1166-L1172def _select_last_carry(sequence: A, seq_lengths: jnp.ndarray) - A: last_idx seq_lengths - 1 def _slice_array(x: jnp.ndarray): return x[last_idx, jnp.arange(x.shape[1])] return jax.tree_util.tree_map(_slice_array, sequence)而实现上采用了 FLIP 未展开的细节优化源码注释原话This uses more memory but is faster than using jnp.where at each iteration当seq_lengths与return_carry同时存在时scan_fn会额外把每一步的 carry 作为输出保留下来形成 carry 历史扫描结束后用上述按行索引一次性挑选避免在每个时间步做jnp.where。见 flax/linen/recurrent.py#L1109-L1147。seq2seq 示例正是这套 masking 的典型消费方examples/seq2seq/models.py 中先计算序列长度再传给编码器def get_seq_lengths(self, inputs: Array) - Array: Get segmentation mask for inputs. # undo one-hot encoding inputs jnp.argmax(inputs, axis-1) # calculate sequence lengths seq_lengths jnp.argmax(inputs self.eos_id, axis-1) return seq_lengths encoder nn.RNN( nn.LSTMCell(self.hidden_size), return_carryTrue, nameencoder) ... seq_lengths self.get_seq_lengths(encoder_inputs) encoder_state, _ encoder(encoder_inputs, seq_lengthsseq_lengths)编码器提取的最终状态随后作为解码器的initial_carry传入——这正是RNNBase.__call__中initial_carry参数存在的意义也印证了 FLIP 三层抽象中「carry 管理交给 RNN 层」的设计意图。6. 反向扫描与 flip_sequences双向的正确性基础FLIP 对 Bidirectional 的要求是反向方向应从真实输入而非 padding 开始。要满足这一点简单地对矩阵做jnp.flip是不行的对于被 padding 的序列naive 翻转后首元素会变成 padding 值。为此 FLIP 引入flip_sequences语义当前实现位于 flax/linen/recurrent.py#L1180-L1238其 docstring 示例清晰地说明了行为inputs [[1, 0, 0], [2, 3, 0], [4, 5, 6]] lengths [1, 2, 3] flip_sequences(inputs, lengths) [[1, 0, 0], [3, 2, 0], [6, 5, 4]]即只翻转每条序列真实长度内的元素padding 保持留在末尾。核心算法是用取模运算构造翻转索引再jnp.take_along_axisidxs jnp.arange(max_steps - 1, -1, -1) # [max_steps] idxs (idxs seq_lengths) % max_steps # [*batch, max_steps] outputs jnp.take_along_axis(inputs, idxs, axistime_axis)在RNN.__call__中reverseTrue时先对输入调用flip_sequences扫描完成后若keep_orderTrue再对输出调用一次翻回见 flax/linen/recurrent.py#L1087-L1158。单元测试test_flip_sequence系列含 batch、多特征维、time_major 各变体见 tests/linen/linen_recurrent_test.py#L390-L428以及test_reverse/test_reverse_but_keep_order逐一验证了语义反向处理时的输出应与逐时间步手动按xs[batch_idx, seq_len - i - 1]喂给 cell 的结果在数值上等价。7. Bidirectional前向/反向编码与结果合并FLIP 给出的Bidirectional伪代码def __call__(self, inputs, seq_lengths): # 前向编码。 carry_forward, outputs_forward self.forward_rnn( inputs, seq_lengthsseq_lengths, return_carryTrue, reverseFalse, ) # 反向编码。 carry_backward, outputs_backward self.backward_rnn( inputs, seq_lengthsseq_lengths, return_carryTrue, reverseTrue, # 按反转顺序处理 keep_orderTrue, # 但按原始顺序返回 ) # 合并两条序列。 outputs jax.tree.map(self.merge_fn, outputs_forward, outputs_backward) return (carry_forward, carry_backward), outputs其中merge_fn是一个接收双向输出并融合的函数默认为concat。提案中的混合用法示例前向 LSTM、反向 GRUforward_rnn nn.RNN(nn.LSTMCell(), cell_size32) backward_rnn nn.RNN(nn.GRUCell(), cell_size32) # 双向组合器。 bi_rnn nn.Bidirectional(forward_rnn, backward_rnn) # 双向编码一批输入序列。 carry, outputs bi_rnn(inputs, seq_lengths)当前仓库的Bidirectional实现flax/linen/recurrent.py#L1262-L1344与伪代码高度吻合且补充了几个提案未细化的工程细节RNG 拆分若传入init_key会用random.split拆成key_forward/key_backward分别初始化两个方向的 carrycarry 拆分initial_carry若非None会被拆为前向/反向两个 carry参数共享警告若forward_rnn is backward_rnn用户误传同一对象会记录一条 warning 提示二者将共享参数——对应测试test_shared_cell见 tests/linen/linen_recurrent_test.py#L454-L467可定制 mergemerge_fn默认为沿最后一维_concatenate测试test_custom_merge_fn验证了merge_fnlambda x, y: x y时输出形状从(batch, seq, 2*out)变为(batch, seq, out)。一个可直接运行的最小示例取自Bidirectionaldocstringlayer nn.Bidirectional(nn.RNN(nn.GRUCell(4)), nn.RNN(nn.GRUCell(4))) x jnp.ones((2, 3)) variables layer.init(jax.random.key(0), x) out layer.apply(variables, x) # out.shape (2, 3, 8)测试test_bidirectional确认默认 concat 合并下输出为(batch, seq, channels_out * 2)test_return_carry确认return_carryTrue时返回的 carry 为(carry_forward, carry_backward)二元组二者各自形状为((batch, out), (batch, out))LSTM 的 (c, h) 对。8. 循环 Dropout用 split_rngs 区分两类 dropoutFLIP 指出 RNN 中 dropout 有两种主要用途Input dropout施加在输入上的常规 dropout每个时间步各不相同Recurrent dropout施加在循环输入/输出状态上的 dropout所有时间步相同。提案认为nn.scan可以天然表达这两种 dropout区别只在split_rngsinput dropout 需要按步拆分 RNGrecurrent dropout 则不需要。配合此前引入的nn.Dropout自定义rng_name能力对应 PR #2540cell 内部可以定义两种 dropoutself.dropout nn.Dropout(...) # input dropout self.recurrent_dropout nn.Dropout(..., rng_collectionrecurrent_dropout)进而nn.scan/nn.RNN可以相应指定split_rngsnn.scan(scan_fn, ..., split_rngs{dropout: True, recurrent_dropout: False})在高层 API 下这等价于构造RNN时传入split_rngs{params: False, dropout: True, recurrent_dropout: False}——seq2seq 示例的解码器就展示了自定义 rng 集合的用法examples/seq2seq/models.py#L113-L119 中nn.RNN(DecoderLSTMCell(...), split_rngs{params: False, lstm: True})其中DecoderLSTMCell通过self.make_rng(lstm)在循环内采样 tokenlstm: True保证每个时间步拿到不同的采样 key。9. 从 FLIP 2396 到落地RNNCell 设计的后续演进FLIP 2396 末尾的「Future ideas」讨论了两个未纳入首期实现的方向其中一个后来真正发生9.1 序列打包Sequence Packing——仍是未来方向允许把多条序列打包以减少 padding、提升内存/空间利用效率。代价是步长时间可能增加每个时间步都要检查是否进入新序列并重置 carry/初始状态但从总体减少 padding 的角度看可能更划算。当前仓库的RNN并未实现 packing仍只支持序列长度掩码与 FLIP 的需求约束一致。9.2 RNNCell 重构——已由 FLIP 3099 实现FLIP 2396 提出的两个替代方案方案 A把initialize_carry变成实例方法。签名变为def initialize_carry(self, sample_input) - Carry超参可直接传给 cell用法简化为LSTM nn.scan(nn.LSTMCell, variable_broadcastparams, split_rngs{dropout: True}) lstm LSTM(features32) carry lstm.initialize_carry(x[:, 0]) carry, y lstm(carry, x)这正是当前仓库的实际形态flax/linen/recurrent.py#L60-L78 中RNNCellBase.initialize_carry(self, rng, input_shape)为实例方法带nowrapLSTMCell、GRUCell等各自实现并新增num_feature_axes属性供RNN推断时间轴——这与 docs/flip/3099-rnnbase-refactor.md 描述的最终落地方案一致。该 FLIP 还解释了动机旧 API类方法 手动拆分 batch 维与特征维如nn.ConvLSTMCell.initialize_carry(key, (16,), (64, 64, 16))容易让人手动计算本应由模块推断的形状新 API 下只需carry lstm.initialize_carry(key1, input_shapex.shape)。方案 B彻底移除initialize_carry把 carry 状态作为一个 collection 处理用法进一步简化为LSTM nn.scan(nn.LSTMCell, variable_broadcastparams, split_rngs{dropout: True}) y LSTM(features32)(carry, x)但 FLIP 指出该方案要求nn.scan支持对 carry collection 的初始化当时尚不可行且即使用户不关心输出 carry 也须显式声明mutable[carry]因此未采纳——当前仓库中initialize_carry依然存在。10. 正确性验证单元测试如何印证设计tests/linen/linen_recurrent_test.py 对该套 API 做了系统验证可作为读者自查正确性的参照形状与参数断言test_rnn_basic_forward、test_rnn_multiple_batch_dims、test_rnn_with_spatial_dimensions验证(*batch, time, *features)输出形状、variables[params][cell]下 kernel/bias 的维度以及 ConvLSTM 的空间维 carry 形状数值等价性test_numerical_equivalence系列把nn.RNN的扫描输出与逐时间步手动调用rnn.cell.apply的结果做assert_allclosertol1e-5并覆盖带 masktest_numerical_equivalence_with_mask验证每个 batch 取length - 1位置的 carry 与 RNN 返回值一致、单 batch、nn.scan手工版、jax.lax.scan手工版等对照路径——这直接证明了「RNN 层只是正确配置了 scan 的 cell」这一设计主张反向与保序test_reverse、test_reverse_but_keep_order验证reverse/keep_order的语义flip_sequences4 个变体测试覆盖 padding 保持、多特征轴与 time_majorBidirectionalBidirectionalTest默认 concat 输出形状、共享 cell 警告路径、自定义merge_fn、return_carry的双元组结构。11. 小结这套抽象给开发者带来了什么回到 FLIP 2396 的核心命题——实现知名循环结构繁琐且易错——当前仓库给出的答案是清晰的分工Cell 层LSTMCell、GRUCell、MGUCell、ConvLSTMCell、OptimizedLSTMCell等继续只管单步数学构造时直接声明features等超参initialize_carry为实例方法、形状可自推断FLIP 3099 的成果nn.RNN层把「carry 初始化、scan 配置、多 batch 维/多特征轴推断、padding 掩码、末位有效 carry 提取、反向翻转」全部封装variable_broadcastparams、split_rngs{params: False}等默认值保证常见 cell 开箱即用其余 scan 参数仍可透传覆盖nn.Bidirectional负责前后向的 RNG/carry 拆分、reverseTrue keep_orderTrue的正确反向编码以及可插拔的merge_fn默认 concat。对于需要更底层控制的用户nn.scan的手工写法仍然可用且被测试证明与nn.RNN数值等价而对于带 padding 的编码器如 seq2seq 示例与双向编码器nn.RNNseq_lengthsnn.Bidirectional就是 FLIP 2396 所承诺的几行代码级别的抽象。适用前提说明本文所有 API 说明均以当前仓库flax/linen/recurrent.py的实际实现为准含 FLIP 3099 重构后的initialize_carry签名与features构造参数FLIP 原文中cell_size、time_axis等提案期参数在当前实现中已被移除若你使用的 Flax 版本早于 0.7.0其 RNN 层 API 可能与本文描述存在差异。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价