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 2396(RNN Flip)为主线,完整梳理「高层循环层」的设计动机、三层抽象结构、RNN/Bidirectional/RNNBase的 API 语义与 masking 机制,并对照当前仓库flax/linen/recurrent.py的实际实现、seq2seq 示例与单元测试,说明这套抽象从提案落地为可运行代码的完整过程。读完你既能看懂这套 API 的设计权衡,也能直接上手编写带 padding、双向扫描、循环 dropout 的 Flax 循环网络。
1. 背景与动机:手动 nn.scan 写 LSTM 有多繁琐
FLIP 2396(2022-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_broadcast="params", split_rngs={"params": False} ) carry = LSTM.initialize_carry( jax.random.key(0), batch_dims=x.shape[:1], size=self.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 / Bidirectional
FLIP 提议采用三层抽象,这是整个设计的骨架:
- 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_size=32) backward_rnn = nn.RNN(nn.LSTMCell(), cell_size=32) # 双向组合器。 bi_rnn = nn.Bidirectional(forward_rnn, backward_rnn) # 双向编码一批输入序列。 carry, outputs = bi_rnn(inputs, seq_lengths)3.1 与当前实现的差异:cell_size 去哪了
需要注意,当前仓库的RNN签名中已没有cell_size参数。这是后续 FLIP 3099(docs/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方法初始化 | None |
init_key | 用于初始化 carry 的 PRNG key;未提供时使用jax.random.key(0)。大多数 cell 会忽略该参数 | None |
seq_lengths | 可选的整型数组,形状为(*batch),指示每条序列的长度;时间维度上索引大于对应长度的元素被视为 padding 并被忽略 | None |
return_carry | False时仅返回输出序列;True时返回(最终 carry, 输出序列)元组 | False |
time_major | False(默认)时期望输入形状为(*batch, time, *features);True时期望(time, *batch, *features) | False |
reverse | False时从左到右处理并按原始顺序返回;True时从右到左处理、按反转顺序返回。若传入seq_lengths,padding 始终留在序列末尾 | False |
keep_order | True且reverse=True时,处理完成后将输出翻回原始顺序,便于在双向 RNN 中对齐序列;默认False保持reverse指定的顺序 | False |
| 返回值 | return_carry=False时仅输出序列;否则为(最终 carry, 输出序列)元组 | — |
RNN的构造函数属性(当前实现,见 flax/linen/recurrent.py#L1001-L1014):
cell: 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_broadcast='params'让参数在时间步间共享,split_rngs={'params': False}防止参数集合被当作逐时间步 RNG 拆分)。当前实现中这组默认值与提案完全一致,并额外暴露了unroll控制展开程度(scan内一次迭代展开的步数,默认 1)。
覆盖 scan 默认值的用法(RNNdocstring 示例):
lstm = nn.RNN( nn.LSTMCell(64), unroll=1, variable_axes={}, variable_broadcast='params', variable_carry=False, split_rngs={'params': False})time_major=True的形态切换同样支持,输出形状随之变为(time, batch, cell_size):
x = jnp.ones((50, 10, 32)) # (time, batch, features) lstm = nn.RNN(nn.LSTMCell(64), time_major=True) 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 packing,Flax 的 LM1B 示例 examples/lm1b/input_pipeline.py 中即有应用)虽然更强大,但实现更复杂,是否值得存疑;最简单的序列长度掩码是最终选择。这一取舍直接体现在nn.RNN的实现中,当前源码的行为可以归纳为:
- 传入
seq_lengths后,padding 步的输出不会被置零(RNNdocstring 明确说明 "The output elements corresponding to padding elements are NOT zeroed out"); - 若同时
return_carry=True,返回的 carry 是每条序列最后一个有效时间步的状态,而非整段 padded 序列末尾的状态。
第 2 点在源码中通过_select_last_carry实现(见 flax/linen/recurrent.py#L1166-L1172):
def _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_carry=True, name='encoder') ... seq_lengths = self.get_seq_lengths(encoder_inputs) encoder_state, _ = encoder(encoder_inputs, seq_lengths=seq_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_axis:
idxs = 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, axis=time_axis)在RNN.__call__中,reverse=True时先对输入调用flip_sequences,扫描完成后若keep_order=True再对输出调用一次翻回(见 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_lengths=seq_lengths, return_carry=True, reverse=False, ) # 反向编码。 carry_backward, outputs_backward = self.backward_rnn( inputs, seq_lengths=seq_lengths, return_carry=True, reverse=True, # 按反转顺序处理 keep_order=True, # 但按原始顺序返回 ) # 合并两条序列。 outputs = jax.tree.map(self.merge_fn, outputs_forward, outputs_backward) return (carry_forward, carry_backward), outputs其中merge_fn是一个接收双向输出并融合的函数,默认为concat。提案中的混合用法示例(前向 LSTM、反向 GRU):
forward_rnn = nn.RNN(nn.LSTMCell(), cell_size=32) backward_rnn = nn.RNN(nn.GRUCell(), cell_size=32) # 双向组合器。 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分别初始化两个方向的 carry; - carry 拆分:
initial_carry若非None,会被拆为前向/反向两个 carry; - 参数共享警告:若
forward_rnn is backward_rnn(用户误传同一对象),会记录一条 warning 提示二者将共享参数——对应测试test_shared_cell(见 tests/linen/linen_recurrent_test.py#L454-L467); - 可定制 merge:
merge_fn默认为沿最后一维_concatenate,测试test_custom_merge_fn验证了merge_fn=lambda x, y: x + y时输出形状从(batch, seq, 2*out)变为(batch, seq, out)。
一个可直接运行的最小示例(取自Bidirectionaldocstring):
layer = 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_carry=True时返回的 carry 为(carry_forward, carry_backward)二元组,二者各自形状为((batch, out), (batch, out))(LSTM 的 (c, h) 对)。
8. 循环 Dropout:用 split_rngs 区分两类 dropout
FLIP 指出 RNN 中 dropout 有两种主要用途:
- Input dropout:施加在输入上的常规 dropout,每个时间步各不相同;
- Recurrent dropout:施加在循环输入/输出(状态)上的 dropout,所有时间步相同。
提案认为nn.scan可以天然表达这两种 dropout,区别只在split_rngs:input dropout 需要按步拆分 RNG,recurrent dropout 则不需要。配合此前引入的nn.Dropout自定义rng_name能力(对应 PR #2540),cell 内部可以定义两种 dropout:
self.dropout = nn.Dropout(...) # input dropout self.recurrent_dropout = nn.Dropout(..., rng_collection='recurrent_dropout')进而nn.scan/nn.RNN可以相应指定split_rngs:
nn.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')在循环内采样 token,'lstm': 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_broadcast='params', split_rngs={'dropout': True}) lstm = LSTM(features=32) carry = lstm.initialize_carry(x[:, 0]) carry, y = lstm(carry, x)这正是当前仓库的实际形态:flax/linen/recurrent.py#L60-L78 中RNNCellBase.initialize_carry(self, rng, input_shape)为实例方法(带@nowrap),LSTMCell、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_shape=x.shape)。
方案 B:彻底移除initialize_carry,把 carry 状态作为一个 collection 处理,用法进一步简化为:
LSTM = nn.scan(nn.LSTMCell, variable_broadcast='params', split_rngs={'dropout': True}) y = LSTM(features=32)(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_allclose(rtol=1e-5),并覆盖带 mask(test_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_sequences(4 个变体测试):覆盖 padding 保持、多特征轴与 time_major;
- Bidirectional(
BidirectionalTest):默认 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_broadcast='params'、split_rngs={'params': False}等默认值保证常见 cell 开箱即用,其余 scan 参数仍可透传覆盖;nn.Bidirectional负责前后向的 RNG/carry 拆分、reverse=True + keep_order=True的正确反向编码,以及可插拔的merge_fn(默认 concat)。
对于需要更底层控制的用户,nn.scan的手工写法仍然可用且被测试证明与nn.RNN数值等价;而对于带 padding 的编码器(如 seq2seq 示例)与双向编码器,nn.RNN+seq_lengths+nn.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),仅供参考