ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Flax Linen 入门实战:用 init/apply、setup/compact 与 JAX 变换构建模块化神经网络

Flax Linen 入门实战:用 init/apply、setup/compact 与 JAX 变换构建模块化神经网络 Flax Linen 入门实战用 init/apply、setup/compact 与 JAX 变换构建模块化神经网络【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax本文是 FlaxJAX 生态的神经网络库中Linen高层 API 的完整入门指南内容以仓库 docs/linen_intro.md 为核心骨架并辅以 flax/linen 源码与 tests/linen 测试进行纵深佐证。读完本文你将掌握如何实例化并调用nn.Moduleinit/apply的完整参数语义、如何用setup()与compact两种方式定义模块、如何用self.param与self.variable管理参数和可变状态、如何区分参数与一般变量集合以及在模块内部嵌套使用jit、remat、vmap、scan等 JAX 变换来组合出可训练、可扩展的模型含多头注意力、LSTM 扫描等实战示例。文档背景与适用前提docs/linen_intro.md最初是 Linen API 的早期预览文档开头带有 CAVEAT PROGRAMMER / alpha API preview 提示。如今该 API 已沉淀为 Flax 的核心稳定接口文档中介绍的概念与调用约定在当前仓库中依然成立init/apply、setup/compact、param/variable、模块内 JAX 变换正是 flax/linen/module.py 与 flax/linen/transforms.py 中实现的核心机制。本文以仓库现状为准将文档中的示例逐一还原并补充源码级细节。环境安装与导入Flax 运行在 JAX 之上需要先安装 JAX含 XLA 编译器再安装 Flax# 升级 JAX / JAXlib !pip install --upgrade -q pip jax jaxlib # 从源码安装最新版 Flax !pip install --upgrade -q githttps://github.com/google/flax.git说明当前仓库本 Flax 仓库即是上述源码安装的对应版本你也可以在本地直接pip install flax使用已发布版本。两种方式下下述 API 用法一致。导入依赖与核心命名空间import functools from typing import Any, Callable, Sequence, Optional import jax from jax import lax, random, numpy as jnp import flax from flax import linen as nnflax.linen常简写为nn是面向对象式的高层 API底层还有一个函数式核心functional core见 flax/coreLinen 模块的变量存储最终落在核心的Scope/VariableDict机制之上。调用模块init 与 apply 的分工与许多框架不同Linen 的Module是真实的对象实例化时传入的是构造参数而不是前向输入。实例化构造参数以下代码创建了一个输出维度为 3 的Dense层model nn.Dense(features3)Dense的完整构造参数可在 flax/linen/linear.py 中查看包括features输出特征数必填use_bias是否加偏置默认Truedtype计算 dtype默认从输入与参数推断param_dtype传给参数初始化器的 dtype默认float32precisionjax.lax.Precision数值精度kernel_init/bias_init权重与偏置的初始化函数默认分别为default_kernel_init与initializers.zeros_init()。init初始化变量参数 状态模块的变量包括参数与其它状态需要在首次调用前初始化。若模块__call__签名为(self, *args, **kwargs)则init的签名为(rngs, *args, **kwargs)# 生成 RNG Key 与假输入 key1, key2 random.split(random.key(0), 2) x random.uniform(key1, (4, 4)) # 传入 key 与假输入得到初始化后的变量 init_variables model.init(key2, x)init返回的是按集合collection分组的变量字典外层键是变量种类如params内层是参数名到数组的映射。对上面的Dense得到的是{params: {kernel: (4, 3), bias: (3,)}}形状的参数树可参考 flax/linen/linear.py 的 docstring 示例。在 flax/linen/module.py 中可以看到init的实现实际上调用了init_with_output即初始化并返回变量丢弃输出其mutable默认值为DenyList(intermediates)。apply用已有变量执行前向apply的签名为(variables, *args, rngsRNGS, mutableMUTABLEKINDS, **kwargs)其中RNGS调用时需要的 RNG例如 dropout。简单模块只需一个 key若模块含多种种类kind的数据则需要传字典如{params: key0, dropout: key1}含 dropout 层的模块。多个 RNG 流由self.make_rng(name)在模块内部按名称索取未提供的名称会回退到params流见 flax/linen/module.py 的 docstringMUTABLEKINDS可选的可变集合名列表例如[batch_stats]表示调用期间会更新 batchnorm 统计量。mutable可为bool/str/listTrue表示所有集合可变见 flax/linen/module.py若指定了可变集合apply返回(输出, 更新后的变量)二元组否则仅返回输出。本例不涉及可变集合直接(variables, input)y model.apply(init_variables, x)调用非__call__方法method 参数如果要对encode/decode等方法而不是__call__执行init/apply需传入methodinit_variables model.init(key2, x, methodencode) y model.apply(init_variables, x, methoddecode)method支持字符串按名查找模块方法、绑定/未绑定的函数对象甚至外部定义的、首个参数接收模块实例的函数详见 flax/linen/module.py 的 docstring 示例其中展示了Transformer的encode用法。定义基础模块两种风格组合子模块setup() 惰性初始化在setup()中声明子模块仍可享受形状推断带来的便利Linen 使用惰性初始化变量只在第一次被使用的位置、以该处的形状信息完成创建见文档 Declaring and using variables 一节以及 flax/linen/module.py 的setup定义。class ExplicitMLP(nn.Module): features: Sequence[int] def setup(self): # 自动处理子模块的 list / dict self.layers [nn.Dense(feat) for feat in self.features] # 单个子模块直接写 # self.layer1 nn.Dense(feat1) def __call__(self, inputs): x inputs for i, lyr in enumerate(self.layers): x lyr(x) if i ! len(self.layers) - 1: x nn.relu(x) return x key1, key2 random.split(random.key(0), 2) x random.uniform(key1, (4, 4)) model ExplicitMLP(features[3, 4, 5]) init_variables model.init(key2, x) y model.apply(init_variables, x) print(initialized parameter shapes:\n, jax.tree_util.tree_map(jnp.shape, flax.core.unfreeze(init_variables))) print(output:\n, y)要点setup()中通过属性赋值注册子模块list、dict 等容器也会被递归识别features是声明在类体中的字段由nn.Module的 dataclass 机制自动生成构造参数。输出可见每层参数形状(4,3)、(3,4)、(4,5)均由输入形状(4,4)推断得出。等价紧凑形式compactcompact装饰器允许在__call__内部内联声明子模块写法更简洁class SimpleMLP(nn.Module): features: Sequence[int] nn.compact def __call__(self, inputs): x inputs for i, feat in enumerate(self.features): x nn.Dense(feat, nameflayers_{i})(x) if i ! len(self.features) - 1: x nn.relu(x) # 名称是可选的 # 默认自动命名规则为 Dense_0, Dense_1, ... # x nn.Dense(feat)(x) return x key1, key2 random.split(random.key(0), 2) x random.uniform(key1, (4, 4)) model SimpleMLP(features[3, 4, 5]) init_variables model.init(key2, x) y model.apply(init_variables, x) print(initialized parameter shapes:\n, jax.tree_util.tree_map(jnp.shape, flax.core.unfreeze(init_variables))) print(output:\n, y)两种写法产生等价的计算图与变量树setup方式把子模块存放在self.layers等属性中compact方式在调用路径上按名称显式name或自动命名记录子模块。仓库中的设计测试 examples/linen_design_test/mlp_explicit.py 与 examples/linen_design_test/mlp_inline.py 正是这两种风格的对照示例。声明和使用变量param 与 variable参数self.param参数是不会被模型内部修改、只由梯度下降更新的变量使用语法self.param(parameter_name, parameter_init_fn, *init_args, **init_kwargs)参数含义parameter_name字符串名称parameter_init_fn接收 RNG key 与任意其它参数的初始化函数即fn(rng, *args)。nn.initializers中的初始化器通常接收rng与shape两个参数其余参数会在初始化时原样传给 init 函数。文档用compact内联实现了一个SimpleDense与仓库 flax/linen/linear.py 中Dense的实现思路一致那里的self.param(kernel, self.kernel_init, (jnp.shape(inputs)[-1], self.features), self.param_dtype)正是同一模式class SimpleDense(nn.Module): features: int kernel_init: Callable nn.initializers.lecun_normal() bias_init: Callable nn.initializers.zeros_init() nn.compact def __call__(self, inputs): kernel self.param(kernel, self.kernel_init, # RNG 隐式传入 (inputs.shape[-1], self.features)) # 形状信息 y lax.dot_general(inputs, kernel, (((inputs.ndim - 1,), (0,)), ((), ())),) bias self.param(bias, self.bias_init, (self.features,)) y y bias return y key1, key2 random.split(random.key(0), 2) x random.uniform(key1, (4, 4)) model SimpleDense(features3) init_variables model.init(key2, x) y model.apply(init_variables, x) print(initialized parameters:\n, init_variables) print(output:\n, y)注意param的 init 函数所需的 RNG key 是隐式传入的来自init时的paramsRNG 流用户只需在init时提供一个 key 即可真正的Dense实现还通过self.promote_dtype统一输入/参数 dtype并支持use_biasFalse见 flax/linen/linear.py。setup 中的参数需显式形状在setup()中声明参数无法享受形状推断必须给出显式形状。文档示例class ExplicitDense(nn.Module): features_in: int # -- 显式输入形状 features: int kernel_init: Callable nn.initializers.lecun_normal() bias_init: Callable nn.initializers.zeros_init() def setup(self): self.kernel self.param(kernel, self.kernel_init, (self.features_in, self.features)) self.bias self.param(bias, self.bias_init, (self.features,)) def __call__(self, inputs): y lax.dot_general(inputs, self.kernel, (((inputs.ndim - 1,), (0,)), ((), ())),) y y self.bias return y key1, key2 random.split(random.key(0), 2) x random.uniform(key1, (4, 4)) model ExplicitDense(features_in4, features3) init_variables model.init(key2, x) y model.apply(init_variables, x) print(initialized parameters:\n, init_variables) print(output:\n, y)一般可变变量self.variable对于会在模型内部被修改的状态batchnorm 移动统计量batch_stats、自回归缓存cache等使用self.variable(variable_kind, variable_name, variable_init_fn, *init_args, **init_kwargs)参数含义variable_kind变量所属的集合名即顶层变量字典中的内层键。例如batch_stats、cache参数也有集合名默认就是paramsvariable_name字符串名称variable_init_fn接收任意参数的初始化函数fn(*args)。注意这里默认不传 RNG若需要 RNG请显式用self.make_rng(variable_kind)提供其余参数在初始化时传给 init 函数。⚠️ 与参数不同self.variable返回的不是常量而是变量引用用myvariable.value读取原始值、myvariable.value new_value写入新值。文档的计数器示例同时演示了has_variable的用法has_variable(col, name)的实现在 flax/linen/module.py用于判断某集合下变量是否存在这里用来区分正在初始化还是已被调用过class Counter(nn.Module): nn.compact def __call__(self): # 检测是否处于初始化阶段的简单模式 is_initialized self.has_variable(counter, count) counter self.variable(counter, count, lambda: jnp.zeros((), jnp.int32)) if is_initialized: counter.value 1 return counter.value key1 random.key(0) model Counter() init_variables model.init(key1) print(initialized variables:\n, init_variables) y, mutated_variables model.apply(init_variables, mutable[counter]) print(mutated variables:\n, mutated_variables) print(output:\n, y)关键点init阶段has_variable返回False变量被创建为 0不执行 1apply时显式传入mutable[counter]counter.value 1生效返回值变为二元组(y, mutated_variables)若apply不传mutable修改会被丢弃集合不可变。综合示例参数 随机层 可变状态文档用一个刻意混合的Block示例演示三者协作可微参数Dense、随机层Dropout、可变状态BatchNorm的运行统计量class Block(nn.Module): features: int training: bool nn.compact def __call__(self, inputs): x nn.Dense(self.features)(inputs) x nn.Dropout(rate0.5)(x, deterministicnot self.training) x nn.BatchNorm(use_running_averagenot self.training)(x) return x key1, key2, key3, key4 random.split(random.key(0), 4) x random.uniform(key1, (3, 4, 4)) model Block(features3, trainingTrue) init_variables model.init({params: key2, dropout: key3}, x) _, init_params flax.core.pop(init_variables, params) # 传入可变集合调用 apply返回 (输出, 更新后的变量) y, mutated_variables model.apply( init_variables, x, rngs{dropout: key4}, mutable[batch_stats]) # 重新组装完整变量真实训练循环中这里还会带上优化器更新后的 params updated_variables flax.core.freeze(dict(paramsinit_params, **mutated_variables)) print(updated variables:\n, updated_variables) print(initialized variable shapes:\n, jax.tree_util.tree_map(jnp.shape, init_variables)) print(output:\n, y) # 用这些变量进行评估推理 eval_model Block(features3, trainingFalse) y eval_model.apply(updated_variables, x) # 无可变集合单返回值 print(eval output:\n, y)这段代码展示了完整的训练-推理数据流init时传入多 RNG 流字典{params: key2, dropout: key3}params流初始化Dense权重dropout流供Dropout使用apply时rngs{dropout: key4}只提供调用时的随机性mutable[batch_stats]声明BatchNorm的移动均值/方差会被更新此时返回(y, mutated_variables)用flax.core.freeze(dict(params..., **mutated_variables))把初始参数与更新的 batch 统计量重新冻结成完整变量树FrozenDict见 flax/core/frozen_dict.py对应真实训练循环中优化器产出新 params 模型产出新 batch_stats的合并动作推理阶段trainingFalseDropout被禁用deterministicTrue、BatchNorm使用运行平均值use_running_averageTrue且无可变集合apply只返回输出。模块内的 JAX 变换Linen 支持把jit、remat、vmap、scan等变换直接施加在模块/方法上且自动处理其中的参数、可变变量与 RNG。这些变换的函数签名见 flax/linen/transforms.py并有对应的单元测试 tests/linen/linen_transforms_test.py如test_jit、test_remat、test_vmap、test_scan等保证其行为正确。JIT编译子模块nn.jit可以编译特定子模块默认编译其__call__class MLP(nn.Module): features: Sequence[int] nn.compact def __call__(self, inputs): x inputs for i, feat in enumerate(self.features): # 对 Module默认是其 __call__做 JIT x nn.jit(nn.Dense)(feat, nameflayers_{i})(x) if i ! len(self.features) - 1: x nn.relu(x) return x key1, key2 random.split(random.key(3), 2) x random.uniform(key1, (4, 4)) model MLP(features[3, 4, 5]) init_variables model.init(key2, x) y model.apply(init_variables, x) print(initialized parameter shapes:\n, jax.tree_util.tree_map(jnp.shape, flax.core.unfreeze(init_variables))) print(output:\n, y)已知 Gotcha目前该装饰器会轻微改变 RNG 流因此 jit 与未 jit 的初始化结果看起来不同测试test_jit_rng_equivalancetests/linen/linen_module_test.py专门验证了 jit 前后 RNG 行为的一致性约定。nn.jit还支持static_argnums、static_argnames、donate_argnums、device、backend等参数并可作用于整个模块类或某个方法见 flax/linen/transforms.py 中jit的定义。Remat以重算换显存对于内存开销大的计算可用nn.remat让反向传播时重新计算模块输出从而省去保存激活值的显存class RematMLP(nn.Module): features: Sequence[int] # 对所有变换既可以标注方法也可以包装已有 Module 类 # 这里我们标注方法。 nn.remat nn.compact def __call__(self, inputs): x inputs for i, feat in enumerate(self.features): x nn.Dense(feat, nameflayers_{i})(x) if i ! len(self.features) - 1: x nn.relu(x) return x key1, key2 random.split(random.key(3), 2) x random.uniform(key1, (4, 4)) model RematMLP(features[3, 4, 5]) init_variables model.init(key2, x) y model.apply(init_variables, x) print(initialized parameter shapes:\n, jax.tree_util.tree_map(jnp.shape, flax.core.unfreeze(init_variables))) print(output:\n, y)同样有 RNG 流的已知 Gotcha。nn.remat在源码中由checkpoint实现别名支持policy、static_argnums等参数仓库还提供remat_scanremat scan 组合用于长序列的显存优化见 flax/linen/transforms.py 与测试test_remat_scan。测试test_remat、test_remat_decorated验证了 remat 前后输出一致tests/linen/linen_transforms_test.py。Vmap模块级向量化nn.vmap把 JAX 的vmap提升到模块层。除 JAX 常规参数外还针对每种变量集合提供轴规则in_axes每个输入参数对应的映射轴整数或Noneout_axes每个输出对应的映射轴整数或Noneaxis_size需要显式指定时的轴大小针对每种 kind 的变量variable_in_axes字典kind → 整数或None指定该集合的输入映射轴variable_out_axes字典kind → 整数或None指定该集合的输出映射轴split_rngs字典RNG-kind → bool指定是否沿轴拆分 RNG。完整签名见 flax/linen/transforms.py 中vmap的定义含axis_name、spmd_axis_name等。文档用从单头无 batch 注意力推导出批量多头注意力的例子展示 vmap 威力class RawDotProductAttention(nn.Module): attn_dropout_rate: float 0.1 train: bool False nn.compact def __call__(self, query, key, value, biasNone, dtypejnp.float32): assert key.ndim query.ndim assert key.ndim value.ndim n query.ndim attn_weights lax.dot_general( query, key, (((n-1,), (n - 1,)), ((), ()))) if bias is not None: attn_weights bias norm_dims tuple(range(attn_weights.ndim // 2, attn_weights.ndim)) attn_weights jax.nn.softmax(attn_weights, axisnorm_dims) attn_weights nn.Dropout(self.attn_dropout_rate)(attn_weights, deterministicnot self.train) attn_weights attn_weights.astype(dtype) contract_dims ( tuple(range(n - 1, attn_weights.ndim)), tuple(range(0, n - 1))) y lax.dot_general( attn_weights, value, (contract_dims, ((), ()))) return y class DotProductAttention(nn.Module): qkv_features: Optional[int] None out_features: Optional[int] None train: bool False nn.compact def __call__(self, inputs_q, inputs_kv, biasNone, dtypejnp.float32): qkv_features self.qkv_features or inputs_q.shape[-1] out_features self.out_features or inputs_q.shape[-1] QKVDense functools.partial( nn.Dense, featuresqkv_features, use_biasFalse, dtypedtype) query QKVDense(namequery)(inputs_q) key QKVDense(namekey)(inputs_kv) value QKVDense(namevalue)(inputs_kv) y RawDotProductAttention(trainself.train)( query, key, value, biasbias, dtypedtype) y nn.Dense(featuresout_features, dtypedtype, nameout)(y) return y class MultiHeadDotProductAttention(nn.Module): qkv_features: Optional[int] None out_features: Optional[int] None batch_axes: Sequence[int] (0,) num_heads: int 1 broadcast_dropout: bool False train: bool False nn.compact def __call__(self, inputs_q, inputs_kv, biasNone, dtypejnp.float32): qkv_features self.qkv_features or inputs_q.shape[-1] out_features self.out_features or inputs_q.shape[-1] # 从单头实现构造多头沿参数轴 0 映射得到 num_heads 组独立参数 Attn nn.vmap(DotProductAttention, in_axes(None, None, None), out_axes2, axis_sizeself.num_heads, variable_axes{params: 0}, split_rngs{params: True, dropout: not self.broadcast_dropout}) # 沿 batch 维度 vmap for axis in reversed(sorted(self.batch_axes)): Attn nn.vmap(Attn, in_axes(axis, axis, axis), out_axesaxis, variable_axes{params: None}, split_rngs{params: False, dropout: False}) # 运行 vmap 后的类 y Attn(qkv_featuresqkv_features // self.num_heads, out_featuresout_features, trainself.train, nameattention)(inputs_q, inputs_kv, bias) return y.mean(axis-2) key1, key2, key3, key4 random.split(random.key(0), 4) x random.uniform(key1, (3, 13, 64)) model functools.partial( MultiHeadDotProductAttention, broadcast_dropoutFalse, num_heads2, batch_axes(0,)) init_variables model(trainFalse).init({params: key2}, x, x) print(initialized parameter shapes:\n, jax.tree_util.tree_map(jnp.shape, flax.core.unfreeze(init_variables))) y model(trainTrue).apply(init_variables, x, x, rngs{dropout: key4}) print(output:\n, y.shape)逐步拆解第一次nn.vmap沿num_heads轴variable_axes{params: 0}让每个 head 拥有独立的参数split_rngs{params: True, dropout: not broadcast_dropout}表示按 head 拆分参数初始化 RNG且每个 head 使用不同的 dropout 掩码第二次nn.vmap沿batch_axes参数在 batch 间共享variable_axes{params: None}RNG 不拆分头数num_heads2时qkv_features // num_heads自动把特征维均分到各头最后mean(axis-2)合并多头输出初始化用model(trainFalse)关闭 dropout前向用model(trainTrue)并显式提供rngs{dropout: key4}。这一先用 vmap 造多头、再用 vmap 批量化的组合正是 Linen 变换可叠加性的直观体现仓库现代版nn.MultiHeadAttention位于 flax/linen/attention.py。Scan沿轴扫描含参数广播与携带变量nn.scan把lax.scan提升到模块层可沿指定轴迭代执行模块并正确处理其参数与可变变量。对每种 kind 的变量需指定变换方式nn.broadcast把该变量种类作为常量广播到所有扫描步各步共享axis:int沿该轴扫描例如每一步拥有独立参数或者通过variable_carry参数指定该变量种类作为携带状态carry跨步传递。此外对被 scan 的变量种类还可指定是否在每一步拆分 RNG。文档示例用nn.scan沿时间轴扫描LSTMCellclass SimpleScan(nn.Module): features: int nn.compact def __call__(self, xs): LSTM nn.scan(nn.LSTMCell, in_axes1, out_axes1, variable_broadcastparams, split_rngs{params: False}) lstm LSTM(self.features, namelstm_cell) dummy_rng random.key(0) input_shape xs[:, 0].shape init_carry lstm.initialize_carry(dummy_rng, input_shape) return lstm(init_carry, xs) key1, key2 random.split(random.key(0), 2) xs random.uniform(key1, (1, 5, 2)) model SimpleScan(2) init_variables model.init(key2, xs) print(initialized parameter shapes:\n, jax.tree_util.tree_map(jnp.shape, flax.core.unfreeze(init_variables))) y model.apply(init_variables, xs) print(output:\n, y)要点in_axes1, out_axes1沿第 1 维时间步扫描xs形状为(batch1, time5, feat2)输出同样保留时间轴variable_broadcastparamsLSTM 的参数在所有时间步共享这正是 RNN 的权重绑定语义因此变量树中只出现一份参数split_rngs{params: False}参数初始化 RNG 不在步间拆分initialize_carry是LSTMCell提供的接口见 flax/linen/recurrent.py用于按输入形状构造初始隐藏状态/记忆状态lstm(init_carry, xs)返回(carry, outputs)二元组。nn.scan的完整参数variable_axes、variable_carry、length、reverse、unroll等见 flax/linen/transforms.py 中scan的定义测试 tests/linen/linen_transforms_test.py 中的test_scan、test_scan_decorated、test_scan_negative_axes覆盖了广播、负轴等边界情况。进阶学习路径从setup/compact/param/variable的完整实现与 docstring 入手flax/linen/module.py全部模块内变换jit、remat、vmap、scan、grad/vjp 等的签名与语义flax/linen/transforms.py常用层源码Dense、Conv、Attention、BatchNorm、Dropout、LSTMCellflax/linen/linear.py、flax/linen/attention.py、flax/linen/normalization.py、flax/linen/stochastic.py、flax/linen/recurrent.py变换行为测试jit/remat/vmap/scan 的正确性与等价性tests/linen/linen_transforms_test.py两种模块定义风格的对照设计测试examples/linen_design_test把本文概念落地到完整训练脚本仓库 examples/mnist/train.py最简分类训练、examples/imagenet/train.py分布式大模型训练、examples/wmt/train.pySeq2Seq 机器翻译以及 examples/seq2seq 中的编码器-解码器示例。至此你已经掌握了 Linen 的全部核心心智模型模块是对象、变量按集合组织、参数与状态分离、变换可作用于模块与方法。基于这套机制你可以把任意 JAX 变换组合进网络内部如本文的多头注意力 vmap 与 LSTM scan并借助 flax.linen 的高层抽象轻松写出既灵活又易于分布式扩展的模型代码。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表