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

Flax Dropout 实战指南:在 Linen 模型中正确启用与关闭随机失活

Flax Dropout 实战指南在 Linen 模型中正确启用与关闭随机失活【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax这篇指南以 docs/guides/training_techniques/dropout.rst 为主线系统讲解如何在 FlaxJAX 的神经网络库中使用flax.linen.Dropout实现随机失活stochastic regularization正则化技术。你将掌握 PRNG 密钥的拆分、模型定义与初始化、训练/评估切换、以及如何将 dropout 无缝整合进TrainState训练循环并了解仓库中 WMT、SST-2、Seq2Seq 等真实示例的落地用法。Dropout 是一种随机正则化技术它会在训练过程中随机移除置零网络中的隐藏单元与可见单元从而抑制过拟合、增强模型泛化能力。在 Flax 中这一切通过nn.Dropout层配合显式的 PRNG伪随机数生成器密钥流完成。本文所有示例均基于 Flax Linen APIimport flax.linen as nn并全程给出无 Dropout与有 Dropout两版代码对照。1. 前置准备开始前确保已安装 Flax 及其依赖JAX、Optax并导入所需模块import flax.linen as nn import jax.numpy as jnp import jax import optax后续所有代码示例都依赖这些导入。本文中的模型将以MyModel为例其原始版本只是一个不带任何正则化的单层 Dense 网络。2. 拆分 PRNG 密钥Dropout 的随机性来源Dropout 是随机操作因此需要伪随机数生成器PRNG状态。Flax 使用 JAX 的可拆分PRNG 密钥——它具备对神经网络而言非常理想的属性可复制、可派生、可组合。JAX 的密钥通过jax.random.key(seed0)创建并可用jax.random.split()分裂出多个子密钥。2.1 从两个密钥到三个密钥无 Dropout 时只需两个密钥一个主密钥后续可再拆分一个用于初始化参数的params密钥。加入 Dropout 后需要额外分裂出第三个密钥专门用于 Dropout# 无 Dropout root_key jax.random.key(seed0) main_key, params_key jax.random.split(keyroot_key) # 有 Dropout root_key jax.random.key(seed0) main_key, params_key, dropout_key jax.random.split(keyroot_key, num3) #!关键差异在于num3它一次性把根密钥拆成三个互不重叠的子密钥。其中dropout_key将被喂给dropoutPRNG 流。2.2 带名字的 PRNG 流PRNG streamsFlax 的一个重要设计是PRNG 流是带名字的named PRNG streams。在Module内你可以通过流的名字在稍后取用密钥。例如流params用于参数初始化流dropout用于nn.Dropout生成随机掩码。在调用module.apply(variables, x, rngs{dropout: dropout_key})时你实际上是在按名字为这条流注入密钥。这种命名机制让同一模型中可以并存多个独立随机源如 dropout 与 attention dropout、采样等互不干扰。3. 定义带 Dropout 的模型3.1 核心步骤创建一个带 dropout 的模型只需三步继承flax.linen.Module所有神经网络模块的基类在模块内部使用nn.Dropout添加 dropout 层以关键字参数方式传递deterministic—— 可以在构造 Module 时传入也可以在调用init()/apply()时传入底层由flax.linen.module.merge_param负责合并两处取值。deterministic是布尔值其语义如下deterministic行为False训练输入以rate概率被置零masked其余输入按1 / (1 - rate)缩放从而保持输入均值不变True推理/评估不应用掩码输入原样返回dropout 关闭3.2 代码对照# 无 Dropout class MyModel(nn.Module): num_neurons: int nn.compact def __call__(self, x): x nn.Dense(self.num_neurons)(x) return x # 有 Dropout class MyModel(nn.Module): num_neurons: int nn.compact def __call__(self, x, training: bool): #! x nn.Dense(self.num_neurons)(x) # 设置 dropout 层rate 为 50% # 当 deterministicTrue 时 dropout 被关闭。 x nn.Dropout(rate0.5, deterministicnot training)(x) #! return x这里采用了一个非常常见的模式父级 Module 接收布尔参数training或train再通过deterministicnot training把训练/推理语义传递给 Dropout。在其他框架PyTorch、TensorFlow/Keras中这一开关通常由可变状态或调用标志完成例如torch.nn.Module.eval()或tf.keras.Model的training参数Flax 则把它显式化到了函数签名里让训练/推理差异在类型层面就一目了然。3.3 底层实现make_rng与merge_param从源码 flax/linen/stochastic.py 可以看出nn.Dropout的实现要点class Dropout(Module): rate: float broadcast_dims: Sequence[int] () deterministic: bool | None None rng_collection: str dropout compact def __call__(self, inputs, deterministicNone, rngNone): deterministic merge_param(deterministic, self.deterministic, deterministic) if (self.rate 0.0) or deterministic: return inputs # 防止 1.0 边沿情况下的梯度 NaN if self.rate 1.0: return jnp.zeros_like(inputs) keep_prob 1.0 - self.rate if rng is None: rng self.make_rng(self.rng_collection) broadcast_shape list(inputs.shape) for dim in self.broadcast_dims: broadcast_shape[dim] 1 mask random.bernoulli(rng, pkeep_prob, shapebroadcast_shape) mask jnp.broadcast_to(mask, inputs.shape) return lax.select(mask, inputs / keep_prob, jnp.zeros_like(inputs))几个值得注意的细节rate 0.0与deterministic短路当rate0.0时直接返回输入不产生任何开销deterministicTrue同理。rate 1.0特判若rate恰为 1直接返回全零张量避免keep_prob0导致梯度中出现 NaN源码注释明确写明Prevent gradient NaNs in 1.0 edge-case。make_rng保证可复现性Flax 通过Module.make_rng从 PRNG 流中切出一个全新密钥。make_rng保证每次调用都返回唯一的新密钥内部按调用序列推进状态因此同一rngs输入下模型行为完全可复现。这是 Flax 隐式处理随机密钥的方式——你无需手动管理计数器。merge_param合并两处取值deterministic可以在构造时传入也可以在调用时传入merge_param负责判断哪一侧生效调用时的值优先于构造值。broadcast_dims则用于让某些维度共享同一掩码例如序列模型中对整条样本施加同一掩码。补充说明源码中还支持broadcast_dims参数用于指定哪些维度共享同一 dropout 掩码掩码在该维度上广播。例如在 Transformer 中常对 batch 维与序列维之外的特征维做整体丢弃。此外rng_collection默认为dropout即默认从名为dropout的流取密钥。4. 初始化模型训练开关与参数提取创建好模型后按以下顺序初始化实例化模型在init()调用中传入trainingFalse即关闭 dropout从 variable dictionary 中取出params。注意初始化阶段不应启用 dropout因为此时只需要确定参数形状不需要随机掩码。# 无 Dropout my_model MyModel(num_neurons3) x jnp.empty((3, 4, 4)) variables my_model.init(params_key, x) params variables[params] # 有 Dropout my_model MyModel(num_neurons3) x jnp.empty((3, 4, 4)) # 用 trainingFalse 关闭 dropout即 deterministicTrue variables my_model.init(params_key, x, trainingFalse) #! params variables[params]与无 Dropout 版本相比唯一区别是在init()中多传了trainingFalse。如果希望启用 dropout就必须提供training或train参数——这一点在调用apply()时同样成立。5. 训练前向传播开启 Dropout 并注入 PRNG使用module.apply()执行前向传播时需要同时做两件事向apply()传trainingTrue使 Dropout 生效提供 PRNG 密钥以填充dropout流——通过rngs{dropout: dropout_key}传入。# 无 Dropout无需传 training 与 rngs y my_model.apply({params: params}, x) # 有 DropouttrainingTrue 开启 dropout并注入 dropout 流密钥 y my_model.apply({params: params}, x, trainingTrue, rngs{dropout: dropout_key}) #!这里rngs是名字→密钥的字典映射。因为密钥名与Dropout的rng_collection默认为dropout一一对应Flax 会自动把dropout_key路由到每一处nn.Dropout。评估/推理时直接使用与上面相同的代码但以trainingFalse调用此时不需要传递 RNG——deterministicTrue会短路随机路径模型以完全确定性的方式输出。6. 把 Dropout 接入TrainState训练循环在实际训练中Dropout 需要每个训练步都产生新的随机掩码。Flax 的常见模式是用一个 dataclass如flax.training.train_state.TrainState封装整个训练状态参数、优化器状态等并以单一state: TrainState参数传入训练步函数。要让 dropout 融入这一模式需要在自定义TrainState子类中新增key字段把dropout_key传入TrainState.create()。6.1 扩展 TrainStatefrom flax.training import train_state # 无 Dropout state train_state.TrainState.create( apply_fnmy_model.apply, paramsparams, txoptax.adam(1e-3) ) # 有 Dropout class TrainState(train_state.TrainState): #! key: jax.Array #! state TrainState.create( #! apply_fnmy_model.apply, paramsparams, keydropout_key, #! txoptax.adam(1e-3) )key字段会随TrainState一起被 JAX 作为 pytree 处理天然支持jax.jit与梯度更新apply_gradients会自动保留该字段。6.2 训练步每个 step 生成新的 dropout 密钥在train_step中从state.key派生每个训练步专用的密钥。Flax 提供了两种方式jax.random.split()分裂出一个可复用的新密钥之后可以继续复用jax.random.fold_in()一般更快它 1) 折叠进唯一的数据如state.step2) 能生成更长的 PRNG 流序列。推荐在训练循环中使用fold_in因为每步都折叠进当前步号天然保证密钥唯一且随步数推进。# 无 Dropout jax.jit def train_step(state: train_state.TrainState, batch): def loss_fn(params): logits state.apply_fn( {params: params}, xbatch[image], ) loss optax.softmax_cross_entropy_with_integer_labels( logitslogits, labelsbatch[label]) return loss, logits grad_fn jax.value_and_grad(loss_fn, has_auxTrue) (loss, logits), grads grad_fn(state.params) state state.apply_gradients(gradsgrads) return state # 有 Dropout jax.jit def train_step(state: TrainState, batch, dropout_key): #! dropout_train_key jax.random.fold_in(keydropout_key, datastate.step) #! def loss_fn(params): logits state.apply_fn( {params: params}, xbatch[image], trainingTrue, #! rngs{dropout: dropout_train_key} #! ) loss optax.softmax_cross_entropy_with_integer_labels( logitslogits, labelsbatch[label]) return loss, logits grad_fn jax.value_and_grad(loss_fn, has_auxTrue) (loss, logits), grads grad_fn(state.params) state state.apply_gradients(gradsgrads) return state可以看到训练步的改动集中在三点train_step多接收一个dropout_key参数、用fold_in生成dropout_train_key、在loss_fn内部的前向调用中传入trainingTrue与rngs{dropout: dropout_train_key}。由于密钥随state.step折叠每一步掩码都不同且整个train_step仍可被jax.jit编译。评估循环中只需以trainingFalse或无rngs调用state.apply_fn即可得到确定性的推理输出。7. 仓库真实示例WMT、SST-2 与 Seq2Seq7.1 WMT 机器翻译Transformer 中的 dropout 与 attention dropoutexamples/wmt/models.py 中的MlpBlock展示了标准用法——在 MLP 的两个 Dense 层之后各接一个nn.Dropoutx nn.Dense(config.mlp_dim, ...)(inputs) x nn.relu(x) x nn.Dropout(rateconfig.dropout_rate)(x, deterministicconfig.deterministic) output nn.Dense(actual_out_dim, ...)(x) output nn.Dropout(rateconfig.dropout_rate)(output, deterministicconfig.deterministic)同时Transformer 的注意力模块使用nn.MultiHeadDotProductAttention配置了dropout_rateconfig.attention_dropout_rate与broadcast_dropoutFalse见 examples/wmt/models.py实现了注意力 dropout。整个模型的deterministic由config.deterministic统一控制——训练配置下为False推理配置下为True。7.2 SST-2 文本分类自定义 WordDropout 层examples/sst2/models.py 定义了一个WordDropout模块它本质上是nn.Dropout的变体但允许指定被丢弃元素的替换值unk_idx用于对一批输入 ID 施加词级 dropoutclass WordDropout(nn.Module): dropout_rate: float unk_idx: int deterministic: bool | None None nn.compact def __call__(self, inputs, deterministicNone): deterministic nn.module.merge_param(deterministic, self.deterministic, deterministic) if deterministic or self.dropout_rate 0.0: return inputs rng self.make_rng(dropout) mask jax.random.bernoulli(rng, pself.dropout_rate, shapeinputs.shape) return jnp.where(mask, jnp.array([self.unk_idx]), inputs)这个自定义层展示了两个可复用的模式使用nn.module.merge_param合并构造/调用两处的deterministic取值在自定义模块内部直接调用self.make_rng(dropout)从dropout流取密钥。随后它在Embedderexamples/sst2/models.py中与标准nn.Dropout配合使用先做词级 dropout把部分词替换为unk_idx再做嵌入后的特征 dropout两者都接收同一个deterministic参数。7.3 Seq2Seqmake_rng在解码器中的另一种应用make_rng的威力不止于 dropout。examples/seq2seq/models.py 的DecoderLSTM中解码器通过self.make_rng(lstm)获取密钥用jax.random.categorical从 logits 采样预测 tokencategorical_rng self.make_rng(lstm) predicted_token jax.random.categorical(categorical_rng, logits)这说明 Flax 的命名 PRNG 流机制是一种通用抽象——任何需要随机性的模块dropout、采样、正则化都可以声明自己的流名dropout、lstm等并在apply(rngs{...})时统一注入。这也解释了为何 Flax 无需像部分框架那样依赖全局可变随机状态一切随机源都显式、可复现、可组合。8. 小结与最佳实践回顾全文在 Flax Linen 中使用 Dropout 的核心要点密钥先行用jax.random.split(key, num3)为 params 与 dropout 分别准备密钥模型内声明nn.Dropout(rate0.5, deterministicnot training)是标准写法deterministic必须按关键字传入初始化关、训练开init()时传trainingFalse训练apply()时传trainingTrue并附rngs{dropout: dropout_key}训练循环按步换钥在TrainState中保存key每个train_step用jax.random.fold_in(key, datastate.step)派生新密钥更快、流更长评估时完全不传rngs底层语义rate是丢弃概率而非保留概率deterministicTrue或rate0.0时直接短路返回输入rate1.0时返回全零以避免梯度 NaN命名流复用rng_collection默认dropout自定义随机模块可通过self.make_rng(任意流名)取钥实现完全可复现的随机行为。需要进一步探索时可阅读 flax/linen/stochastic.py 的完整实现、examples/wmt/models.py 的 Transformer 应用以及 docs/guides/training_techniques/dropout.rst 的原始文档对 NNX 新 API 感兴趣的读者也可以对照 flax/nnx/nn/stochastic.py 中nnx.Dropout的实现了解deterministic优先级与nnx.view的等价控制方式。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

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

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