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
本文是 Flax(JAX 生态的神经网络库)中Linen高层 API 的完整入门指南,内容以仓库 docs/linen_intro.md 为核心骨架,并辅以 flax/linen 源码与 tests/linen 测试进行纵深佐证。读完本文,你将掌握:如何实例化并调用nn.Module(init/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 git+https://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/core),Linen 模块的变量存储最终落在核心的Scope/VariableDict机制之上。
调用模块:init 与 apply 的分工
与许多框架不同,Linen 的Module是真实的对象:实例化时传入的是构造参数,而不是前向输入。
实例化:构造参数
以下代码创建了一个输出维度为 3 的Dense层:
model = nn.Dense(features=3)Dense的完整构造参数可在 flax/linen/linear.py 中查看,包括:
features:输出特征数(必填);use_bias:是否加偏置,默认True;dtype:计算 dtype,默认从输入与参数推断;param_dtype:传给参数初始化器的 dtype,默认float32;precision:jax.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, rngs=<RNGS>, mutable=<MUTABLEKINDS>, **kwargs),其中:
<RNGS>:调用时需要的 RNG,例如 dropout。简单模块只需一个 key;若模块含多种"种类"(kind)的数据,则需要传字典,如{'params': key0, 'dropout': key1}(含 dropout 层的模块)。多个 RNG 流由self.make_rng(name)在模块内部按名称索取,未提供的名称会回退到'params'流(见 flax/linen/module.py 的 docstring);<MUTABLEKINDS>:可选的可变集合名列表,例如['batch_stats']表示调用期间会更新 batchnorm 统计量。mutable可为bool/str/list,True表示所有集合可变(见 flax/linen/module.py);- 若指定了可变集合,
apply返回(输出, 更新后的变量)二元组;否则仅返回输出。
本例不涉及可变集合,直接(variables, input):
y = model.apply(init_variables, x)调用非__call__方法:method 参数
如果要对encode/decode等方法而不是__call__执行init/apply,需传入method=:
init_variables = model.init(key2, x, method='encode') y = model.apply(init_variables, x, method='decode')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)推断得出。
等价紧凑形式:@compact
@compact装饰器允许在__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, name=f'layers_{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(features=3) 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时的'params'RNG 流),用户只需在init时提供一个 key 即可;真正的Dense实现还通过self.promote_dtype统一输入/参数 dtype,并支持use_bias=False(见 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_in=4, features=3) 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;参数也有集合名,默认就是'params';variable_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,不执行+= 1;apply时显式传入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(rate=0.5)(x, deterministic=not self.training) x = nn.BatchNorm(use_running_average=not self.training)(x) return x key1, key2, key3, key4 = random.split(random.key(0), 4) x = random.uniform(key1, (3, 4, 4)) model = Block(features=3, training=True) 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(params=init_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(features=3, training=False) 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"的合并动作; - 推理阶段
training=False,Dropout被禁用(deterministic=True)、BatchNorm使用运行平均值(use_running_average=True),且无可变集合,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, name=f'layers_{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_equivalance(tests/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, name=f'layers_{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_scan(remat + 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:每个输入参数对应的映射轴(整数或None);out_axes:每个输出对应的映射轴(整数或None);axis_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, bias=None, dtype=jnp.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, axis=norm_dims) attn_weights = nn.Dropout(self.attn_dropout_rate)(attn_weights, deterministic=not 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, bias=None, dtype=jnp.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, features=qkv_features, use_bias=False, dtype=dtype) query = QKVDense(name='query')(inputs_q) key = QKVDense(name='key')(inputs_kv) value = QKVDense(name='value')(inputs_kv) y = RawDotProductAttention(train=self.train)( query, key, value, bias=bias, dtype=dtype) y = nn.Dense(features=out_features, dtype=dtype, name='out')(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, bias=None, dtype=jnp.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_axes=2, axis_size=self.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_axes=axis, variable_axes={'params': None}, split_rngs={'params': False, 'dropout': False}) # 运行 vmap 后的类 y = Attn(qkv_features=qkv_features // self.num_heads, out_features=out_features, train=self.train, name='attention')(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_dropout=False, num_heads=2, batch_axes=(0,)) init_variables = model(train=False).init({'params': key2}, x, x) print('initialized parameter shapes:\n', jax.tree_util.tree_map(jnp.shape, flax.core.unfreeze(init_variables))) y = model(train=True).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_heads=2时,qkv_features // num_heads自动把特征维均分到各头,最后mean(axis=-2)合并多头输出; - 初始化用
model(train=False)关闭 dropout,前向用model(train=True)并显式提供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沿时间轴扫描LSTMCell:
class SimpleScan(nn.Module): features: int @nn.compact def __call__(self, xs): LSTM = nn.scan(nn.LSTMCell, in_axes=1, out_axes=1, variable_broadcast='params', split_rngs={'params': False}) lstm = LSTM(self.features, name="lstm_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_axes=1, out_axes=1:沿第 1 维(时间步)扫描;xs形状为(batch=1, time=5, feat=2),输出同样保留时间轴;variable_broadcast='params':LSTM 的参数在所有时间步共享(这正是 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、LSTMCell):flax/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.py(Seq2Seq 机器翻译)以及 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),仅供参考