ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

Flax NNX LoRA 参数高效微调详解:LoRA 与 LoRALinear 的源码级实战指南

Flax NNX LoRA 参数高效微调详解:LoRA 与 LoRALinear 的源码级实战指南 Flax NNX LoRA 参数高效微调详解LoRA 与 LoRALinear 的源码级实战指南【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flaxFlax NNX 在flax.nnx.nn.lora模块中提供了两套 LoRALow-Rank Adaptation低秩适配实现独立 LoRA 层LoRA和与nnx.Linear状态完全兼容的LoRALinear。本文以 lora.rst 文档为骨架结合 lora.py 源码、lora_test.py 测试用例以及 surgery.md、bridge_guide.md 指南完整讲解两类的全部构造参数、初始化规则、前向计算与参数分区机制并给出可直接复制的层替换、状态恢复与部分初始化实战方案。读完本文你将能在 JAX 生态中独立完成冻结主干权重、只训练低秩适配矩阵的经典微调工作流。一、NNX 中的 LoRA 设计概览LoRA 的核心思想是对预训练权重矩阵W不直接更新W本身而是引入两个低秩矩阵Afan-in 方向与Bfan-out 方向将权重更新量约束为ΔW A B从而把可训练参数量从in_features × out_features压缩到in_features × lora_rank lora_rank × out_features。NNX 的 lora 模块位于 flax/nnx/nn/lora.py公开暴露三个符号见 flax/nnx/init.pynnx.LoRA独立的 LoRA 适配层可包裹任意已有模块nnx.LoRALinearnnx.Linear的子类在保持原 Linear 状态结构kernel/bias不变的前提下附加 LoRA 分支nnx.LoRAParam专用的参数类型variablelib.Param的子类用于把 LoRA 矩阵与其他参数区分开。二、nnx.LoRA独立的低秩适配层2.1 构造签名与参数说明nnx.LoRA( in_features: int, lora_rank: int, out_features: int, *, base_module: tp.Optional[Module] None, dtype: tp.Optional[Dtype] None, param_dtype: Dtype jnp.float32, a_initializer: Initializer default_a_initializer, # 默认 he_uniform b_initializer: Initializer default_b_initializer, # 默认 zeros lora_param_type: tp.Type[variablelib.Variable] LoRAParam, promote_dtype: PromoteDtypeFn dtypes.promote_dtype, rngs: rnglib.Rngs, a_metadata: tp.Mapping[str, tp.Any] MappingProxyType({}), b_metadata: tp.Mapping[str, tp.Any] MappingProxyType({}), )各参数的作用如下与 lora.py 文档字符串一致参数含义默认值in_features输入特征数必填lora_rankLoRA 低秩维度矩阵 A 的列数 / 矩阵 B 的行数必填out_features输出特征数必填base_module被适配的基础模块前向时若可调用则将其输出与 LoRA 输出相加Nonedtype计算的 dtype默认由输入与参数推断Noneparam_dtype传给参数初始化器的 dtypejnp.float32a_initializerfan-in 矩阵 A 的初始化器he_uniformb_initializerfan-out 矩阵 B 的初始化器zeros零初始化lora_param_typeLoRA 参数的变量类型LoRAParampromote_dtype统一输入与lora_a、lora_b的 dtype 的函数dtypes.promote_dtyperngs随机数源nnx.Rngs必填a_metadata初始化 fan-in 矩阵时的可选元数据字典空字典b_metadata初始化 fan-out 矩阵时的可选元数据字典空字典2.2 内部实现参数形状与默认初始化LoRA在初始化时创建两个参数矩阵lora.pyself.lora_a形状(in_features, lora_rank)默认用he_uniform随机初始化self.lora_b形状(lora_rank, out_features)默认用零初始化器置零。B 初始化为零是关键设计它保证适配开始前A B 0从而 LoRA 分支不改变预训练模型的前向输出——这与 LoRA 原论文的初始化约定一致。模块级默认初始化器定义在 lora.pydefault_a_initializer initializers.he_uniform() default_b_initializer initializers.zeros2.3 前向计算逻辑__call__的完整实现lora.pydef __call__(self, x: jax.Array, **kwargs: tp.Any): x, lora_a, lora_b self.promote_dtype( (x, self.lora_a[...], self.lora_b[...]), dtypeself.dtype ) out x lora_a lora_b if self.base_module is not None: if not callable(self.base_module): raise ValueError(self.base_module must be callable.) out self.base_module(x, **kwargs) return out前向流程分三步用promote_dtype统一x、lora_a、lora_b的 dtype未显式指定dtype时按默认提升规则推断计算低秩分支输出x lora_a lora_b若存在base_module将其前向结果可透传**kwargs累加进去若base_module不可调用则抛出ValueError。注意self.lora_a[...]这种取值写法[...]从Variable包装中取出底层jax.Array是 NNX 中访问变量值的标准方式。2.4 最小可运行示例来自 lora.py 的文档示例完整演示形状关系与包裹已有模块两种用法from flax import nnx import jax, jax.numpy as jnp # 1) 独立 LoRA 层输入 3 维低秩 2输出 4 维 layer nnx.LoRA(3, 2, 4, rngsnnx.Rngs(0)) layer.lora_a.shape # (3, 2) layer.lora_b.shape # (2, 4) # 2) 包裹已有模块base_module 被原样保存并参与前向 linear nnx.Linear(3, 4, rngsnnx.Rngs(0)) wrapper nnx.LoRA(3, 2, 4, base_modulelinear, rngsnnx.Rngs(1)) assert wrapper.base_module linear wrapper.lora_a.shape # (3, 2) wrapper.lora_b.shape # (2, 4) # 3) 前向批量 16、输入特征 3 - 输出特征 4 y layer(jnp.ones((16, 3))) y.shape # (16, 4)三、nnx.LoRALinear与Linear状态兼容的 LoRA 化全连接层3.1 设计目标LoRALinear继承自nnx.Linearlora.py其文档明确指出模型状态结构将与 Linear 兼容The model state structure will be compatible with that of Linear。这意味着它同时拥有kernel、bias来自Linear父类即原始权重结构self.lora内部持有的LoRA子模块包含lora_a、lora_b。因此在不解冻原有权重、只额外添加低秩分支的场景下LoRALinear是比手动组合Linear LoRA更直接的替代品。3.2 构造签名与参数说明nnx.LoRALinear( in_features: int, out_features: int, *, lora_rank: int, # 必填关键字参数 lora_dtype: tp.Optional[Dtype] None, lora_param_dtype: Dtype jnp.float32, a_initializer: Initializer default_a_initializer, b_initializer: Initializer default_b_initializer, lora_param_type: tp.Type[variablelib.Variable] LoRAParam, lora_promote_dtype: PromoteDtypeFn dtypes.promote_dtype, rngs: rnglib.Rngs, a_metadata: Mapping {}, b_metadata: Mapping {}, **kwargs, # 透传给 Linear 的参数如 use_bias、precision、kernel_init、bias_init )与LoRA的差异在于以lora_rank为必填关键字参数且 LoRA 相关的 dtype/初始化参数统一带lora_前缀与Linear自身的dtype、param_dtype、precision等参数通过**kwargs透传互不干扰。3.3 初始化与前向实现初始化时lora.py先调用super().__init__(in_features, out_features, rngsrngs, **kwargs)构建标准的Linearkernel 可选 bias再创建内部LoRA子模块挂到self.lora上。前向实现lora.py相当简洁def __call__(self, x: jax.Array, out_sharding None): y super().__call__(x, out_shardingout_sharding) y self.lora(x) return y即先走标准Linear前向再把 LoRA 低秩分支的输出逐元素累加。这也印证了状态结构兼容 Linear——任意接受nnx.Linear的代码序列化、nnx.split/merge、状态遍历都能直接接受LoRALinear。3.4 文档示例kernel 兼容性验证lora.py 的文档示例验证了LoRALinear与Linear的 kernel 完全一致因为内核参数由同一个初始化流程生成from flax import nnx import jax, jax.numpy as jnp linear nnx.Linear(3, 4, rngsnnx.Rngs(0)) lora_linear nnx.LoRALinear(3, 4, lora_rank2, rngsnnx.Rngs(0)) linear.kernel.shape # (3, 4) lora_linear.kernel.shape # (3, 4) lora_linear.lora.lora_a.shape # (3, 2) jnp.allclose(linear.kernel[...], lora_linear.kernel[...]) # Array(True, dtypebool) y lora_linear(jnp.ones((16, 3))) y.shape # (16, 4)四、LoRAParam与参数分区精准控制可训练子集LoRAParam是variablelib.Param的空子类lora.py它存在的意义是让lora_a、lora_b在 NNX 状态图中拥有独立类型标签。这样就能借助nnx.split精确地把 LoRA 参数与普通参数分开。lora_test.py 的test_lora_param_type展示了完整的分区用法model nnx.LoRA(3, 4, 2, lora_param_typennx.LoRAParam, rngsrngs) _, lora_params, params nnx.split(model, nnx.LoRAParam, nnx.Param) # params 为空说明 LoRAParam 与 Param 被严格区分 assert params {} assert (lora_a in lora_params) and (lora_b in lora_params)反过来把lora_param_type设为nnx.Param则lora_a/lora_b会被划入普通参数集合。这一机制是后续冻结主干、只训练 LoRA的基石优化器只作用于LoRAParam分区即可。值得一提的细节自定义LoRA基模块场景下若 base 模块内部还有其他参数如自定义 kernelLoRAParam分区同样只包含lora_a与lora_b。test_lora_base_modulelora_test.py验证了base_module持有原Linear的kernel、bias且前向结果等于x linear.kernel x lora_a lora_b。五、实战一替换模型中的全连接层层交换最典型的应用是把训练好的模型中某个Linear原位替换成 LoRA 版本替换后立即生效且输出保持不变因为lora_b初始为零。5.1 用LoRA包裹原层test_layer_swap_loralora_test.py演示了无需重新初始化整个模型的做法class MLP(nnx.Module): def __init__(self, dim, rngs: nnx.Rngs): self.linear1 nnx.Linear(dim, dim, rngsrngs) self.linear2 nnx.Linear(dim, dim, rngsrngs) def __call__(self, x): x self.linear1(x) return self.linear2(x) rngs nnx.Rngs(0) model MLP(3, rngsrngs) x jax.random.normal(jax.random.key(0), (1, 3)) y model(x) # 原位替换把 linear2 包进 LoRA原层作为 base_module 继续参与前向 model.linear2 nnx.LoRA(3, 4, 3, base_modulemodel.linear2, rngsrngs) lora_y model(x) np.testing.assert_allclose(y, lora_y) # 替换前后输出一致B 初始为零注意这里nnx.LoRA(3, 4, 3, ...)的lora_rank4低秩矩阵形状为(3, 4)与(4, 3)。5.2 用LoRALinear替换并恢复原权重test_layer_swap_loralinearlora_test.py演示了更精细的做法先用nnx.split抽出原linear2的 kernel/bias 状态再换上新LoRALinear并用nnx.update把旧权重写回去从而零成本保留预训练权重_, state nnx.split(model.linear2) # 保留 linear2 的 kernel 与 bias model.linear2 nnx.LoRALinear(3, 3, lora_rank4, rngsrngs) nnx.update(model.linear2, state) # 恢复原始权重 lora_y model(x) np.testing.assert_allclose(y, lora_y) # 输出不变六、实战二LoRA 与 Linen 桥接多集合导出LoRA 参数类型在 NNX→Linen 桥接时会被自动映射为独立 collection。在 bridge_guide.md 的示例中NNX 模型内部使用nnx.LoRA(din, 3, dout, rngsrngs)作为子模块经bridge.to_linen转换后Linen 变量字典中出现独立的LoRAParamcollectionnnx.register_variable_name(counts, overwriteTrue) class Count(nnx.Variable): pass class NNXMultiCollections(nnx.Module): def __init__(self, din, dout, rngs): self.w nnx.Param(nnx.initializers.lecun_normal()(rngs.params(), (din, dout))) self.lora nnx.LoRA(din, 3, dout, rngsrngs) self.count Count(jnp.array(0)) def __call__(self, x): self.count.value 1 return (x self.w.value) self.lora(x) model bridge.to_linen(NNXMultiCollections, 4, 3) var model.init({params: pkey, dropout: dkey}, x) print(All Linen collections:, list(var.keys())) # All Linen collections: [LoRAParam, params, counts]也就是说LoRAParam变量在桥接后会进入名为LoRAParam的独立 Linen collection与params、counts等并存方便在 Linen 侧按 collection 决定是否冻结或参与优化。七、实战三LoRA 场景下的部分初始化内存优化微调 LoRA 时通常只想随机初始化lora_a/lora_b而 kernel/bias 使用预训练权重。surgery.md 专门给出了两种方案。7.1 朴素部分初始化可能浪费内存直接初始化整个模型再覆盖预训练权重——但nnx.LoRALinear创建时会同时生成kernel、bias、lora_a、lora_b四个数组其中 kernel/bias 随后会被old_state覆盖而成为垃圾对象old_state nnx.state(TwoLayerMLP(4, rngsnnx.Rngs(0))) simple_model nnx.eval_shape(lambda: TwoLayerMLP(4, rngsnnx.Rngs(42))) print(fNumber of jax arrays in memory at start: {len(jax.live_arrays())}) # 下面这行会额外创建 kernel、bias、lora_a、lora_b 四个数组 simple_model.linear1 nnx.LoRALinear(4, 4, lora_rank3, rngsnnx.Rngs(42)) print(fNumber of jax arrays in memory midway: {len(jax.live_arrays())} (4 new created in LoRALinear - kernel, bias, lora_a lora_b)) nnx.update(simple_model, old_state) print(fNumber of jax arrays in memory at end: {len(jax.live_arrays())} (2 discarded - only lora_a lora_b are used in model))7.2 内存高效的部分初始化推荐用nnx.jit包装jax.jit包裹初始化函数编译器会自动跳过最终被丢弃的 kernel/bias 分配old_state nnx.state(TwoLayerMLP(4, rngsnnx.Rngs(0))) nnx.jit(donate_argnums0) def partial_init(old_state, rngs): model TwoLayerMLP(4, rngsrngs) model.linear1 nnx.LoRALinear(4, 4, lora_rank3, rngsrngs) # 只创建 LoRA 分支 nnx.update(model, old_state) # 复用预训练状态 return model print(fNumber of JAX Arrays in memory at start: {len(jax.live_arrays())}) good_model partial_init(old_state, nnx.Rngs(42)) print(fNumber of JAX Arrays in memory at end: {len(jax.live_arrays())} (2 new created - lora_a and lora_b))实测前后对比最终只新增lora_a、lora_b两个数组这正是 LoRA 微调场景下冻结主干、补充轻量可训练参数的内存最优路径。八、dtype 行为验证LoRA支持计算 dtype 与参数 dtype 分离。lora_test.py 的test_dtype验证了model nnx.LoRA(3, 4, 2, dtypejnp.float16, param_dtypejnp.float32, rngsrngs) assert model.lora_a.dtype jnp.float32 # 参数以 float32 存储 y model(jnp.ones((1, 3)).astype(jnp.float32)) assert y.dtype jnp.float16 # 计算/输出为 float16即参数初始化 dtype 由param_dtype控制实际运算与输出 dtype 由dtype控制二者解耦——这对于混合精度训练如 AMP、bf16 推理非常实用。LoRALinear侧对应的参数为lora_param_dtype与lora_dtype语义完全相同。九、可复现性说明LoRA 前向是纯矩阵乘法与加法x lora_a lora_b无随机性随机性只来自初始化阶段的rngs.params()lora.py因此固定nnx.Rngs(seed)即可复现lora_a/lora_b。测试中所有数值断言如np.testing.assert_allclose(y, lora_y)均基于b_initializer默认零初始化导致初始 LoRA 分支输出为零这一特性这也应作为你落地微调实验时的第一个正确性检查替换 LoRA 后、开始训练前模型输出应与替换前严格一致。十、总结与选型建议场景推荐类型理由给任意已有模块不限于 Linear附加 LoRA 分支nnx.LoRA(base_module...)通过base_module包裹任意可调用模块替换模型中现有nnx.Linear并保持状态结构兼容nnx.LoRALinearkernel/bias 结构不变可用nnx.split/nnx.update直接恢复权重在 Linen 桥接中生成独立LoRAParamcollection两者皆可内部均使用LoRAParam自动映射为独立 Linen collection内存敏感的大模型微调初始化LoRALinearnnx.jit部分初始化只分配lora_a/lora_b复用预训练状态进一步阅读LoRA 在前向/状态层面与nnx.split、nnx.update、nnx.jit的完整配合见 surgery.md部分初始化章节与 bridge_guide.mdNNX 到 Linen 多集合转换章节模块级单元测试可运行 lora_test.py 作为正确性参照LoRA与LoRALinear的 API 摘要文档位于 lora.rst。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进