拓冰建站拓冰建站
首页 / 资讯中心 / 正文

Flax NNX 循环神经网络模块完全指南:LSTM / GRU / RNN / Bidirectional 从源码到实战

Flax NNX 循环神经网络模块完全指南LSTM / GRU / RNN / Bidirectional 从源码到实战【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flaxFlax 的 NNX 子库flax.nnx在 flax/nnx/nn/recurrent.py 中提供了一套完整、面向对象风格的循环神经网络RNN组件涵盖LSTMCell、OptimizedLSTMCell、GRUCell、SimpleCell四种细胞Cell以及基于scan的序列层RNN和双向组合层Bidirectional。本文以 docs_nnx/api_reference/flax.nnx/nn/recurrent.rst 索引的 API 为核心结合源码实现与 tests/nnx/nn/recurrent_test.py 中的测试用例系统讲解每种组件的数学定义、参数语义、初始化协议与真实用法帮助你在 JAX 生态中快速落地定长/变长序列建模任务。一、架构总览Cell、Layer、Bidirectional 三层抽象设计文档 docs_nnx/flip/2396-rnn.md 明确提出 RNN 组件应分为三层抽象Cell细胞逐时间步逻辑RNNCellBase及其子类LSTMCell、OptimizedLSTMCell、GRUCell、SimpleCell负责定义给定当前 carry 与当前时刻输入如何产出新 carry 与输出的单步计算。Layer层序列扫描RNN类接收任意RNNCellBase实例通过flax.nnx.scan沿时间轴展开序列计算并内置对 paddingseq_lengths、反向reverse、时间轴布局time_major等细节的处理。Bidirectional双向组合Bidirectional类同时驱动一个前向RNN与一个反向RNN并将两个方向的输出按merge_fn合并。这一分层结构的好处是Cell 保持零状态管理的纯粹单步逻辑而序列级的 padding、翻转、carry 取舍等易错细节被统一收敛到RNN层。你完全可以基于这层抽象自定义自己的 Cell只需实现RNNCellBase协议RNN会自动完成扫描。所有组件都通过 flax/nnx/init.py 从flax.nnx顶层导出RNNCellBase、LSTMCell、GRUCell、OptimizedLSTMCell、SimpleCell、RNN、Bidirectional而工具函数flip_sequences则需要从flax.nnx.nn.recurrent导入。二、RNNCellBase所有细胞的基础协议源码中RNNCellBaseflax/nnx/nn/recurrent.py定义了所有 RNN 细胞必须实现的三件事initialize_carry(input_shape, rngsNone, carry_initNone) - Carry根据输入形状不含特征维初始化隐藏状态。input_shape的最后一个维度被当作特征维剔除前面的维度全部视为 batch 维这为 ConvLSTM 之类的多维细胞预留了空间。__call__(carry, inputs) - (new_carry, output)单步前向。inputs的所有非最后一维均视为 batch 维。num_feature_axes属性返回特征轴数量普通 RNN 细胞恒为1用于让RNN层正确推断时间轴位置见 docs_nnx/flip/3099-rnnbase-refactor.md 对该属性的讨论。carry 初始化的统一规则所有内置 Cell 的initialize_carry遵循同一套逻辑以LSTMCell为例见 flax/nnx/nn/recurrent.py若调用时未显式传入rngs则回退到构造时保存在self.rngs上的随机源若仍为None抛出ValueError(RNGs must be provided to initialize the cell carry.)。carry 初始化器优先级调用时传入的carry_init 构造时的self.carry_init 默认initializers.zeros_init()。LSTMCell的 carry 是(c, h)二元组cmemory与hhidden形状均为batch_dims (hidden_features,)GRUCell、SimpleCell的 carry 则只是一个h数组。注意__init__中传入carry_init会触发弃用警告。源码提示若把carry_init放在构造参数里两个配置相同但carry_init不同的实例会拥有不同的 graphdef破坏模块图的可序列化一致性。正确的做法是把carry_init传给initialize_carry方法。三、LSTMCell经典长短期记忆细胞数学定义LSTMCellflax/nnx/nn/recurrent.py遵循标准 LSTM 更新规则i σ(W_ii·x W_hi·h b_hi) f σ(W_if·x W_hf·h b_hf) g tanh(W_ig·x W_hg·h b_hg) o σ(W_io·x W_ho·h b_ho) c f ⊙ c i ⊙ g h o ⊙ tanh(c)其中x为当前输入、h为上一时刻输出、c为记忆单元。构造参数与默认值参数默认值语义in_features必填输入特征数hidden_features必填隐藏状态 / 输出特征数gate_fnsigmoid门控激活函数i/f/o 门activation_fntanh记忆更新与输出激活gkernel_initlecun_normal()输入变换核初始化器recurrent_kernel_initmodified_orthogonal隐藏状态变换核初始化器兼容半精度bias_initzeros_init()偏置初始化器dtypeNone计算精度默认从输入与参数推断param_dtypejnp.float32参数初始化精度carry_initNone已弃用请用initialize_carry(carry_init...)promote_dtypedtypes.promote_dtypedtype 提升策略keep_rngsFalse已弃用改由__call__的rngs参数控制rngs必填nnx.Rngs随机源kernel_metadata/recurrent_kernel_metadata/bias_metadata{}参数元数据字典实现要点为什么输入层不设偏置源码将输入变换与循环变换拆成两组Linearself.ii/if_/ig/io与self.hi/hf/hg/hoflax/nnx/nn/recurrent.py。关键细节是输入侧四个 Linear 使用use_biasFalse偏置只放在隐藏侧因为两者最终按门求和ii(x) hi(h)一个门只需一个偏置避免冗余参数。if_采用末尾下划线命名是因为if是 Python 关键字。前向计算中四个门分别由独立的线性变换求和后过激活得到最后按数学公式更新(new_c, new_h)并返回((new_c, new_h), new_h)——即输出就是新的隐藏状态。一个可运行的完整示例import jax.numpy as jnp from flax import nnx module nnx.LSTMCell( in_features3, hidden_features4, rngsnnx.Rngs(0), ) x jnp.ones((2, 3)) # (batch, in_features) carry module.initialize_carry(x.shape, nnx.Rngs(0)) new_carry, y module(carry, x) # y.shape (2, 4)这正是 tests/nnx/nn/recurrent_test.py 中test_basic的用法test_lstm_with_different_dtypes则验证了dtypejnp.bfloat16, param_dtypejnp.bfloat16时输出 dtype 保持一致test_lstm_with_custom_activations展示了用gate_fnjax.nn.relu、activation_fnjax.nn.elu替换默认激活test_lstm_initialize_carry验证了carry_initinitializers.ones可把c、h全部初始化为 1。这些测试全部位于 tests/nnx/nn/recurrent_test.py可作为上手参考。四、OptimizedLSTMCell合并矩阵乘的加速变体OptimizedLSTMCellflax/nnx/nn/recurrent.py与LSTMCell数学定义完全相同、参数完全兼容区别仅在实现策略LSTMCell需要 8 个独立Linear4 输入 × 4 隐藏OptimizedLSTMCell只建两个Lineardense_i输出4 * hidden_features、dense_h输出4 * hidden_features前向时先算y dense_i(inputs) dense_h(h)再用jnp.split(y, 4, axis-1)一次切出 i/f/g/o 四个门flax/nnx/nn/recurrent.py。这种先拼接矩阵、后切分门的方式显著减少了矩阵乘调用次数。源码 docstring 特别说明只要隐藏单元数约不超过 2048该细胞通常比LSTMCell更快在超大隐藏维度下两者的性能差异需要实测。参数兼容意味着可以随时在两种实现间切换而不影响已保存的 checkpoint。五、GRUCell门控循环单元GRUCellflax/nnx/nn/recurrent.py实现标准 GRU 更新规则r σ(W_ir·x b_ir W_hr·h) z σ(W_iz·x b_iz W_hz·h) n tanh(W_in·x b_in r ⊙ (W_hn·h b_hn)) h (1 - z) ⊙ n z ⊙ h其中rreset 门、zupdate 门、n候选隐藏状态。与OptimizedLSTMCell类似的融合策略dense_i输出3 * hidden_features对应 r/z/ndense_h输出3 * hidden_features前向时用jnp.split分别切出xi_r/xi_z/xi_n与hh_r/hh_z/hh_n然后计算r σ(xi_r hh_r)、z σ(xi_z hh_z)、n tanh(xi_n r ⊙ hh_n)最终new_h (1 - z) ⊙ n z ⊙ hflax/nnx/nn/recurrent.py。GRUCell的构造参数与LSTMCell对齐gate_fn、activation_fn、kernel_init、recurrent_kernel_init、bias_init、dtype、param_dtype、promote_dtype、rngs及三组 metadata默认激活仍为 sigmoid tanhcarry 只是一个h数组。测试 tests/nnx/nn/recurrent_test.py 展示了它与RNN配合处理(batch, time, features)输入的用法。六、SimpleCell极简单层细胞可选残差连接SimpleCellflax/nnx/nn/recurrent.py是最简单的 RNN 细胞数学定义为h tanh(W_i·x b_i W_h·h)若residualTrue则变成带残差的形式h tanh(W_i·x b_i W_h·h h)实现上由两个Linear组成dense_h隐藏到隐藏use_biasFalse与dense_i输入到隐藏use_biasTrue加和过activation_fn默认tanh后直接作为新 carry 与输出flax/nnx/nn/recurrent.py。它适合做教学示例或需要极简循环单元的快速原型。七、RNN 层基于 scan 的序列处理器RNN模块flax/nnx/nn/recurrent.py是这一组件的核心它接收任意RNNCellBase实例用flax.nnx.scan沿时间轴扫描整个序列。构造参数RNN( cell, # 任意 RNNCellBase 实例 *, time_majorFalse, # 时间轴是否在首位 return_carryFalse, # 是否同时返回最终 carry reverseFalse, # 是否从右向左处理 keep_orderFalse, # reverse 时是否把输出翻回原顺序 unroll1, # scan 的展开步数 state_axesNone, # scan 状态轴映射 broadcast_rngsNone, # 需跨时间步广播的 RNG 集合如循环 dropout rngsTrue, # carry 初始化随机源 )time_majorFalse时输入形状为(*batch, time, *features)True时形状为(time, *batch, *features)。源码通过inputs.ndim - (cell.num_feature_axes 1)自动推算时间轴位置flax/nnx/nn/recurrent.py因此支持多维 batch 与多维特征如 ConvLSTM。return_carry为True时返回(final_carry, outputs)二元组。reverse/keep_orderreverseTrue时序列被flip_sequences翻转后处理若同时keep_orderTrue输出会再次翻回原始时间顺序用于双向网络对齐。unroll传给nnx.scan的展开系数控制编译与运行时开销的平衡。state_axes控制哪些状态跨时间步传递默认{...: Carry}即所有状态按 Carry 语义传递broadcast_rngs则让指定 RNG 集合如recurrent_dropout在每一步共享同一随机数——这正是循环 dropout 每步相同、输入 dropout 每步不同的实现基础docs_nnx/flip/2396-rnn.md 的 Recurrent Dropout 一节对此有专门讨论。变长序列padding支持调用时传入seq_lengths: Array形状(*batch,)元素为每条序列的真实长度即可处理尾部 padding 的批量序列x jnp.ones((2, 5, 3)) # (batch2, time5, features3) seq_lengths jnp.array([3, 5]) # 第一条有效长度 3第二条 5 final_carry, outputs rnn(x, initial_carrycarry, seq_lengthsseq_lengths)在return_carryTrue且提供seq_lengths的情况下RNN会让 scan 额外输出每一时间步的 carry 历史再用_select_last_carry按seq_lengths精确切出每条序列的真实最终 carryflax/nnx/nn/recurrent.py避免 padding 步污染结果。这一设计正是 docs_nnx/flip/2396-rnn.md 中 Masking 一节所定的方向——只支持尾部连续 padding 的 sequence-length 掩码格式以保证性能。对应测试见test_rnn_with_seq_lengthstests/nnx/nn/recurrent_test.py。完整的 RNN LSTM 序列建模示例import jax.numpy as jnp from flax import nnx cell nnx.LSTMCell(in_features3, hidden_features4, rngsnnx.Rngs(0)) rnn nnx.RNN(cell) # 默认 time-majorFalse x jnp.ones((2, 5, 3)) # (batch, time, features) carry cell.initialize_carry((2, 3), nnx.Rngs(0)) outputs rnn(x, initial_carrycarry) # outputs.shape (2, 5, 4)当不传入initial_carry时RNN会自动调用cell.initialize_carry用输入形状去掉时间轴后的部分完成初始化flax/nnx/nn/recurrent.py。复用性与变量 batch size由于 carry 与输入形状都由调用时动态决定同一个RNN实例可以接受不同 batch size 的输入test_rnn_with_variable_batch_sizetests/nnx/nn/recurrent_test.pytest_rnn_with_unrollunroll2与test_rnn_time_major则分别验证了展开系数与时间优先布局的用法。八、Bidirectional双向编码器Bidirectional模块flax/nnx/nn/recurrent.py接收forward_rnn与backward_rnn两个RNN实例前向编码、反向编码后合并输出前向forward_rnn(inputs, reverseFalse)反向backward_rnn(inputs, reverseTrue, keep_orderTrue)keep_orderTrue保证反向输出仍按原始时间顺序对齐从而能与前向输出按时间步一一合并合并默认merge_fn_concatenate即在最后一维拼接两个方向的输出flax/nnx/nn/recurrent.py。merge_fn可替换为jnp.add等任意二元函数。需要注意两点共享参数警告若forward_rnn is backward_rnn同一个对象两个方向会共享参数源码会打印 warningflax/nnx/nn/recurrent.py。通常应分别构造两个独立的 RNN。返回形式return_carryTrue时返回((carry_forward, carry_backward), outputs)。源码 docstring 中给出了完整示例flax/nnx/nn/recurrent.pyfrom flax import nnx import jax.numpy as jnp forward_rnn nnx.RNN(nnx.GRUCell(in_features3, hidden_features4, rngsnnx.Rngs(0))) backward_rnn nnx.RNN(nnx.GRUCell(in_features3, hidden_features4, rngsnnx.Rngs(0))) layer nnx.Bidirectional(forward_rnnforward_rnn, backward_rnnbackward_rnn) x jnp.ones((2, 3, 3)) out layer(x) print(out.shape) # (2, 3, 8)两方向各 4 维拼接九、flip_sequencespadding 感知的序列翻转工具flip_sequencesflax/nnx/nn/recurrent.py是RNN(reverseTrue)与Bidirectional反向分支的底层支撑解决了一个关键痛点直接翻转带 padding 的批量序列会把 padding 翻到序列开头而正确做法是让 padding 永远留在末尾、仅翻转有效元素。它的签名如下flip_sequences(inputs, seq_lengths, num_batch_dims, time_major) - Array源码 docstring 自带可复现示例flax/nnx/nn/recurrent.py from flax.nnx.nn.recurrent import flip_sequences from jax import numpy as jnp inputs jnp.array([[1, 0, 0], [2, 3, 0], [4, 5, 6]]) lengths jnp.array([1, 2, 3]) flip_sequences(inputs, lengths, 1, False) Array([[1, 0, 0], [3, 2, 0], [6, 5, 4]], dtypeint32)实现思路构造倒序索引idxs arange(max_steps-1, -1, -1)与展宽后的seq_lengths相加后对max_steps取模得到翻转后各元素应去的位置最后用jnp.take_along_axis完成取值。若seq_lengths为None则直接jnp.flip。该函数需要从flax.nnx.nn.recurrent子模块导入未在flax.nnx顶层导出。十、自定义 Cell扩展 RNN 的三种正确姿势方式一继承RNNCellBase协议RNN只要求 cell 具备__call__(carry, inputs)、initialize_carry(...)与num_feature_axes属性。tests/nnx/nn/recurrent_test.py 的test_rnn_with_custom_cell演示了一个拼接输入与隐藏状态后过一个 Linear tanh的自定义 Cell 如何在RNN中直接工作class CustomRNNCell(nnx.Module): def __init__(self, in_features, hidden_features, rngs): self.dense nnx.Linear( in_featuresin_features hidden_features, out_featureshidden_features, rngsrngs, ) def __call__(self, carry, inputs): h carry x jnp.concatenate([inputs, h], axis-1) new_h jax.nn.tanh(self.dense(x)) return new_h, new_h def initialize_carry(self, input_shape, rngs): return jnp.zeros((input_shape[0], self.hidden_features)) property def num_feature_axes(self): return 1方式二子类化内置 Cell 叠加循环 dropoutflax/nnx/nn/recurrent.py 的test_recurrent_dropout展示了在OptimizedLSTMCell子类中加入recurrent_dropoutRNG 集合并用RNN(..., broadcast_rngsrecurrent_dropout)让该随机源跨时间步共享每步用同一随机掩码从而实现标准的循环 dropout测试还断言了model.lstm.cell.recurrent_dropout.rngs.count在调用前后从 0 变为 1验证 RNG 被正确消费。方式三利用 StateAxes 精细控制 scan 语义RNN构造参数中的state_axes默认把全部状态按iteration.Carry传递结合flax.nnx的split_rngs装饰器源码中RNN.__call__内部使用nnx.split_rngs(splits1, onlyself.broadcast_rngs, squeezeTrue)保证每次调用都获得独立随机源见 flax/nnx/nn/recurrent.py你可以在不改动 Cell 的情况下调整状态与随机源的扫描行为。十一、与 Linen 的等价性NNX 与旧 API 的可迁移性测试文件用两个用例直接对比了 NNX 与 Linen 的实现test_lstm_equivalence_with_flax_linentests/nnx/nn/recurrent_test.py把linen.LSTMCell的参数ii/if_/ig/io、hi/hf/hg/ho的 kernel/bias逐门拷入nnx.LSTMCell两者输出与 carry 的数值误差在1e-5以内。test_rnn_equivalence_with_flax_linentests/nnx/nn/recurrent_test.py对nnx.RNN(cell)与linen.RNN(cell)施加相同参数后序列输出同样在1e-5误差内一致。这从数值层面印证了 NNX 循环组件与 Linen 在数学上完全等价便于从旧代码迁移Linen 中LSTMCell(featuresN)对应 NNX 的LSTMCell(in_features..., hidden_featuresN)参数命名与结构保持兼容。十二、路径速查与进一步阅读API 索引文档docs_nnx/api_reference/flax.nnx/nn/recurrent.rst核心源码约 1150 行含全部数学公式 docstringflax/nnx/nn/recurrent.py单元测试覆盖各 Cell、RNN 各参数、padding、dropout、Linen 等价性tests/nnx/nn/recurrent_test.py设计提案三层抽象与 masking 取舍docs_nnx/flip/2396-rnn.mdCell 基类重构提案num_feature_axes由来docs_nnx/flip/3099-rnnbase-refactor.md顶层导出确认flax.nnx命名空间flax/nnx/init.py序列模型端到端示例seq2seqexamples/seq2seq/models.py综上Flax NNX 的循环神经网络组件以Cell RNN 层 Bidirectional的三层设计把单步计算、序列扫描与双向融合清晰解耦同时通过seq_lengths、reverse、keep_order、unroll、broadcast_rngs等参数覆盖了变长序列、反向编码、循环 dropout 等真实训练场景。无论是快速原型还是生产级序列模型你都可以直接从flax.nnx顶层导入这些组件或基于RNNCellBase协议快速定制自己的循环单元。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门