从 flax.nn 迁移到 flax.linen:Flax 0.4.0 代码库升级实战指南
从 flax.nn 迁移到 flax.linenFlax 0.4.0 代码库升级实战指南【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax自 Flax v0.4.0 起旧版flax.nn模块已从库中移除取而代之的是全新的 Linen APIflax.linen。本文基于官方升级指南 docs/guides/converting_and_upgrading/linen_upgrade_guide.rst系统梳理从旧 API 迁移到 Linen 的每一步改造要点模块定义方式、子模块组合、参数与状态管理、顶层训练循环、检查点加载与随机性处理等。读完本文你将能把自己的 Flax 代码库一次性、无痛地升级到 Linen并理解其底层设计动机如变量集合、可变性控制、RNG 分流在源码中的落地方式。升级背景为什么flax.nn变成了flax.linen旧 API 中模块通过继承base.Module并重写apply方法来实现而 Linen 中模块继承nn.Module本质是一个 dataclass通过__call__方法定义前向逻辑。最直观的变化是 import 语句# 旧写法 from flax import nn # 新写法 from flax import linen as nn这一行替换背后是整个模块设计哲学的转变模块实例不再在apply时临时创建而是像普通 Python 对象一样被构造、共享和传递。从源码看nn.Module在 flax/linen/module.py 中定义其compact装饰器module.py#L477-L502只是给方法打上fun.compact True标记允许在方法体内内联定义子模块——底层实现仍然复用 Scope 机制但用户接口已经彻底对象化。定义简单 Flax Modulesapply到__call__的转换这是迁移中最频繁的改动。旧代码把配置参数放在apply的参数列表里新代码把它们提升为 dataclass 字段前向方法从apply改名为__call__。官方指南给出了一个 Dense 层的完整对照# ---------------- 旧 Flax ---------------- from flax import nn class Dense(base.Module): def apply(self, inputs, features, use_biasTrue, kernel_initdefault_kernel_init, bias_initinitializers.zeros_init()): kernel self.param(kernel, (inputs.shape[-1], features), kernel_init) y jnp.dot(inputs, kernel) if use_bias: bias self.param( bias, (features,), bias_init) y y bias return y # ---------------- Linen ---------------- from flax import linen as nn class Dense(nn.Module): features: int use_bias: bool True kernel_init: Callable[[PRNGKey, Shape, Dtype], Array] default_kernel_init bias_init: Callable[[PRNGKey, Shape, Dtype], Array] initializers.zeros_init() nn.compact def __call__(self, inputs): kernel self.param(kernel, self.kernel_init, (inputs.shape[-1], self.features)) y jnp.dot(inputs, kernel) if self.use_bias: bias self.param( bias, self.bias_init, (self.features,)) y y bias return y逐条对应关系如下import 替换from flax import nn→from flax import linen as nn。参数移到 dataclass 属性apply的位置参数features、use_bias等变成类属性建议加类型注解不需要类型时可用Any跳过。方法改名apply→__call__并用nn.compact装饰可选。只有被compact装饰的方法可以在方法体内直接内联定义子模块且每个模块最多只能有一个compact方法源码中 compact 的实现直接以标记位实现setup_or_nncompact文档中也提到多方法场景会触发MultipleMethodsCompactError。另一种方式是定义setup方法二者的取舍可参考 docs/guides/flax_fundamentals/setup_or_nncompact.rst。属性访问方法体内通过self.attr读取 dataclass 字段如self.features。参数初始化顺序调整self.param的签名变为param(name, init_fn, *init_args)形状参数移到初始化函数之后初始化函数可接受任意参数列表。从 module.py 中 param 的源码 可以看到init_fn的第一个参数是自动注入的 PRNG key不需要显式传入。在模块内使用其他模块构造函数返回实例而非输出旧 API 中nn.Dense(x, 500)直接返回前向结果Linen 中模块构造函数返回模块实例需要再调用一次才能得到输出。官方对照如下# ---------------- 旧 Flax ---------------- class Encoder(nn.Module): def apply(self, x): x nn.Dense(x, 500) x nn.relu(x) z nn.Dense(x, 500, namelatents) return z # ---------------- Linen ---------------- class Encoder(nn.Module): nn.compact def __call__(self, x): x nn.Dense(500)(x) x nn.relu(x) z nn.Dense(500, namelatents)(x) return z两个关键点模块实例可以像普通 Python 对象一样被共享复用替代旧 API 的.shared()机制。所有模块构造函数都可以通过name显式命名可选。不传name时子模块按类名_序号自动命名——这正是后续加载 pre-Linen 检查点时需要注意的命名差异来源见下文加载 pre-Linen 检查点一节。共享子模块与多方法模块setup的引入当一个模块需要多个前向方法如自编码器的__call__和generate、或需要把子模块预定义后复用如 ResNet 中的共享 Block时用setup替代旧 API 的_create_submodules模式。官方示例# ---------------- 旧 Flax ---------------- class AutoEncoder(nn.Module): def _create_submodules(self): return Decoder.shared(nameencoder) def apply(self, x, z_rng, latents20): decoder self._create_decoder() z Encoder(x, latents, nameencoder) return decoder(z) nn.module_method def generate(self, z, **unused_kwargs): decoder self._create_decoder() return nn.sigmoid(decoder(z)) # ---------------- Linen ---------------- class AutoEncoder(nn.Module): latents: int 20 def setup(self): self.encoder Encoder(self.latents) self.decoder Decoder() def __call__(self, x): z self.encoder(x) return self.decoder(z) def generate(self, z): return nn.sigmoid(self.decoder(z))要点说明用setup替代__init____init__已被 dataclass 机制占用Flax 会在模块准备好使用后自动调用setup。所有模块都可以用setup风格不用compact但官方更推荐compact因为它把子模块的定义和使用同位放置在存在循环或条件分支时代码更清晰。子模块共享在初始化时把子模块赋值给self.encoder它就自动以属性名encoder命名与 PyTorch 的约定一致。对同一属性重复赋值即可实现子模块共享。不内联定义子模块时无需compact本例所有子模块都在setup中定义因此__call__不添加装饰器。附加方法generate就是普通 Python 方法可以被顶层apply(method...)或init(method...)调用。Module.partial的替代使用标准库functools.partial旧 API 用nn.Conv.partial(biasFalse)预绑定模块超参数Linen 直接使用 Python 标准库的functools.partial。官方 ResNet 示例# ---------------- 旧 Flax ---------------- class ResNet(nn.Module): ResNetV1. def apply(self, x, stage_sizes, num_filters64, trainTrue): conv nn.Conv.partial(biasFalse) norm nn.BatchNorm.partial( use_running_averagenot train, momentum0.9, epsilon1e-5) x conv(x, num_filters, (7, 7), (2, 2), padding[(3, 3), (3, 3)], nameconv_init) x norm(x, namebn_init) # [...] return x # ---------------- Linen ---------------- from functools import partial class ResNet(nn.Module): ResNetV1. stage_sizes: Sequence[int] num_filters: int 64 train: bool True nn.compact def __call__(self, x): conv partial(nn.Conv, use_biasFalse) norm partial(nn.BatchNorm, use_running_averagenot self.train, momentum0.9, epsilon1e-5) x conv(self.num_filters, (7, 7), (2, 2), padding[(3, 3), (3, 3)], nameconv_init)(x) x norm(namebn_init)(x) # [...] return xpartial返回的仍然是模块构造函数因此用法保持不变调用时返回实例再调用。注意BatchNorm的momentum0.9, epsilon1e-5与 Linen 内置 BatchNorm 默认值 一致use_running_averagenot self.train的写法将train字段绑定进超参数迁移后无需重复传参。顶层训练代码模式从nn.Model到TrainState旧 API 中nn.Model把参数与模型绑定在一起配合optim.Momentum构造优化器。Linen 不再提供Model抽象而是直接传递参数通常封装在一个TrainState对象中该对象可以直接传入 JAX 变换jax.jit/jax.grad等。官方对照# ---------------- 旧 Flax ---------------- def create_model(key): _, initial_params CNN.init_by_shape( key, [((1, 28, 28, 1), jnp.float32)]) model nn.Model(CNN, initial_params) return model def create_optimizer(model, learning_rate): optimizer_def optim.Momentum(learning_ratelearning_rate) optimizer optimizer_def.create(model) return optimizer def loss_fn(model): logits model(batch[image]) one_hot jax.nn.one_hot(batch[label], num_classes10) loss -jnp.mean(jnp.sum(one_hot_labels * batch[label], axis-1)) return loss, logits # ---------------- Linen ---------------- def create_train_state(rng, config): variables CNN().init(rng, jnp.ones([1, 28, 28, 1])) params variables[params] tx optax.sgd(config.learning_rate, config.momentum) return train_state.TrainState.create( apply_fnCNN.apply, paramsparams, txtx) def loss_fn(params): logits CNN().apply({params: params}, batch[image]) one_hot jax.nn.one_hot(batch[label], 10) loss jnp.mean(optax.softmax_cross_entropy(logitslogits, labelsone_hot)) return loss, logits迁移要点弃用Model抽象参数直接传递TrainState是对参数 优化器状态的轻量封装TrainState源码见 flax/training/train_state.py包含step、apply_fn、params、tx、opt_state五个字段apply_gradients内部调用tx.update与optax.apply_updates。它的apply_fn通常就是model.apply。初始化参数用init/init_with_output构造模块实例后调用init(rng, 具体输入)。官方明确没有移植init_by_shape因为它按形状求值却返回真实数值语义混乱。因此 Linen 要求传入具体数值做初始化并强烈建议用jax.jit包裹初始化以跳过完整前向传播的开销init实现见 module.py#L2316-L2351。参数是变量集合之一Linen 把参数推广为变量。变量是嵌套字典顶层键是不同变量集合名params只是其中一个集合。详见 docs/api_reference/flax.linen/variable.rst。优化器推荐 Optax迁移到 Optax 的细节参考 docs/guides/converting_and_upgrading/optax_update_guide.rst。推理在顶层构造一个模块实例只是构造属性的轻量包装几乎零成本调用其apply方法内部会调用__call__。非可训练变量状态模块内部的定义方式BatchNorm 的running_mean/running_var是典型非可训练状态。旧 API 用self.state(...)Linen 用self.variable(collection, name, init_fn, *init_args)# ---------------- 旧 Flax ---------------- class BatchNorm(nn.Module): def apply(self, x): # [...] ra_mean self.state( mean, (x.shape[-1], ), initializers.zeros_init()) ra_var self.state( var, (x.shape[-1], ), initializers.ones_init()) # [...] # ---------------- Linen ---------------- class BatchNorm(nn.Module): def __call__(self, x): # [...] ra_mean self.variable( batch_stats, mean, initializers.zeros_init(), (x.shape[-1], )) ra_var self.variable( batch_stats, var, initializers.ones_init(), (x.shape[-1], )) # [...]self.variable的第一个参数是变量集合名——params是唯一始终可用的集合self.param就是variable(params, ...)的简写见 module.py#L1677-L1784。不同集合在顶层训练代码中可被区别对待为可变或不可变在模块内部使用 JAX 变换时每个集合也可以被单独处理通过flax.linen的提升变换。非可训练变量顶层训练代码模式mutable控制官方对照展示了训练时更新 batch 统计、评估时只读的完整模式# ---------------- 旧 Flax ---------------- # 初始化参数与状态 def initial_model(key, init_batch): with nn.stateful() as initial_state: _, initial_params ResNet.init(key, init_batch) model nn.Model(ResNet, initial_params) return model, init_state # 训练时更新 batch 统计 def loss_fn(model, model_state): with nn.stateful(model_state) as new_model_state: logits model(batch[image]) # [...] # 评估时只读 batch 统计 def eval_step(model, model_state, batch): with nn.stateful(model_state, mutableFalse): logits model(batch[image], trainFalse) return compute_metrics(logits, batch[label]) # ---------------- Linen ---------------- # 初始化变量 ({param: ..., batch_stats: ...}) def initial_variables(key, init_batch): return ResNet().init(key, init_batch) # 训练时更新 batch 统计 def loss_fn(params, batch_stats): variables {params: params, batch_stats: batch_stats} logits, new_variables ResNet(trainTrue).apply( variables, batch[image], mutable[batch_stats]) new_batch_stats new_variables[batch_stats] # [...] # 评估时只读 batch 统计 def eval_step(params, batch_stats, batch): variables {params: params, batch_stats: batch_stats} logits ResNet(trainFalse).apply( variables, batch[image], mutableFalse) return compute_metrics(logits, batch[label])四个关键机制init返回完整变量字典如{params: ..., batch_stats: ...}参见变量文档 docs/api_reference/flax.linen/variable.rst。旧 API 的nn.stateful()上下文管理器被彻底移除。手动合并集合把params与batch_stats拼成变量字典传给apply。mutable[batch_stats]声明训练中batch_stats集合可变。此时module.apply的返回值变成二元组(output, new_variables)从中取new_variables[batch_stats]即可获得更新后的统计量。mutable接受 bool / str / list 三种形式bool 表示全部/全不可变str 为单个集合名list 为集合名列表详见 module.py 中 apply 的签名与文档。mutableFalse评估时强制所有集合只读若误用了训练模式下的 BatchNorm 会直接报错。因为没有任何集合被修改返回值就只是输出本身。加载 pre-Linen 检查点子模块命名差异与convert_pre_linen大部分 Linen 模块可以直接加载 pre-Linen 权重但有一个命名差异必须处理旧 API 中子模块按出现顺序全局递增编号与类无关Linen 改为按模块类分别计数。官方示例pre-Linen{Conv_0: { ... }, Dense_1: { ... } }Linen{Conv_0: { ... }, Dense_0: { ... } }迁移工具位于 flax/training/checkpoints.py 的convert_pre_linen。从源码可以看到其实现逻辑对参数 pytree 按键做自然排序用正则MODULE_NUM_RE匹配类名_序号形式的键然后按类名分别重新计数并递归处理子层同时它会安全地跳过已是 Linen 格式的 pytree可直接对任意已加载检查点调用。典型用法from flax.training import checkpoints params checkpoints.convert_pre_linen(pre_linen_params)官方还提示该工具也适用于转换 pre-Linen 的其他变量集合但旧集合是扁平结构需要先用flax.traverse_util.unflatten_dict展开为嵌套字典再转换batch_stats checkpoints.convert_pre_linen(flax.traverse_util.unflatten_dict({ tuple(k.split(/)[1:]): v for k, v in pre_linen_model_state.as_dict().items() }))随后即可构造 Linen 变量字典variables {params: params, batch_stats: batch_stats}随机性从nn.stochastic上下文到 RNG 流make_rng与rngs旧 API 通过nn.stochastic(dropout_rng)上下文管理器注入随机源Linen 中随机源显式通过apply(..., rngs...)传递且 RNG 有种类kinds。官方 Dropout 对照# ---------------- 旧 Flax ---------------- def dropout(inputs, rate, deterministicFalse): keep_prob 1. - rate if deterministic: return inputs else: mask random.bernoulli( make_rng(), pkeep_prob, shapeinputs.shape) return lax.select( mask, inputs / keep_prob, jnp.zeros_like(inputs)) def loss_fn(model, dropout_rng): with nn.stochastic(dropout_rng): logits model(inputs) # ---------------- Linen ---------------- class Dropout(nn.Module): rate: float nn.compact def __call__(self, inputs, deterministicFalse): keep_prob 1. - self.rate if deterministic: return inputs else: mask random.bernoulli( self.make_rng(dropout), pkeep_prob, shapeinputs.shape) return lax.select( mask, inputs / keep_prob, jnp.zeros_like(inputs)) def loss_fn(params, dropout_rng): logits Transformer().apply( {params: params}, inputs, rngs{dropout: dropout_rng})要点RNG 种类kindsself.make_rng(dropout)中的dropout是 RNG 流名称。不同种类在 JAX 变换中可以区别对待——例如序列模型中每个时间步是共享同一个 dropout mask 还是各自独立。从 module.py 中 make_rng 的源码 看每次调用都会从对应 RNG 序列中分裂出一个新 key保证完全可复现。显式传入rngsapply/init接受rngs{dropout: key}字典替代旧上下文管理器。评估时不传 RNG一旦误用非确定性 dropoutself.make_rng(dropout)就会抛错。源码还说明如果调用了一个未被传入的 RNG 流名称会默认回退到params流见 apply 的文档直接传单个PRNGKey等价于{params: key}。提升变换Lifted transformationsLinen 中不再直接使用 JAX 变换而是使用提升变换lifted transforms——即作用于 Flax Module 的 JAX 变换例如nn.scan、nn.vmap、nn.jit、nn.remat等。它们能正确处理模块内的变量集合与 RNG 流例如让nn.scan决定序列各时间步共享还是各自独立的 RNG。设计原理可参考仓库中的设计笔记 docs/developer_notes/lift.md。官方指南中关于jax.scan_in_dim旧与nn.scan新的对照示例仍标记为 TODO迁移时建议直接参考该设计文档与 docs/api_reference/flax.linen/transformations.rst 的 API 说明。迁移自检清单完成迁移后可按以下清单逐项确认代码库已与 Linen 完全对齐所有from flax import nn已替换为from flax import linen as nn模块继承nn.Module配置参数改为带类型注解的 dataclass 字段前向方法统一为__call__单方法用compact多方法用setup子模块组合使用构造实例再调用模式name显式命名关键子模块Module.partial已替换为functools.partial顶层训练改为initTrainState Optax 优化器参数直接传入jax.grad/jax.jitself.state(...)已改为self.variable(collection, ...)训练/评估分别用mutable[batch_stats]与mutableFalse旧检查点已通过checkpoints.convert_pre_linen转换命名扁平集合先unflatten_dict随机性改用self.make_rng(kind)apply(..., rngs{...})评估时不传 RNG。参考实现仓库中的 MNIST 示例 examples/mnist/train.py、ImageNet 训练 examples/imagenet/train.py 以及 seq2seq examples/seq2seq/train.py 都是完整的 Linen 迁移后代码范例可直接对照阅读相关单元测试如 tests/linen/linen_module_test.py 与 tests/linen/linen_transforms_test.py覆盖了init/apply/mutable/RNG 等核心行为可作为迁移正确性的验证参照。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考